feat: add structured output ladder with bounded feedback retries
This commit is contained in:
@@ -0,0 +1,112 @@
|
||||
"""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
|
||||
@@ -0,0 +1,57 @@
|
||||
"""JsonRepairStrategy: 阶梯②修复(D7/D14;蓝本 VT core/agent/loop.py:341-377)。
|
||||
|
||||
处理链: 围栏剥离 → json_repair → json.loads → 可选 normalize 钩子。
|
||||
业务 schema 相关的变体归一化(如 VT `_normalize_action` 的 DeepSeek 平铺
|
||||
收拢)**不入库**——零业务假设铁律;业务侧经 `normalize` 注入自带函数。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from polygateway.errors import ResultInvalidError
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
try:
|
||||
from json_repair import repair_json
|
||||
except ImportError as _exc: # pragma: no cover - 依赖缺失路径
|
||||
repair_json = None
|
||||
_IMPORT_ERROR = _exc
|
||||
else:
|
||||
_IMPORT_ERROR = None
|
||||
|
||||
# VT loop.py:37 同款围栏正则
|
||||
_CODE_FENCE_RE = re.compile(r"^\s*```(?:json)?\s*\n?|\n?\s*```\s*$")
|
||||
|
||||
|
||||
class JsonRepairStrategy:
|
||||
"""prompt 约定 + 事后修复策略;不改请求体(request_overlay 恒空)。"""
|
||||
|
||||
def __init__(self, normalize: Callable[[Any], Any] | None = None) -> None:
|
||||
if repair_json is None:
|
||||
raise ImportError(
|
||||
"结构化输出需要 json_repair 包: pip install 'polygateway[structured]'"
|
||||
) from _IMPORT_ERROR
|
||||
self._normalize = normalize
|
||||
|
||||
def request_overlay(self, schema: dict[str, Any] | None) -> dict[str, Any]:
|
||||
return {}
|
||||
|
||||
def parse(self, text: str) -> Any:
|
||||
"""修复并解析;失败抛 ResultInvalidError(坏结果 ≠ 坏服务)。"""
|
||||
stripped = _CODE_FENCE_RE.sub("", text).strip()
|
||||
if not stripped:
|
||||
raise ResultInvalidError("结构化输出为空", raw_text=text, repair_error="empty content")
|
||||
try:
|
||||
data = json.loads(repair_json(stripped))
|
||||
except (json.JSONDecodeError, ValueError) as exc:
|
||||
raise ResultInvalidError(
|
||||
"JSON 修复失败", raw_text=text, repair_error=str(exc)
|
||||
) from exc
|
||||
if self._normalize is not None:
|
||||
data = self._normalize(data)
|
||||
return data
|
||||
@@ -0,0 +1,32 @@
|
||||
"""NativeSchemaStrategy: 阶梯①预防(D7/D14)——网关支持时用协议级约束。
|
||||
|
||||
请求侧注入 response_format(有 schema 用 json_schema 严格模式,无 schema
|
||||
退为 json_object);响应侧仍复用修复链兜底——原生约束下模型偶发的围栏/
|
||||
噪声不至于直接判死。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from polygateway.structured.json_repair import JsonRepairStrategy
|
||||
|
||||
|
||||
class NativeSchemaStrategy:
|
||||
"""response_format 注入策略;由装配层按 provider 注册表能力选择。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._repair = JsonRepairStrategy()
|
||||
|
||||
def request_overlay(self, schema: dict[str, Any] | None) -> dict[str, Any]:
|
||||
if schema is None:
|
||||
return {"response_format": {"type": "json_object"}}
|
||||
return {
|
||||
"response_format": {
|
||||
"type": "json_schema",
|
||||
"json_schema": {"name": "structured_output", "strict": True, "schema": schema},
|
||||
}
|
||||
}
|
||||
|
||||
def parse(self, text: str) -> Any:
|
||||
return self._repair.parse(text)
|
||||
Reference in New Issue
Block a user