"""结构化输出阶梯测试(D14/设计 §5): 三档分派、修复链、有界带反馈重问。""" import dataclasses import pytest from pydantic import BaseModel from polygateway.errors import ResultInvalidError from polygateway.middleware.structured import StructuredMW from polygateway.structured.json_repair import JsonRepairStrategy from polygateway.structured.native_schema import NativeSchemaStrategy from polygateway.types import ChatRequest, LLMResponse _MSGS = [{"role": "user", "content": "give json"}] class Verdict(BaseModel): answer: int reason: str def _resp(content): return LLMResponse( content=content, thinking="", model="m", provider="p", prompt_tokens=1, completion_tokens=2, latency_ms=10, ttft_ms=None, max_inter_token_ms=None, cache_hit=False, call_id="cid", source_name="s1", usage_source="measured", ) class ScriptedTerminal: """按脚本逐次返回 content;记录收到的 ChatRequest 序列。""" def __init__(self, contents): self.contents = list(contents) self.requests = [] async def __call__(self, request): self.requests.append(request) return _resp(self.contents.pop(0)) class TestJsonRepairStrategy: @pytest.mark.parametrize( "dirty", [ '```json\n{"answer": 1, "reason": "ok"}\n```', # 围栏 '{"answer": 1, "reason": "ok",}', # 尾逗号 "{'answer': 1, 'reason': 'ok'}", # 单引号 '{"answer": 1, "reason": "ok"', # 缺右括号 ], ) def test_repairs_real_world_dirt(self, dirty): assert JsonRepairStrategy().parse(dirty) == {"answer": 1, "reason": "ok"} def test_unrepairable_raises_result_invalid(self): with pytest.raises(ResultInvalidError) as ei: JsonRepairStrategy().parse("I refuse to answer in JSON.") assert ei.value.raw_text def test_normalize_hook_applied(self): strategy = JsonRepairStrategy(normalize=lambda d: {**d, "tagged": True}) assert strategy.parse('{"a": 1}') == {"a": 1, "tagged": True} def test_request_overlay_empty(self): assert JsonRepairStrategy().request_overlay({"type": "object"}) == {} class TestNativeSchemaStrategy: def test_overlay_with_schema(self): overlay = NativeSchemaStrategy().request_overlay(Verdict.model_json_schema()) rf = overlay["response_format"] assert rf["type"] == "json_schema" assert rf["json_schema"]["schema"]["required"] == ["answer", "reason"] def test_overlay_without_schema_is_json_object(self): assert NativeSchemaStrategy().request_overlay(None) == { "response_format": {"type": "json_object"} } def _mw(**kwargs): defaults = {"strategy": JsonRepairStrategy(), "max_retries": 1, "escalation": None} defaults.update(kwargs) return StructuredMW(**defaults) class TestThreeTiers: async def test_tier_none_passthrough(self): terminal = ScriptedTerminal(["free text"]) resp = await _mw()(ChatRequest(messages=_MSGS), terminal) assert resp.structured_data is None assert terminal.requests[0].overlay == {} async def test_tier_json_repair_only_no_retry(self): terminal = ScriptedTerminal(["not json at all"]) with pytest.raises(ResultInvalidError): await _mw()(ChatRequest(messages=_MSGS, structured="json"), terminal) assert len(terminal.requests) == 1 # "json" 档失败不重问(CHS 语义由 0 档覆盖) async def test_tier_json_success(self): terminal = ScriptedTerminal(['```json\n{"x": 1}\n```']) resp = await _mw()(ChatRequest(messages=_MSGS, structured="json"), terminal) assert resp.structured_data == {"x": 1} async def test_tier_model_full_ladder_success(self): terminal = ScriptedTerminal(['{"answer": 7, "reason": "sure"}']) resp = await _mw()(ChatRequest(messages=_MSGS, structured=Verdict), terminal) assert isinstance(resp.structured_data, Verdict) assert resp.structured_data.answer == 7 class TestFeedbackRetry: async def test_validation_failure_triggers_feedback_reask(self): terminal = ScriptedTerminal( ['{"answer": "not-an-int"}', '{"answer": 7, "reason": "fixed"}'] ) resp = await _mw()(ChatRequest(messages=_MSGS, structured=Verdict), terminal) assert resp.structured_data.answer == 7 assert len(terminal.requests) == 2 reask = terminal.requests[1].messages # 反馈模板: 原 messages + assistant 坏输出 + user 纠错指令(设计 §5 细则 1) assert reask[0] == _MSGS[0] assert reask[1]["role"] == "assistant" and "not-an-int" in reask[1]["content"] assert reask[2]["role"] == "user" and "valid JSON" in reask[2]["content"] assert "answer" in reask[2]["content"] # 校验错误进入反馈 async def test_exhaustion_raises_with_diagnosis(self): terminal = ScriptedTerminal(['{"answer": "a"}', '{"answer": "b"}']) with pytest.raises(ResultInvalidError) as ei: await _mw()(ChatRequest(messages=_MSGS, structured=Verdict), terminal) assert len(terminal.requests) == 2 # 首次 + 1 次重问 assert ei.value.raw_text == '{"answer": "b"}' assert ei.value.validation_errors async def test_zero_retries_is_chs_policy(self): terminal = ScriptedTerminal(['{"answer": "bad"}']) with pytest.raises(ResultInvalidError): await _mw(max_retries=0)(ChatRequest(messages=_MSGS, structured=Verdict), terminal) assert len(terminal.requests) == 1 async def test_reask_escalates_to_native_schema(self): terminal = ScriptedTerminal(['{"answer": "bad"}', '{"answer": 1, "reason": "r"}']) mw = _mw(escalation=NativeSchemaStrategy()) await mw(ChatRequest(messages=_MSGS, structured=Verdict), terminal) assert "response_format" not in terminal.requests[0].overlay # 首发 JsonRepair 无 overlay assert terminal.requests[1].overlay["response_format"]["type"] == "json_schema" async def test_error_feedback_truncated(self): huge_reason = "x" * 5000 terminal = ScriptedTerminal( [f'{{"answer": "{huge_reason}"}}', '{"answer": 1, "reason": "r"}'] ) await _mw()(ChatRequest(messages=_MSGS, structured=Verdict), terminal) feedback = terminal.requests[1].messages[-1]["content"] assert len(feedback) < 2000 # 每条错误截断 200 字符,防 prompt 膨胀 class TestNativeOverlayFirstAttempt: async def test_native_strategy_shapes_first_request(self): terminal = ScriptedTerminal(['{"answer": 1, "reason": "r"}']) mw = _mw(strategy=NativeSchemaStrategy()) await mw(ChatRequest(messages=_MSGS, structured=Verdict), terminal) assert terminal.requests[0].overlay["response_format"]["type"] == "json_schema" async def test_response_immutability_preserved(self): terminal = ScriptedTerminal(['{"answer": 1, "reason": "r"}']) resp = await _mw()(ChatRequest(messages=_MSGS, structured=Verdict), terminal) with pytest.raises(dataclasses.FrozenInstanceError): resp.structured_data = None