Files
PolyGateway/src/polygateway/middleware/structured.py
T

113 lines
4.5 KiB
Python

"""StructuredMW: D14 五级阶梯的编排层(遥测→缓存→**结构化**→重试)。
重问 = 再次调用 call_next(内层重试循环)——天然照过限流/熔断门并逐次
遥测;`ResultInvalidError` 不会进入重试循环的失败计数(RetryMW 对其记
成功后上抛,由本层决定是否带反馈重问)。缓存在本层之外,只固化阶梯
通过的最终响应。
"""
from __future__ import annotations
import dataclasses
from typing import TYPE_CHECKING
from polygateway.errors import ResultInvalidError
if TYPE_CHECKING:
from polygateway.ports import CallNext, StructuredOutputStrategy
from polygateway.types import ChatRequest, LLMResponse
# 反馈模板: 库内常量,英文、零业务词(设计 §5 细则 1)
_FEEDBACK_TEMPLATE = (
"Your previous reply was not valid JSON matching the required schema. "
"Errors: {errors}. Reply with ONLY the corrected JSON object."
)
_MAX_FEEDBACK_ERRORS = 3
_MAX_ERROR_CHARS = 200
def _format_errors(errors: list[str]) -> str:
clipped = [e[:_MAX_ERROR_CHARS] for e in errors[:_MAX_FEEDBACK_ERRORS]]
return "; ".join(clipped) if clipped else "output could not be parsed"
class StructuredMW:
"""三档分派: 不传直通 / "json" 仅修复 / pydantic 模型走完整阶梯。"""
def __init__(
self,
*,
strategy: StructuredOutputStrategy,
max_retries: int = 1,
escalation: StructuredOutputStrategy | None = None,
) -> None:
if max_retries < 0:
raise ValueError("max_structured_retries 不能为负(0 = 不重问转人工)")
self._strategy = strategy
self._max_retries = max_retries
self._escalation = escalation
async def __call__(self, request: ChatRequest, call_next: CallNext) -> LLMResponse:
if request.structured is None:
return await call_next(request)
if request.structured == "json":
response = await call_next(self._shape(request, self._strategy, schema=None))
return dataclasses.replace(
response, structured_data=self._strategy.parse(response.content)
)
return await self._run_ladder(request, call_next)
async def _run_ladder(self, request: ChatRequest, call_next: CallNext) -> LLMResponse:
model_cls = request.structured
schema = model_cls.model_json_schema()
current = self._shape(request, self._strategy, schema=schema)
response = await call_next(current)
reasks = 0
while True:
errors: list[str] = []
repair_error: str | None = None
try:
parsed = self._strategy.parse(response.content)
except ResultInvalidError as exc:
repair_error = exc.repair_error
errors.append(exc.repair_error or "not parseable as JSON")
else:
try:
validated = model_cls.model_validate(parsed)
except ValueError as exc: # pydantic ValidationError 继承 ValueError
errors.append(str(exc))
else:
return dataclasses.replace(response, structured_data=validated)
if reasks >= self._max_retries:
raise ResultInvalidError(
"结构化输出阶梯耗尽",
raw_text=response.content,
repair_error=repair_error,
validation_errors=tuple(errors),
)
reasks += 1
current = self._with_feedback(current, response.content, errors, schema)
response = await call_next(current)
def _shape(
self, request: ChatRequest, strategy: StructuredOutputStrategy, *, schema: dict | None
) -> ChatRequest:
overlay = strategy.request_overlay(schema)
if not overlay:
return request
return dataclasses.replace(request, overlay={**request.overlay, **overlay})
def _with_feedback(
self, current: ChatRequest, bad_content: str, errors: list[str], schema: dict | None
) -> ChatRequest:
"""构造带反馈的重问请求;可升级到原生 schema 策略(设计 §5 细则 2)。"""
messages = [
*current.messages,
{"role": "assistant", "content": bad_content},
{"role": "user", "content": _FEEDBACK_TEMPLATE.format(errors=_format_errors(errors))},
]
reask = dataclasses.replace(current, messages=messages)
if self._escalation is not None:
reask = self._shape(reask, self._escalation, schema=schema)
return reask