From ec4b2faead5311570a1c68224c71887fe5e609a6 Mon Sep 17 00:00:00 2001 From: RKS Date: Mon, 14 Sep 2026 07:59:09 -0400 Subject: [PATCH] fix(core): preserve structured guardrail state during persistence Reuse output normalization for persisted guardrail diagnostics and answers while retaining best-effort fallback when custom serialization fails. Closes #5006 --- src/agents/run_state.py | 18 +++++-- tests/test_run_state.py | 114 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 128 insertions(+), 4 deletions(-) diff --git a/src/agents/run_state.py b/src/agents/run_state.py index d79e0781a4..9ce4a33413 100644 --- a/src/agents/run_state.py +++ b/src/agents/run_state.py @@ -2512,7 +2512,7 @@ def _ensure_json_compatible(value: Any) -> Any: def _serialize_output_value(value: Any) -> Any: - """Convert a tool output value, including containers of models, to plain data. + """Convert an output value, including containers of models, to plain data. ``_ensure_json_compatible`` stringifies anything ``json.dumps`` cannot handle, so Pydantic models and dataclasses nested in containers would otherwise degrade to @@ -2851,6 +2851,16 @@ class _DeserializedFunctionAction: nested_agent_run_state_data: Mapping[str, Any] | None +def _serialize_guardrail_payload(value: Any) -> Any: + """Preserve structured payloads without losing best-effort serialization.""" + try: + value = _serialize_output_value(value) + except Exception: + # Retain the original payload for the existing JSON/string fallback. + pass + return _ensure_json_compatible(value) + + def _serialize_guardrail_results( results: Sequence[InputGuardrailResult | OutputGuardrailResult], *, @@ -2866,11 +2876,11 @@ def _serialize_guardrail_results( }, "output": { "tripwireTriggered": result.output.tripwire_triggered, - "outputInfo": _ensure_json_compatible(result.output.output_info), + "outputInfo": _serialize_guardrail_payload(result.output.output_info), }, } if isinstance(result, OutputGuardrailResult): - entry["agentOutput"] = _ensure_json_compatible(result.agent_output) + entry["agentOutput"] = _serialize_guardrail_payload(result.agent_output) entry["agent"] = _serialize_agent_reference( result.agent, agent_identity_keys_by_id=agent_identity_keys_by_id, @@ -2896,7 +2906,7 @@ def _serialize_tool_guardrail_results( { "guardrail": {"type": type_label, "name": guardrail_name}, "output": { - "outputInfo": _ensure_json_compatible(result.output.output_info), + "outputInfo": _serialize_guardrail_payload(result.output.output_info), "behavior": result.output.behavior, }, } diff --git a/tests/test_run_state.py b/tests/test_run_state.py index cd2daa51b6..6d61d45137 100644 --- a/tests/test_run_state.py +++ b/tests/test_run_state.py @@ -11582,3 +11582,117 @@ async def test_schema_1_13_hosted_mcp_orphaned_call_decisions_require_reapproval ) == "legacy exact denial" ) + + +@pytest.mark.asyncio +async def test_runner_guardrail_models_survive_state_serialization() -> None: + from agents.testing import assistant_message + + class Verdict(BaseModel): + allowed: bool + reason: str + + class Answer(BaseModel): + text: str + + async def check(*args: Any) -> GuardrailFunctionOutput: + return GuardrailFunctionOutput( + output_info=Verdict(allowed=True, reason="approved"), + tripwire_triggered=False, + ) + + agent = Agent( + name="Audit", + model=ScriptedModel([[assistant_message('{"text":"hello"}')]]), + output_type=Answer, + input_guardrails=[InputGuardrail(check)], + output_guardrails=[OutputGuardrail(check)], + ) + result = await Runner.run(agent, "hello", run_config=RunConfig(tracing_disabled=True)) + restored = await RunState.from_json(agent, json.loads(result.to_state().to_string())) + + assert restored._input_guardrail_results[0].output.output_info == { + "allowed": True, + "reason": "approved", + } + assert restored._output_guardrail_results[0].output.output_info == { + "allowed": True, + "reason": "approved", + } + assert restored._output_guardrail_results[0].agent_output == {"text": "hello"} + + +@pytest.mark.asyncio +async def test_tool_guardrail_dataclasses_survive_state_serialization() -> None: + @dataclass + class Verdict: + allowed: bool + reason: str + + verdict = Verdict(allowed=True, reason="approved") + output = ToolGuardrailFunctionOutput( + output_info={"checks": [verdict]}, behavior=AllowBehavior(type="allow") + ) + agent = Agent(name="Audit") + state = make_state(agent, context=RunContextWrapper(context=None)) + state._tool_input_guardrail_results = [ + ToolInputGuardrailResult( + guardrail=ToolInputGuardrail(lambda data: output, name="input"), + output=output, + ) + ] + state._tool_output_guardrail_results = [ + ToolOutputGuardrailResult( + guardrail=ToolOutputGuardrail(lambda data: output, name="output"), + output=output, + ) + ] + restored = await RunState.from_json(agent, json.loads(state.to_string())) + expected = {"checks": [{"allowed": True, "reason": "approved"}]} + assert restored._tool_input_guardrail_results[0].output.output_info == expected + assert restored._tool_output_guardrail_results[0].output.output_info == expected + + +@pytest.mark.asyncio +async def test_guardrail_state_keeps_fallback_when_model_serializer_raises() -> None: + class Diagnostic(BaseModel): + reason: str + + @model_serializer + def serialize(self) -> dict[str, Any]: + raise ValueError("Serializer unavailable.") + + diagnostic = Diagnostic(reason="approved") + output = GuardrailFunctionOutput(output_info=diagnostic, tripwire_triggered=False) + agent = Agent(name="Audit") + state = make_state(agent, context=RunContextWrapper(context=None)) + state._input_guardrail_results = [ + InputGuardrailResult(guardrail=InputGuardrail(lambda *args: output), output=output) + ] + state._output_guardrail_results = [ + OutputGuardrailResult( + guardrail=OutputGuardrail(lambda *args: output), + agent=agent, + agent_output=diagnostic, + output=output, + ) + ] + tool_output = ToolGuardrailFunctionOutput( + output_info=diagnostic, behavior=AllowBehavior(type="allow") + ) + state._tool_input_guardrail_results = [ + ToolInputGuardrailResult( + guardrail=ToolInputGuardrail(lambda data: tool_output), output=tool_output + ) + ] + state._tool_output_guardrail_results = [ + ToolOutputGuardrailResult( + guardrail=ToolOutputGuardrail(lambda data: tool_output), output=tool_output + ) + ] + restored = await RunState.from_json(agent, json.loads(state.to_string())) + assert restored._input_guardrail_results[0].output.output_info == "reason='approved'" + assert restored._output_guardrail_results[0].output.output_info == "reason='approved'" + assert restored._output_guardrail_results[0].agent_output == "reason='approved'" + assert restored._tool_input_guardrail_results[0].output.output_info == "reason='approved'" + assert restored._tool_output_guardrail_results[0].output.output_info == "reason='approved'"