From b928be773979de6e3f6add02abccf333c2793538 Mon Sep 17 00:00:00 2001 From: Sean Date: Tue, 8 Sep 2026 09:20:13 +0800 Subject: [PATCH 1/8] fix(sessions): reconcile pending input appends --- .../packaging/test_run_state_compatibility.py | 4 +- src/agents/result.py | 18 +- .../run_internal/session_persistence.py | 58 +++++- src/agents/run_state.py | 59 ++++++- tests/sandbox/test_docker.py | 2 +- tests/test_run_state.py | 1 + tests/test_run_state_compatibility_corpus.py | 8 +- tests/test_run_state_pending_input.py | 167 +++++++++++++++++- 8 files changed, 294 insertions(+), 23 deletions(-) diff --git a/integration_tests/packaging/test_run_state_compatibility.py b/integration_tests/packaging/test_run_state_compatibility.py index 0d139854c9..b009807cef 100644 --- a/integration_tests/packaging/test_run_state_compatibility.py +++ b/integration_tests/packaging/test_run_state_compatibility.py @@ -5,7 +5,7 @@ import pytest from agents import Agent, RunState -from agents.run_state import SUPPORTED_SCHEMA_VERSIONS +from agents.run_state import CURRENT_SCHEMA_VERSION, SUPPORTED_SCHEMA_VERSIONS from integration_tests._contract_state import ( _deserialize_common_sandbox_session_state, _redaction_observables, @@ -21,7 +21,7 @@ def test_installed_distribution_supports_the_historical_fixture_corpus() -> None: - assert frozenset(SOURCES["versions"]) == SUPPORTED_SCHEMA_VERSIONS + assert frozenset(SOURCES["versions"]) | {CURRENT_SCHEMA_VERSION} == SUPPORTED_SCHEMA_VERSIONS @pytest.mark.parametrize( diff --git a/src/agents/result.py b/src/agents/result.py index 70d48fe3ef..58b857a2f0 100644 --- a/src/agents/result.py +++ b/src/agents/result.py @@ -119,12 +119,25 @@ def _populate_state_from_result( ) -> RunState[Any]: """Populate a RunState with common fields from a RunResult.""" state._current_agent = result.last_agent + source_state = getattr(result, "_state", None) + pending_input_write = ( + source_state._pending_session_write + if isinstance(source_state, RunState) + and source_state._pending_session_write is not None + and "pending_input" in source_state._pending_session_write + else None + ) model_input_items = getattr(result, "_model_input_items", None) - if isinstance(model_input_items, list): + if pending_input_write is not None: + assert isinstance(source_state, RunState) + state._generated_items = list(source_state._generated_items) + state._session_items = list(source_state._session_items) + elif isinstance(model_input_items, list): state._generated_items = list(model_input_items) else: state._generated_items = result.new_items - state._session_items = list(result.new_items) + if pending_input_write is None: + state._session_items = list(result.new_items) snapshot_refs = _state_snapshot_owned_item_refs(result, state._original_input) live_refs = rebase_nested_history_owned_item_refs( state._original_input, @@ -144,7 +157,6 @@ def _populate_state_from_result( state._conversation_id = conversation_id state._previous_response_id = previous_response_id state._auto_previous_response_id = auto_previous_response_id - source_state = getattr(result, "_state", None) if isinstance(source_state, RunState): state._generated_prompt_cache_key = source_state._generated_prompt_cache_key state._pending_input = copy.deepcopy(source_state._pending_input) diff --git a/src/agents/run_internal/session_persistence.py b/src/agents/run_internal/session_persistence.py index de5d3e6b3c..e5ba860fe5 100644 --- a/src/agents/run_internal/session_persistence.py +++ b/src/agents/run_internal/session_persistence.py @@ -121,8 +121,10 @@ async def admit_pending_input( None, store=store, wrapper=wrapper, + resumed_write_state=run_state, + pending_input_snapshot=pending_input, ) - if server_conversation_tracker is None: + elif server_conversation_tracker is None: run_state.clear_pending_input() return admission_items @@ -591,6 +593,7 @@ async def save_result_to_session( store: bool | None = None, wrapper: RunContextWrapper[Any] | None = None, resumed_write_state: RunState | None = None, + pending_input_snapshot: list[TResponseInputItem] | None = None, ) -> int: """ Persist a turn to the session store, keeping track of what was already saved so retries @@ -682,6 +685,22 @@ async def save_result_to_session( item for item in items_to_save if not _is_unpersistable_for_openai_conversation(item) ] + if pending_input_snapshot is not None: + if resumed_write_state is None: + raise UserError("Pending input Session writes require a resumable RunState") + if len(new_items) != len(pending_input_snapshot) or not all( + isinstance(item, InputItem) for item in new_items + ): + raise UserError("Pending input Session writes must contain only admission items") + if ( + resumed_write_state.pending_input[: len(pending_input_snapshot)] + != pending_input_snapshot + ): + raise UserError("Pending input changed before its Session write could be checkpointed") + if not items_to_save: + del resumed_write_state._pending_input[: len(pending_input_snapshot)] + return 0 + if len(items_to_save) == 0: if run_state is not None: run_state._current_turn_persisted_item_count = already_persisted + saved_run_items_count @@ -694,10 +713,15 @@ async def save_result_to_session( "session_id": session.session_id, "items": copy.deepcopy(items_to_save), "before": None, - "persisted_count": ( - resumed_write_state._current_turn_persisted_item_count + saved_run_items_count - ), + "persisted_count": resumed_write_state._current_turn_persisted_item_count + + (0 if pending_input_snapshot is not None else saved_run_items_count), } + if pending_input_snapshot is not None: + resumed_write_state._pending_session_write["pending_input"] = copy.deepcopy( + new_items_as_input + ) + resumed_write_state._generated_items.extend(new_items) + resumed_write_state._session_items.extend(new_items) await resume_pending_session_write( resumed_write_state, session, @@ -817,6 +841,30 @@ def digests(items: Sequence[TResponseInputItem]) -> list[str]: for item in items ] + pending_input = pending.get("pending_input") + if pending_input is not None: + expected_pending_items = deduplicate_input_items_preferring_latest(pending_input) + if isinstance(session, OpenAIConversationsSession): + expected_pending_items = [ + _sanitize_openai_conversation_item(item) for item in expected_pending_items + ] + expected_pending_items = [ + item + for item in expected_pending_items + if not _is_unpersistable_for_openai_conversation(item) + ] + if [digest_input_item(item) for item in expected_pending_items] != [ + digest_input_item(item) for item in pending["items"] + ]: + raise UserError( + "Cannot reconcile the pending Session write: its staged input batch changed." + ) + prefix = run_state._pending_input[: len(pending_input)] + if len(prefix) != len(pending_input) or [digest_input_item(item) for item in prefix] != [ + digest_input_item(item) for item in pending_input + ]: + raise UserError("Cannot reconcile the pending Session write: its staged input changed.") + run_state._session_write_in_progress = True try: before = pending["before"] @@ -852,6 +900,8 @@ def digests(items: Sequence[TResponseInputItem]) -> list[str]: if append: # Backends may retain or transform their input; the durable checkpoint stays detached. await _session_add_items(session, copy.deepcopy(pending["items"]), wrapper=wrapper) + if pending_input is not None: + del run_state._pending_input[: len(pending_input)] run_state._current_turn_persisted_item_count = pending["persisted_count"] run_state._pending_session_write = None finally: diff --git a/src/agents/run_state.py b/src/agents/run_state.py index d79e0781a4..b79d57674f 100644 --- a/src/agents/run_state.py +++ b/src/agents/run_state.py @@ -36,7 +36,7 @@ ProgramOutput, ) from pydantic import BaseModel, StringConstraints, TypeAdapter, ValidationError -from typing_extensions import TypedDict, TypeVar +from typing_extensions import NotRequired, TypedDict, TypeVar from ._run_state_agent_identity import ( _build_agent_identity_keys_by_id, @@ -174,6 +174,7 @@ class _PendingSessionWrite(TypedDict): items: list[TResponseInputItem] before: list[str] | None persisted_count: int + pending_input: NotRequired[list[TResponseInputItem]] def _default_run_state_validation_error( @@ -190,7 +191,7 @@ def _default_run_state_validation_error( # 3. to_json() always emits CURRENT_SCHEMA_VERSION. # 4. Forward compatibility is intentionally fail-fast (older SDKs reject newer or unsupported # versions). -CURRENT_SCHEMA_VERSION = "1.17" +CURRENT_SCHEMA_VERSION = "1.18" _PROGRAMMATIC_TOOL_CALLING_MIN_SCHEMA_VERSION = "1.13" _HOSTED_MCP_APPROVALS_MIN_SCHEMA_VERSION = "1.14" _CURRENT_RESPONSE_OWNERSHIP_MIN_SCHEMA_VERSION = "1.17" @@ -229,6 +230,7 @@ def _default_run_state_validation_error( "Persists Docker container labels and current-response generated-item ownership across " "resume flows, including pending resumed Session writes and terminal-unrecoverable runs." ), + "1.18": "Persists ownership of pending input during unresolved Session writes.", } SUPPORTED_SCHEMA_VERSIONS = frozenset(SCHEMA_VERSION_SUMMARIES) @@ -1020,6 +1022,13 @@ def add_input(self, input: str | list[TResponseInputItem]) -> None: def clear_pending_input(self) -> None: """Remove all input staged for the next resumed model call.""" + if ( + self._pending_session_write is not None + and "pending_input" in self._pending_session_write + ): + raise UserError( + "Cannot clear pending input while its Session write is awaiting reconciliation" + ) self._pending_input = [] def get_interruptions(self) -> list[ToolApprovalItem]: @@ -4128,10 +4137,12 @@ async def _build_run_state_from_json( pending_input_raw = state_json.get("pending_input", []) if not isinstance(pending_input_raw, list): raise validation_error_factory("Run state pending_input must be a list", UserError) - state._pending_input = cast( - list[TResponseInputItem], - [dict(item) if isinstance(item, Mapping) else item for item in pending_input_raw], - ) + try: + state._pending_input = [ + _HANDOFF_OUTPUT_ADAPTER.validate_python(item) for item in pending_input_raw + ] + except ValidationError: + raise validation_error_factory("Run state pending_input is invalid", UserError) from None state._model_responses = _deserialize_model_responses(state_json.get("model_responses", [])) serialized_generated_items = state_json.get("generated_items", []) state._generated_items, generated_source_indexes = _deserialize_items_with_source_indexes( @@ -4372,15 +4383,43 @@ async def _build_run_state_from_json( if pending_write is not None: from .run_internal.run_steps import NextStepInterruption, NextStepRunAgain + pending_input_write = ( + pending_write.get("pending_input") if isinstance(pending_write, dict) else None + ) + pending_write_items = ( + pending_write.get("items") if isinstance(pending_write, dict) else None + ) + try: + validated_pending_write_items = ( + [_HANDOFF_OUTPUT_ADAPTER.validate_python(item) for item in pending_write_items] + if isinstance(pending_write_items, list) + else None + ) + validated_pending_input_write = ( + [_HANDOFF_OUTPUT_ADAPTER.validate_python(item) for item in pending_input_write] + if isinstance(pending_input_write, list) + else None + ) + except ValidationError: + validated_pending_write_items = None + validated_pending_input_write = None + valid_pending_input = pending_input_write is None or ( + (schema_major, schema_minor) >= (1, 18) and bool(validated_pending_input_write) + ) if ( (schema_major, schema_minor) < (1, 17) or not isinstance(state._current_step, NextStepRunAgain | NextStepInterruption) or not isinstance(pending_write, dict) - or set(pending_write) != {"session_id", "items", "before", "persisted_count"} + or set(pending_write) + not in ( + {"session_id", "items", "before", "persisted_count"}, + {"session_id", "items", "before", "persisted_count", "pending_input"}, + ) or not isinstance(pending_write.get("session_id"), str) - or not isinstance(pending_write.get("items"), list) + or not isinstance(pending_write_items, list) or not pending_write["items"] or not all(isinstance(item, dict) for item in pending_write["items"]) + or not validated_pending_write_items or ( pending_write.get("before") is not None and ( @@ -4390,9 +4429,13 @@ async def _build_run_state_from_json( ) or type(pending_write.get("persisted_count")) is not int or pending_write["persisted_count"] < 0 + or not valid_pending_input ): raise validation_error_factory("Run state pending Session write is invalid", UserError) state._pending_session_write = copy.deepcopy(cast(_PendingSessionWrite, pending_write)) + state._pending_session_write["items"] = validated_pending_write_items + if validated_pending_input_write is not None: + state._pending_session_write["pending_input"] = validated_pending_input_write terminal_unrecoverable = state_json.get("terminal_unrecoverable") if terminal_unrecoverable is not None: # An older label never wrote this marker, so honoring one would let a snapshot claim a diff --git a/tests/sandbox/test_docker.py b/tests/sandbox/test_docker.py index e4c7cc812f..d32a670659 100644 --- a/tests/sandbox/test_docker.py +++ b/tests/sandbox/test_docker.py @@ -1926,7 +1926,7 @@ async def test_docker_labels_roundtrip_through_run_state() -> None: serialized = run_state.to_json() restored = await RunState.from_json(agent, serialized) - assert serialized["$schemaVersion"] == CURRENT_SCHEMA_VERSION == "1.17" + assert serialized["$schemaVersion"] == CURRENT_SCHEMA_VERSION assert restored._sandbox is not None restored_session_state = restored._sandbox["session_state"] assert isinstance(restored_session_state, dict) diff --git a/tests/test_run_state.py b/tests/test_run_state.py index cd2daa51b6..4505ac8939 100644 --- a/tests/test_run_state.py +++ b/tests/test_run_state.py @@ -8611,6 +8611,7 @@ def test_supported_schema_versions_match_released_boundary(self): "1.14", "1.15", "1.16", + "1.17", CURRENT_SCHEMA_VERSION, } ) diff --git a/tests/test_run_state_compatibility_corpus.py b/tests/test_run_state_compatibility_corpus.py index 52996188cc..37314afca7 100644 --- a/tests/test_run_state_compatibility_corpus.py +++ b/tests/test_run_state_compatibility_corpus.py @@ -19,7 +19,7 @@ from agents import Agent, RunState, UserError from agents.run_context import RunContextWrapper -from agents.run_state import SUPPORTED_SCHEMA_VERSIONS +from agents.run_state import CURRENT_SCHEMA_VERSION, SUPPORTED_SCHEMA_VERSIONS from agents.sandbox.entries.mounts.patterns import FuseMountConfig from integration_tests._contract_state import ( _deserialize_common_sandbox_session_state, @@ -164,10 +164,12 @@ def test_historical_state_comparison_preserves_json_scalar_types() -> None: def test_historical_fixture_corpus_matches_supported_schema_versions() -> None: assert SOURCES["baseline"] == "v0.19.4" - assert frozenset(SOURCES["versions"]) == SUPPORTED_SCHEMA_VERSIONS + assert frozenset(SOURCES["versions"]) | {CURRENT_SCHEMA_VERSION} == SUPPORTED_SCHEMA_VERSIONS assert all(entry["commit"] for entry in SOURCES["versions"].values()) assert {entry["version"] for entry in SOURCES["features"]} == { - version for version in SUPPORTED_SCHEMA_VERSIONS if version not in {"1.0", "1.1"} + version + for version in SUPPORTED_SCHEMA_VERSIONS + if version not in {"1.0", "1.1", CURRENT_SCHEMA_VERSION} } assert {entry["provenance"] for entry in SOURCES["features"]} == { "historical_writer", diff --git a/tests/test_run_state_pending_input.py b/tests/test_run_state_pending_input.py index fdf1d2cd71..c07e0ceacc 100644 --- a/tests/test_run_state_pending_input.py +++ b/tests/test_run_state_pending_input.py @@ -1,7 +1,8 @@ from __future__ import annotations +import copy import json -from typing import Any, cast +from typing import Any, Literal, cast import pytest from openai.types.responses.response_computer_tool_call import ( @@ -14,6 +15,7 @@ from agents.guardrail import GuardrailFunctionOutput, InputGuardrail from agents.items import ModelResponse, TResponseInputItem from agents.lifecycle import AgentHooks, RunHooks +from agents.memory import OpenAIConversationsSession, Session from agents.run import CallModelData, ModelInputData from agents.run_context import RunContextWrapper from agents.run_internal.oai_conversation import OpenAIServerConversationTracker @@ -29,6 +31,40 @@ from .utils.simple_session import SimpleListSession +class _PendingInputWriteFailureSession(SimpleListSession): + def __init__(self) -> None: + super().__init__() + self.failure: Literal["before", "after"] | None = None + self.error = RuntimeError("pending input Session append failed") + + async def add_items(self, items: list[TResponseInputItem]) -> None: + failure, self.failure = self.failure, None + if failure == "before": + raise self.error + await super().add_items(items) + if failure == "after": + raise self.error + + +class _RecordingConversationsSession(OpenAIConversationsSession): + def __init__(self) -> None: + self._session_id = "test" + self.items: list[TResponseInputItem] = [] + self.failure: Literal["before", "after"] | None = None + self.error = RuntimeError("conversation append failed") + + async def get_items(self, limit: int | None = None) -> list[TResponseInputItem]: + return list(self.items if limit is None else self.items[-limit:]) + + async def add_items(self, items: list[TResponseInputItem]) -> None: + failure, self.failure = self.failure, None + if failure == "before": + raise self.error + self.items.extend(copy.deepcopy(items)) + if failure == "after": + raise self.error + + def _item_type(item: TResponseInputItem) -> str | None: if not isinstance(item, dict): return getattr(item, "type", None) @@ -52,7 +88,7 @@ def _message_text(item: TResponseInputItem) -> str | None: async def _make_after_turn_state( *, - session: SimpleListSession | None = None, + session: Session | None = None, auto_previous_response_id: bool = False, ) -> tuple[ScriptedModel, Agent[Any], RunState[Any], list[str]]: calls: list[str] = [] @@ -89,6 +125,22 @@ def record_destination(destination: str) -> str: return model, agent, state, calls +async def _resume_pending_input_state( + agent: Agent[Any], + state: RunState[Any], + session: Session, + *, + streamed: bool, +) -> Any: + run_config = RunConfig(tracing_disabled=True) + if not streamed: + return await Runner.run(agent, state, session=session, run_config=run_config) + result = Runner.run_streamed(agent, state, session=session, run_config=run_config) + async for _event in result.stream_events(): + pass + return result + + @pytest.mark.asyncio async def test_pending_input_preserves_order_and_serialization_round_trips() -> None: agent = Agent(name="assistant") @@ -168,6 +220,117 @@ async def test_after_turn_resume_admits_input_after_tool_output_exactly_once() - terminal_state.add_input("Too late") +@pytest.mark.asyncio +@pytest.mark.parametrize("failing_streamed", [False, True]) +@pytest.mark.parametrize("recovery_streamed", [False, True]) +@pytest.mark.parametrize("round_trip", [False, True]) +@pytest.mark.parametrize("failure", ["before", "after"]) +async def test_pending_input_session_append_reconciles_once_without_rerunning_guardrails( + failing_streamed: bool, + recovery_streamed: bool, + round_trip: bool, + failure: Literal["before", "after"], +) -> None: + session = _PendingInputWriteFailureSession() + model, agent, state, _calls = await _make_after_turn_state(session=session) + guardrail_calls = 0 + + def inspect_pending_input( + _context: RunContextWrapper[Any], + _agent: Agent[Any], + _input: str | list[TResponseInputItem], + ) -> GuardrailFunctionOutput: + nonlocal guardrail_calls + guardrail_calls += 1 + return GuardrailFunctionOutput(output_info="accepted", tripwire_triggered=False) + + agent.input_guardrails = [InputGuardrail(guardrail_function=inspect_pending_input)] + state.add_input("Late input") + model.enqueue([get_text_message("Recovered")]) + model_calls_before = len(model.calls) + session.failure = failure + + with pytest.raises(RuntimeError, match="pending input Session append failed"): + if failing_streamed: + failed_result = Runner.run_streamed( + agent, state, session=session, run_config=RunConfig(tracing_disabled=True) + ) + async for _event in failed_result.stream_events(): + pass + state = failed_result.to_state() + else: + await Runner.run( + agent, state, session=session, run_config=RunConfig(tracing_disabled=True) + ) + + assert guardrail_calls == 1 + assert len(model.calls) == model_calls_before + assert state._pending_session_write is not None + assert [_message_text(item) for item in state.pending_input] == ["Late input"] + if round_trip: + state = await RunState.from_json(agent, state.to_json()) + + result = await _resume_pending_input_state(agent, state, session, streamed=recovery_streamed) + + assert result.final_output == "Recovered" + assert len(model.calls) == model_calls_before + 1 + assert guardrail_calls == 1 + assert state.pending_input == [] + assert state._pending_session_write is None + assert [_message_text(item) for item in await session.get_items()].count("Late input") == 1 + + +@pytest.mark.asyncio +async def test_pending_input_session_checkpoint_rejects_malformed_input_before_resume() -> None: + session = _PendingInputWriteFailureSession() + model, agent, state, _calls = await _make_after_turn_state(session=session) + state.add_input("Late input") + model.enqueue([get_text_message("Recovered")]) + session.failure = "before" + + with pytest.raises(RuntimeError, match="pending input Session append failed"): + await Runner.run(agent, state, session=session, run_config=RunConfig(tracing_disabled=True)) + + payload = state.to_json() + malformed = {"role": ["user"]} + payload["pending_input"] = [malformed] + pending_write = cast(dict[str, Any], payload["pending_session_write"]) + pending_write["pending_input"] = [malformed] + pending_write["items"] = [malformed] + + with pytest.raises(RuntimeError, match="Error details are redacted"): + await RunState.from_json(agent, payload) + + +@pytest.mark.asyncio +async def test_pending_input_conversations_session_reconciles_sanitized_message_id() -> None: + session = _RecordingConversationsSession() + model, agent, state, _calls = await _make_after_turn_state(session=session) + state.add_input( + [ + { + "id": "user-id", + "type": "message", + "role": "user", + "content": "Late input", + } + ] + ) + model.enqueue([get_text_message("Recovered")]) + session.failure = "after" + + with pytest.raises(RuntimeError, match="conversation append failed"): + await Runner.run(agent, state, session=session, run_config=RunConfig(tracing_disabled=True)) + + state = await RunState.from_json(agent, state.to_json()) + result = await Runner.run( + agent, state, session=session, run_config=RunConfig(tracing_disabled=True) + ) + + assert result.final_output == "Recovered" + assert [_message_text(item) for item in session.items].count("Late input") == 1 + + @pytest.mark.asyncio async def test_streamed_resume_matches_pending_input_ordering() -> None: model, agent, state, calls = await _make_after_turn_state() From 1d5fcb56cdc45d7e35d42512a92517fe424a1ea4 Mon Sep 17 00:00:00 2001 From: Sean Date: Tue, 8 Sep 2026 09:20:13 +0800 Subject: [PATCH 2/8] fix(sessions): reconcile pending input appends --- .../packaging/test_run_state_compatibility.py | 4 +- src/agents/result.py | 18 +- .../run_internal/session_persistence.py | 156 ++++++++++-- src/agents/run_state.py | 69 ++++- tests/sandbox/test_docker.py | 2 +- tests/test_run_impl_resume_paths.py | 26 ++ tests/test_run_state.py | 1 + tests/test_run_state_compatibility_corpus.py | 8 +- tests/test_run_state_pending_input.py | 237 +++++++++++++++++- 9 files changed, 488 insertions(+), 33 deletions(-) diff --git a/integration_tests/packaging/test_run_state_compatibility.py b/integration_tests/packaging/test_run_state_compatibility.py index 0d139854c9..b009807cef 100644 --- a/integration_tests/packaging/test_run_state_compatibility.py +++ b/integration_tests/packaging/test_run_state_compatibility.py @@ -5,7 +5,7 @@ import pytest from agents import Agent, RunState -from agents.run_state import SUPPORTED_SCHEMA_VERSIONS +from agents.run_state import CURRENT_SCHEMA_VERSION, SUPPORTED_SCHEMA_VERSIONS from integration_tests._contract_state import ( _deserialize_common_sandbox_session_state, _redaction_observables, @@ -21,7 +21,7 @@ def test_installed_distribution_supports_the_historical_fixture_corpus() -> None: - assert frozenset(SOURCES["versions"]) == SUPPORTED_SCHEMA_VERSIONS + assert frozenset(SOURCES["versions"]) | {CURRENT_SCHEMA_VERSION} == SUPPORTED_SCHEMA_VERSIONS @pytest.mark.parametrize( diff --git a/src/agents/result.py b/src/agents/result.py index 70d48fe3ef..58b857a2f0 100644 --- a/src/agents/result.py +++ b/src/agents/result.py @@ -119,12 +119,25 @@ def _populate_state_from_result( ) -> RunState[Any]: """Populate a RunState with common fields from a RunResult.""" state._current_agent = result.last_agent + source_state = getattr(result, "_state", None) + pending_input_write = ( + source_state._pending_session_write + if isinstance(source_state, RunState) + and source_state._pending_session_write is not None + and "pending_input" in source_state._pending_session_write + else None + ) model_input_items = getattr(result, "_model_input_items", None) - if isinstance(model_input_items, list): + if pending_input_write is not None: + assert isinstance(source_state, RunState) + state._generated_items = list(source_state._generated_items) + state._session_items = list(source_state._session_items) + elif isinstance(model_input_items, list): state._generated_items = list(model_input_items) else: state._generated_items = result.new_items - state._session_items = list(result.new_items) + if pending_input_write is None: + state._session_items = list(result.new_items) snapshot_refs = _state_snapshot_owned_item_refs(result, state._original_input) live_refs = rebase_nested_history_owned_item_refs( state._original_input, @@ -144,7 +157,6 @@ def _populate_state_from_result( state._conversation_id = conversation_id state._previous_response_id = previous_response_id state._auto_previous_response_id = auto_previous_response_id - source_state = getattr(result, "_state", None) if isinstance(source_state, RunState): state._generated_prompt_cache_key = source_state._generated_prompt_cache_key state._pending_input = copy.deepcopy(source_state._pending_input) diff --git a/src/agents/run_internal/session_persistence.py b/src/agents/run_internal/session_persistence.py index de5d3e6b3c..7a378a1c6a 100644 --- a/src/agents/run_internal/session_persistence.py +++ b/src/agents/run_internal/session_persistence.py @@ -121,8 +121,10 @@ async def admit_pending_input( None, store=store, wrapper=wrapper, + resumed_write_state=run_state, + pending_input_snapshot=pending_input, ) - if server_conversation_tracker is None: + elif server_conversation_tracker is None: run_state.clear_pending_input() return admission_items @@ -591,6 +593,7 @@ async def save_result_to_session( store: bool | None = None, wrapper: RunContextWrapper[Any] | None = None, resumed_write_state: RunState | None = None, + pending_input_snapshot: list[TResponseInputItem] | None = None, ) -> int: """ Persist a turn to the session store, keeping track of what was already saved so retries @@ -682,6 +685,22 @@ async def save_result_to_session( item for item in items_to_save if not _is_unpersistable_for_openai_conversation(item) ] + if pending_input_snapshot is not None: + if resumed_write_state is None: + raise UserError("Pending input Session writes require a resumable RunState") + if len(new_items) != len(pending_input_snapshot) or not all( + isinstance(item, InputItem) for item in new_items + ): + raise UserError("Pending input Session writes must contain only admission items") + if ( + resumed_write_state.pending_input[: len(pending_input_snapshot)] + != pending_input_snapshot + ): + raise UserError("Pending input changed before its Session write could be checkpointed") + if not items_to_save: + del resumed_write_state._pending_input[: len(pending_input_snapshot)] + return 0 + if len(items_to_save) == 0: if run_state is not None: run_state._current_turn_persisted_item_count = already_persisted + saved_run_items_count @@ -690,14 +709,30 @@ async def save_result_to_session( if resumed_write_state is not None: if resumed_write_state._pending_session_write is not None: raise UserError("Resolve the pending Session write before saving another batch") + if isinstance(session, OpenAIConversationsSession): + try: + pending_session_id = session.session_id + except ValueError: + # Conversations sessions create their ID lazily. A zero-item read follows the + # public initialization path without appending history, so the ID is checkpointed + # before the fallible write begins. + await _session_get_items(session, limit=0, wrapper=wrapper) + pending_session_id = session.session_id + else: + pending_session_id = session.session_id resumed_write_state._pending_session_write = { - "session_id": session.session_id, + "session_id": pending_session_id, "items": copy.deepcopy(items_to_save), "before": None, - "persisted_count": ( - resumed_write_state._current_turn_persisted_item_count + saved_run_items_count - ), + "persisted_count": resumed_write_state._current_turn_persisted_item_count + + (0 if pending_input_snapshot is not None else saved_run_items_count), } + if pending_input_snapshot is not None: + resumed_write_state._pending_session_write["pending_input"] = copy.deepcopy( + new_items_as_input + ) + resumed_write_state._generated_items.extend(new_items) + resumed_write_state._session_items.extend(new_items) await resume_pending_session_write( resumed_write_state, session, @@ -807,7 +842,7 @@ async def resume_pending_session_write( if session is None or session.session_id != pending["session_id"]: raise UserError("Resume the pending Session write with the original Session and session ID") - def digests(items: Sequence[TResponseInputItem]) -> list[str]: + def legacy_digests(items: Sequence[TResponseInputItem]) -> list[str]: return [ hashlib.sha256( _fingerprint_or_repr( @@ -817,6 +852,37 @@ def digests(items: Sequence[TResponseInputItem]) -> list[str]: for item in items ] + def appended_digests(items: Sequence[TResponseInputItem]) -> list[str]: + if not isinstance(session, OpenAIConversationsSession): + return legacy_digests(items) + return legacy_digests( + [_canonicalize_openai_conversation_item_for_reconciliation(item) for item in items] + ) + + pending_input = pending.get("pending_input") + if pending_input is not None: + expected_pending_items = deduplicate_input_items_preferring_latest(pending_input) + if isinstance(session, OpenAIConversationsSession): + expected_pending_items = [ + _sanitize_openai_conversation_item(item) for item in expected_pending_items + ] + expected_pending_items = [ + item + for item in expected_pending_items + if not _is_unpersistable_for_openai_conversation(item) + ] + if [digest_input_item(item) for item in expected_pending_items] != [ + digest_input_item(item) for item in pending["items"] + ]: + raise UserError( + "Cannot reconcile the pending Session write: its staged input batch changed." + ) + prefix = run_state._pending_input[: len(pending_input)] + if len(prefix) != len(pending_input) or [digest_input_item(item) for item in prefix] != [ + digest_input_item(item) for item in pending_input + ]: + raise UserError("Cannot reconcile the pending Session write: its staged input changed.") + run_state._session_write_in_progress = True try: before = pending["before"] @@ -825,33 +891,49 @@ def digests(items: Sequence[TResponseInputItem]) -> list[str]: tail = await _session_get_items( session, limit=len(pending["items"]) + 1, wrapper=wrapper ) - pending["before"] = digests(tail) + pending["before"] = legacy_digests(tail) append = True else: - expected = before + digests(pending["items"]) + expected_length = len(before) + len(pending["items"]) committed_generation: int | None = None get_with_generation = getattr(session, "_get_items_with_generation", None) if wrapper is not None and callable(get_with_generation): tail, committed_generation = await _call_session_method( get_with_generation, - limit=len(expected), + limit=expected_length, ) else: - tail = await _session_get_items(session, limit=len(expected), wrapper=wrapper) - observed = digests(tail) - committed = observed == expected + tail = await _session_get_items(session, limit=expected_length, wrapper=wrapper) + observed = legacy_digests(tail) + committed = ( + len(tail) == expected_length + and observed[: len(before)] == before + and appended_digests(tail[len(before) :]) == appended_digests(pending["items"]) + ) unchanged = observed[-len(before) :] == before if before else not observed - if committed == unchanged: + if committed: + before_was_complete = len(before) < len(pending["items"]) + 1 + if unchanged and not before_was_complete: + raise UserError( + "Cannot reconcile the pending Session write: history changed or is " + "ambiguous. Repair the original Session before resuming; do not rerun " + "the completed tool." + ) + append = False + elif unchanged: + append = True + else: raise UserError( "Cannot reconcile the pending Session write: history changed or is ambiguous. " "Repair the original Session before resuming; do not rerun the completed tool." ) - append = unchanged if committed and committed_generation is not None and wrapper is not None: wrapper._session_compaction_generation = committed_generation # type: ignore[attr-defined] if append: # Backends may retain or transform their input; the durable checkpoint stays detached. await _session_add_items(session, copy.deepcopy(pending["items"]), wrapper=wrapper) + if pending_input is not None: + del run_state._pending_input[: len(pending_input)] run_state._current_turn_persisted_item_count = pending["persisted_count"] run_state._pending_session_write = None finally: @@ -1064,6 +1146,52 @@ def _sanitize_openai_conversation_item(item: TResponseInputItem) -> TResponseInp return item +def _canonicalize_openai_conversation_item_for_reconciliation( + item: TResponseInputItem, +) -> TResponseInputItem: + """Normalize Conversations API response defaults for lost-ack matching only.""" + normalized = ensure_input_item_format(item) + if not isinstance(normalized, dict): + return normalized + + clean = cast(dict[str, Any], _sanitize_openai_conversation_item(normalized)) + clean.pop("created_by", None) + if clean.get("status") == "completed": + clean.pop("status", None) + + item_type = clean.get("type") + role = clean.get("role") + if item_type not in (None, "message") or role not in { + "user", + "assistant", + "system", + "developer", + }: + return cast(TResponseInputItem, clean) + + clean.pop("type", None) + if clean.get("phase") is None: + clean.pop("phase", None) + content = clean.get("content") + if not isinstance(content, list) or len(content) != 1 or not isinstance(content[0], dict): + return cast(TResponseInputItem, clean) + + expected_text_type = "output_text" if role == "assistant" else "input_text" + text_part = dict(content[0]) + if text_part.get("type") != expected_text_type: + return cast(TResponseInputItem, clean) + if expected_text_type == "output_text": + if text_part.get("annotations") == []: + text_part.pop("annotations", None) + if text_part.get("logprobs") in (None, []): + text_part.pop("logprobs", None) + elif text_part.get("prompt_cache_breakpoint") is None: + text_part.pop("prompt_cache_breakpoint", None) + if set(text_part) == {"type", "text"} and isinstance(text_part.get("text"), str): + clean["content"] = text_part["text"] + return cast(TResponseInputItem, clean) + + def _openai_conversation_item_requires_id(item: dict[str, Any]) -> bool: """Return whether the Conversations create-item schema requires this item's top-level ID.""" return item.get("type") in _OPENAI_CONVERSATION_ITEM_TYPES_WITH_REQUIRED_ID diff --git a/src/agents/run_state.py b/src/agents/run_state.py index d79e0781a4..d576bfa82f 100644 --- a/src/agents/run_state.py +++ b/src/agents/run_state.py @@ -36,7 +36,7 @@ ProgramOutput, ) from pydantic import BaseModel, StringConstraints, TypeAdapter, ValidationError -from typing_extensions import TypedDict, TypeVar +from typing_extensions import NotRequired, TypedDict, TypeVar from ._run_state_agent_identity import ( _build_agent_identity_keys_by_id, @@ -174,6 +174,7 @@ class _PendingSessionWrite(TypedDict): items: list[TResponseInputItem] before: list[str] | None persisted_count: int + pending_input: NotRequired[list[TResponseInputItem]] def _default_run_state_validation_error( @@ -190,7 +191,7 @@ def _default_run_state_validation_error( # 3. to_json() always emits CURRENT_SCHEMA_VERSION. # 4. Forward compatibility is intentionally fail-fast (older SDKs reject newer or unsupported # versions). -CURRENT_SCHEMA_VERSION = "1.17" +CURRENT_SCHEMA_VERSION = "1.18" _PROGRAMMATIC_TOOL_CALLING_MIN_SCHEMA_VERSION = "1.13" _HOSTED_MCP_APPROVALS_MIN_SCHEMA_VERSION = "1.14" _CURRENT_RESPONSE_OWNERSHIP_MIN_SCHEMA_VERSION = "1.17" @@ -229,6 +230,7 @@ def _default_run_state_validation_error( "Persists Docker container labels and current-response generated-item ownership across " "resume flows, including pending resumed Session writes and terminal-unrecoverable runs." ), + "1.18": "Persists ownership of pending input during unresolved Session writes.", } SUPPORTED_SCHEMA_VERSIONS = frozenset(SCHEMA_VERSION_SUMMARIES) @@ -279,6 +281,16 @@ class _LocalShellCallOutputPayload(TypedDict): {"__dict__", "__pydantic_extra__", "__pydantic_fields_set__"} ) _MISSING_CONTEXT_SENTINEL = object() + + +def _validate_pending_session_write_item(item: Any) -> TResponseInputItem: + """Validate a checkpoint item, including SDK-generated local-shell replay outputs.""" + try: + return _HANDOFF_OUTPUT_ADAPTER.validate_python(item) + except ValidationError: + return cast(TResponseInputItem, _LOCAL_SHELL_OUTPUT_ADAPTER.validate_python(item)) + + _ALLOWED_MISSING_MESSAGE_FIELDS = frozenset({"status"}) @@ -1020,6 +1032,13 @@ def add_input(self, input: str | list[TResponseInputItem]) -> None: def clear_pending_input(self) -> None: """Remove all input staged for the next resumed model call.""" + if ( + self._pending_session_write is not None + and "pending_input" in self._pending_session_write + ): + raise UserError( + "Cannot clear pending input while its Session write is awaiting reconciliation" + ) self._pending_input = [] def get_interruptions(self) -> list[ToolApprovalItem]: @@ -4128,10 +4147,12 @@ async def _build_run_state_from_json( pending_input_raw = state_json.get("pending_input", []) if not isinstance(pending_input_raw, list): raise validation_error_factory("Run state pending_input must be a list", UserError) - state._pending_input = cast( - list[TResponseInputItem], - [dict(item) if isinstance(item, Mapping) else item for item in pending_input_raw], - ) + try: + state._pending_input = [ + _HANDOFF_OUTPUT_ADAPTER.validate_python(item) for item in pending_input_raw + ] + except ValidationError: + raise validation_error_factory("Run state pending_input is invalid", UserError) from None state._model_responses = _deserialize_model_responses(state_json.get("model_responses", [])) serialized_generated_items = state_json.get("generated_items", []) state._generated_items, generated_source_indexes = _deserialize_items_with_source_indexes( @@ -4372,15 +4393,43 @@ async def _build_run_state_from_json( if pending_write is not None: from .run_internal.run_steps import NextStepInterruption, NextStepRunAgain + pending_input_write = ( + pending_write.get("pending_input") if isinstance(pending_write, dict) else None + ) + pending_write_items = ( + pending_write.get("items") if isinstance(pending_write, dict) else None + ) + try: + validated_pending_write_items = ( + [_validate_pending_session_write_item(item) for item in pending_write_items] + if isinstance(pending_write_items, list) + else None + ) + validated_pending_input_write = ( + [_validate_pending_session_write_item(item) for item in pending_input_write] + if isinstance(pending_input_write, list) + else None + ) + except ValidationError: + validated_pending_write_items = None + validated_pending_input_write = None + valid_pending_input = pending_input_write is None or ( + (schema_major, schema_minor) >= (1, 18) and bool(validated_pending_input_write) + ) if ( (schema_major, schema_minor) < (1, 17) or not isinstance(state._current_step, NextStepRunAgain | NextStepInterruption) or not isinstance(pending_write, dict) - or set(pending_write) != {"session_id", "items", "before", "persisted_count"} + or set(pending_write) + not in ( + {"session_id", "items", "before", "persisted_count"}, + {"session_id", "items", "before", "persisted_count", "pending_input"}, + ) or not isinstance(pending_write.get("session_id"), str) - or not isinstance(pending_write.get("items"), list) + or not isinstance(pending_write_items, list) or not pending_write["items"] or not all(isinstance(item, dict) for item in pending_write["items"]) + or not validated_pending_write_items or ( pending_write.get("before") is not None and ( @@ -4390,9 +4439,13 @@ async def _build_run_state_from_json( ) or type(pending_write.get("persisted_count")) is not int or pending_write["persisted_count"] < 0 + or not valid_pending_input ): raise validation_error_factory("Run state pending Session write is invalid", UserError) state._pending_session_write = copy.deepcopy(cast(_PendingSessionWrite, pending_write)) + state._pending_session_write["items"] = validated_pending_write_items + if validated_pending_input_write is not None: + state._pending_session_write["pending_input"] = validated_pending_input_write terminal_unrecoverable = state_json.get("terminal_unrecoverable") if terminal_unrecoverable is not None: # An older label never wrote this marker, so honoring one would let a snapshot claim a diff --git a/tests/sandbox/test_docker.py b/tests/sandbox/test_docker.py index e4c7cc812f..d32a670659 100644 --- a/tests/sandbox/test_docker.py +++ b/tests/sandbox/test_docker.py @@ -1926,7 +1926,7 @@ async def test_docker_labels_roundtrip_through_run_state() -> None: serialized = run_state.to_json() restored = await RunState.from_json(agent, serialized) - assert serialized["$schemaVersion"] == CURRENT_SCHEMA_VERSION == "1.17" + assert serialized["$schemaVersion"] == CURRENT_SCHEMA_VERSION assert restored._sandbox is not None restored_session_state = restored._sandbox["session_state"] assert isinstance(restored_session_state, dict) diff --git a/tests/test_run_impl_resume_paths.py b/tests/test_run_impl_resume_paths.py index 482a13edd8..fe97eb02c5 100644 --- a/tests/test_run_impl_resume_paths.py +++ b/tests/test_run_impl_resume_paths.py @@ -514,6 +514,32 @@ async def test_pending_session_write_rejects_invalid_serialized_checkpoint(inval await RunState.from_json(agent, payload) +@pytest.mark.asyncio +async def test_pending_session_write_accepts_local_shell_replay_output() -> None: + agent, _, session, state, _ = await _approved_session_state(False) + session.failure = "before" + with pytest.raises(RuntimeError): + await _run_session_resume(agent, state, session, False) + + payload = state.to_json() + payload["pending_session_write"]["items"] = [ + { + "type": "local_shell_call_output", + "call_id": "shell-1", + "output": "replayed", + } + ] + + restored = await RunState.from_json(agent, payload) + assert restored.to_json()["pending_session_write"]["items"] == [ + { + "type": "local_shell_call_output", + "call_id": "shell-1", + "output": "replayed", + } + ] + + @pytest.mark.asyncio async def test_resumed_session_append_partial_commit_fails_closed() -> None: agent, model, session, state, effects = await _approved_session_state(False) diff --git a/tests/test_run_state.py b/tests/test_run_state.py index cd2daa51b6..4505ac8939 100644 --- a/tests/test_run_state.py +++ b/tests/test_run_state.py @@ -8611,6 +8611,7 @@ def test_supported_schema_versions_match_released_boundary(self): "1.14", "1.15", "1.16", + "1.17", CURRENT_SCHEMA_VERSION, } ) diff --git a/tests/test_run_state_compatibility_corpus.py b/tests/test_run_state_compatibility_corpus.py index 52996188cc..37314afca7 100644 --- a/tests/test_run_state_compatibility_corpus.py +++ b/tests/test_run_state_compatibility_corpus.py @@ -19,7 +19,7 @@ from agents import Agent, RunState, UserError from agents.run_context import RunContextWrapper -from agents.run_state import SUPPORTED_SCHEMA_VERSIONS +from agents.run_state import CURRENT_SCHEMA_VERSION, SUPPORTED_SCHEMA_VERSIONS from agents.sandbox.entries.mounts.patterns import FuseMountConfig from integration_tests._contract_state import ( _deserialize_common_sandbox_session_state, @@ -164,10 +164,12 @@ def test_historical_state_comparison_preserves_json_scalar_types() -> None: def test_historical_fixture_corpus_matches_supported_schema_versions() -> None: assert SOURCES["baseline"] == "v0.19.4" - assert frozenset(SOURCES["versions"]) == SUPPORTED_SCHEMA_VERSIONS + assert frozenset(SOURCES["versions"]) | {CURRENT_SCHEMA_VERSION} == SUPPORTED_SCHEMA_VERSIONS assert all(entry["commit"] for entry in SOURCES["versions"].values()) assert {entry["version"] for entry in SOURCES["features"]} == { - version for version in SUPPORTED_SCHEMA_VERSIONS if version not in {"1.0", "1.1"} + version + for version in SUPPORTED_SCHEMA_VERSIONS + if version not in {"1.0", "1.1", CURRENT_SCHEMA_VERSION} } assert {entry["provenance"] for entry in SOURCES["features"]} == { "historical_writer", diff --git a/tests/test_run_state_pending_input.py b/tests/test_run_state_pending_input.py index fdf1d2cd71..f447d33b22 100644 --- a/tests/test_run_state_pending_input.py +++ b/tests/test_run_state_pending_input.py @@ -1,7 +1,11 @@ from __future__ import annotations +import asyncio +import copy import json -from typing import Any, cast +from types import SimpleNamespace +from typing import Any, Literal, cast +from unittest.mock import AsyncMock import pytest from openai.types.responses.response_computer_tool_call import ( @@ -14,6 +18,7 @@ from agents.guardrail import GuardrailFunctionOutput, InputGuardrail from agents.items import ModelResponse, TResponseInputItem from agents.lifecycle import AgentHooks, RunHooks +from agents.memory import OpenAIConversationsSession, Session from agents.run import CallModelData, ModelInputData from agents.run_context import RunContextWrapper from agents.run_internal.oai_conversation import OpenAIServerConversationTracker @@ -29,6 +34,70 @@ from .utils.simple_session import SimpleListSession +class _PendingInputWriteFailureSession(SimpleListSession): + def __init__(self) -> None: + super().__init__() + self.failure: Literal["before", "after"] | None = None + self.error = RuntimeError("pending input Session append failed") + self.block_next_add = False + self.add_started = asyncio.Event() + self.release_add = asyncio.Event() + + async def add_items(self, items: list[TResponseInputItem]) -> None: + failure, self.failure = self.failure, None + if failure == "before": + raise self.error + if self.block_next_add: + self.block_next_add = False + self.add_started.set() + await self.release_add.wait() + await super().add_items(items) + if failure == "after": + raise self.error + + +class _RecordingConversationsSession(OpenAIConversationsSession): + def __init__(self) -> None: + self.create_conversation = AsyncMock(return_value=SimpleNamespace(id="test")) + client = SimpleNamespace( + conversations=SimpleNamespace(create=self.create_conversation), + ) + super().__init__(openai_client=cast(Any, client)) + self.items: list[TResponseInputItem] = [] + self.failure: Literal["before", "after"] | None = None + self.error = RuntimeError("conversation append failed") + + async def get_items(self, limit: int | None = None) -> list[TResponseInputItem]: + await self._get_session_id() + if limit == 0: + return [] + return list(self.items if limit is None else self.items[-limit:]) + + async def add_items(self, items: list[TResponseInputItem]) -> None: + failure, self.failure = self.failure, None + if failure == "before": + raise self.error + for item in items: + normalized = copy.deepcopy(item) + if ( + isinstance(normalized, dict) + and normalized.get("type") == "message" + and isinstance(normalized.get("content"), str) + ): + normalized["content"] = [ + { + "type": "input_text" + if normalized.get("role") != "assistant" + else "output_text", + "text": normalized["content"], + } + ] + normalized["id"] = "msg_server_assigned" + self.items.append(normalized) + if failure == "after": + raise self.error + + def _item_type(item: TResponseInputItem) -> str | None: if not isinstance(item, dict): return getattr(item, "type", None) @@ -52,7 +121,7 @@ def _message_text(item: TResponseInputItem) -> str | None: async def _make_after_turn_state( *, - session: SimpleListSession | None = None, + session: Session | None = None, auto_previous_response_id: bool = False, ) -> tuple[ScriptedModel, Agent[Any], RunState[Any], list[str]]: calls: list[str] = [] @@ -89,6 +158,22 @@ def record_destination(destination: str) -> str: return model, agent, state, calls +async def _resume_pending_input_state( + agent: Agent[Any], + state: RunState[Any], + session: Session, + *, + streamed: bool, +) -> Any: + run_config = RunConfig(tracing_disabled=True) + if not streamed: + return await Runner.run(agent, state, session=session, run_config=run_config) + result = Runner.run_streamed(agent, state, session=session, run_config=run_config) + async for _event in result.stream_events(): + pass + return result + + @pytest.mark.asyncio async def test_pending_input_preserves_order_and_serialization_round_trips() -> None: agent = Agent(name="assistant") @@ -168,6 +253,154 @@ async def test_after_turn_resume_admits_input_after_tool_output_exactly_once() - terminal_state.add_input("Too late") +@pytest.mark.asyncio +@pytest.mark.parametrize("failing_streamed", [False, True]) +@pytest.mark.parametrize("recovery_streamed", [False, True]) +@pytest.mark.parametrize("round_trip", [False, True]) +@pytest.mark.parametrize("failure", ["before", "after"]) +async def test_pending_input_session_append_reconciles_once_without_rerunning_guardrails( + failing_streamed: bool, + recovery_streamed: bool, + round_trip: bool, + failure: Literal["before", "after"], +) -> None: + session = _PendingInputWriteFailureSession() + model, agent, state, _calls = await _make_after_turn_state(session=session) + guardrail_calls = 0 + + def inspect_pending_input( + _context: RunContextWrapper[Any], + _agent: Agent[Any], + _input: str | list[TResponseInputItem], + ) -> GuardrailFunctionOutput: + nonlocal guardrail_calls + guardrail_calls += 1 + return GuardrailFunctionOutput(output_info="accepted", tripwire_triggered=False) + + agent.input_guardrails = [InputGuardrail(guardrail_function=inspect_pending_input)] + state.add_input("Late input") + model.enqueue([get_text_message("Recovered")]) + model_calls_before = len(model.calls) + session.failure = failure + + with pytest.raises(RuntimeError, match="pending input Session append failed"): + if failing_streamed: + failed_result = Runner.run_streamed( + agent, state, session=session, run_config=RunConfig(tracing_disabled=True) + ) + async for _event in failed_result.stream_events(): + pass + state = failed_result.to_state() + else: + await Runner.run( + agent, state, session=session, run_config=RunConfig(tracing_disabled=True) + ) + + assert guardrail_calls == 1 + assert len(model.calls) == model_calls_before + assert state._pending_session_write is not None + assert [_message_text(item) for item in state.pending_input] == ["Late input"] + if round_trip: + state = await RunState.from_json(agent, state.to_json()) + + result = await _resume_pending_input_state(agent, state, session, streamed=recovery_streamed) + + assert result.final_output == "Recovered" + assert len(model.calls) == model_calls_before + 1 + assert guardrail_calls == 1 + assert state.pending_input == [] + assert state._pending_session_write is None + assert [_message_text(item) for item in await session.get_items()].count("Late input") == 1 + + +@pytest.mark.asyncio +async def test_pending_input_session_checkpoint_rejects_malformed_input_before_resume() -> None: + session = _PendingInputWriteFailureSession() + model, agent, state, _calls = await _make_after_turn_state(session=session) + state.add_input("Late input") + model.enqueue([get_text_message("Recovered")]) + session.failure = "before" + + with pytest.raises(RuntimeError, match="pending input Session append failed"): + await Runner.run(agent, state, session=session, run_config=RunConfig(tracing_disabled=True)) + + payload = state.to_json() + malformed = {"role": ["user"]} + payload["pending_input"] = [malformed] + pending_write = cast(dict[str, Any], payload["pending_session_write"]) + pending_write["pending_input"] = [malformed] + pending_write["items"] = [malformed] + + with pytest.raises(RuntimeError, match="Error details are redacted"): + await RunState.from_json(agent, payload) + + +@pytest.mark.asyncio +async def test_pending_input_conversations_session_reconciles_sanitized_message_id() -> None: + session = _RecordingConversationsSession() + model, agent, state, _calls = await _make_after_turn_state(session=session) + state.add_input( + [ + { + "id": "user-id", + "type": "message", + "role": "user", + "content": "Late input", + } + ] + ) + model.enqueue([get_text_message("Recovered")]) + session.failure = "after" + + with pytest.raises(RuntimeError, match="conversation append failed"): + await Runner.run(agent, state, session=session, run_config=RunConfig(tracing_disabled=True)) + + state = await RunState.from_json(agent, state.to_json()) + result = await Runner.run( + agent, state, session=session, run_config=RunConfig(tracing_disabled=True) + ) + + assert result.final_output == "Recovered" + assert [_message_text(item) for item in session.items].count("Late input") == 1 + session.create_conversation.assert_awaited_once_with(items=[]) + + +@pytest.mark.asyncio +async def test_pending_input_added_during_session_write_survives_stream_checkpoint() -> None: + session = _PendingInputWriteFailureSession() + model, agent, state, _calls = await _make_after_turn_state(session=session) + state.add_input("Before write") + model.enqueue([get_text_message("Recovered")]) + session.failure = "after" + session.block_next_add = True + + failed_result = Runner.run_streamed( + agent, state, session=session, run_config=RunConfig(tracing_disabled=True) + ) + + async def consume_stream() -> None: + async for _event in failed_result.stream_events(): + pass + + consume_task = asyncio.create_task(consume_stream()) + await session.add_started.wait() + state.add_input("During write") + session.release_add.set() + with pytest.raises(RuntimeError, match="pending input Session append failed"): + await consume_task + checkpoint = failed_result.to_state() + assert [_message_text(item) for item in checkpoint.pending_input] == [ + "Before write", + "During write", + ] + + result = await _resume_pending_input_state(agent, checkpoint, session, streamed=False) + + assert result.final_output == "Recovered" + assert [_message_text(item) for item in await session.get_items()].count("Before write") == 1 + assert [_message_text(item) for item in await session.get_items()].count("During write") == 1 + + @pytest.mark.asyncio async def test_streamed_resume_matches_pending_input_ordering() -> None: model, agent, state, calls = await _make_after_turn_state() From 78784a32c68dbb3ef6226130f3526c82ad014e69 Mon Sep 17 00:00:00 2001 From: Sean Date: Tue, 8 Sep 2026 11:52:35 +0800 Subject: [PATCH 3/8] fix(sessions): preserve pending write ownership metadata --- .../run_internal/session_persistence.py | 8 ++-- src/agents/run_state.py | 9 +++- tests/test_run_impl_resume_paths.py | 46 ++++++++++++++++++- 3 files changed, 56 insertions(+), 7 deletions(-) diff --git a/src/agents/run_internal/session_persistence.py b/src/agents/run_internal/session_persistence.py index 7a378a1c6a..360bc1a96c 100644 --- a/src/agents/run_internal/session_persistence.py +++ b/src/agents/run_internal/session_persistence.py @@ -889,7 +889,10 @@ def appended_digests(items: Sequence[TResponseInputItem]) -> list[str]: if before is None: # No append has started. Retain the batch even if this first read fails. tail = await _session_get_items( - session, limit=len(pending["items"]) + 1, wrapper=wrapper + session, + limit=len(pending["items"]) + 1, + wrapper=wrapper, + capture_compaction_generation=True, ) pending["before"] = legacy_digests(tail) append = True @@ -902,6 +905,7 @@ def appended_digests(items: Sequence[TResponseInputItem]) -> list[str]: get_with_generation, limit=expected_length, ) + wrapper._session_compaction_generation = committed_generation # type: ignore[attr-defined] else: tail = await _session_get_items(session, limit=expected_length, wrapper=wrapper) observed = legacy_digests(tail) @@ -927,8 +931,6 @@ def appended_digests(items: Sequence[TResponseInputItem]) -> list[str]: "Cannot reconcile the pending Session write: history changed or is ambiguous. " "Repair the original Session before resuming; do not rerun the completed tool." ) - if committed and committed_generation is not None and wrapper is not None: - wrapper._session_compaction_generation = committed_generation # type: ignore[attr-defined] if append: # Backends may retain or transform their input; the durable checkpoint stays detached. await _session_add_items(session, copy.deepcopy(pending["items"]), wrapper=wrapper) diff --git a/src/agents/run_state.py b/src/agents/run_state.py index d576bfa82f..c4d8be878e 100644 --- a/src/agents/run_state.py +++ b/src/agents/run_state.py @@ -4396,6 +4396,9 @@ async def _build_run_state_from_json( pending_input_write = ( pending_write.get("pending_input") if isinstance(pending_write, dict) else None ) + pending_input_field_present = ( + isinstance(pending_write, dict) and "pending_input" in pending_write + ) pending_write_items = ( pending_write.get("items") if isinstance(pending_write, dict) else None ) @@ -4413,8 +4416,10 @@ async def _build_run_state_from_json( except ValidationError: validated_pending_write_items = None validated_pending_input_write = None - valid_pending_input = pending_input_write is None or ( - (schema_major, schema_minor) >= (1, 18) and bool(validated_pending_input_write) + valid_pending_input = not pending_input_field_present or ( + (schema_major, schema_minor) >= (1, 18) + and isinstance(pending_input_write, list) + and bool(validated_pending_input_write) ) if ( (schema_major, schema_minor) < (1, 17) diff --git a/tests/test_run_impl_resume_paths.py b/tests/test_run_impl_resume_paths.py index fe97eb02c5..27c2f52ae9 100644 --- a/tests/test_run_impl_resume_paths.py +++ b/tests/test_run_impl_resume_paths.py @@ -342,6 +342,46 @@ async def compact(**kwargs: Any) -> SimpleNamespace: backend.close() +@pytest.mark.asyncio +@pytest.mark.parametrize("round_trip", [False, True], ids=["live", "json"]) +async def test_pending_input_recovery_refreshes_compaction_generation(round_trip: bool) -> None: + backend = _FailingResumeSession() + compaction_inputs: list[list[TResponseInputItem]] = [] + compact_enabled = False + + async def compact(**kwargs: Any) -> SimpleNamespace: + items = copy.deepcopy(kwargs["input"]) + compaction_inputs.append(items) + return SimpleNamespace(output=items, usage=None) + + session = OpenAIResponsesCompactionSession( + backend.session_id, + underlying_session=backend, + client=cast(Any, SimpleNamespace(responses=SimpleNamespace(compact=compact))), + compaction_mode="input", + should_trigger_compaction=lambda _: compact_enabled, + ) + agent, model, _, state, effects = await _approved_session_state(False, session) + await session.run_compaction() + state.add_input("Late input") + backend.failure = "before" + + try: + with pytest.raises(RuntimeError) as error: + await _run_session_resume(agent, state, session, False) + assert error.value is backend.error + if round_trip: + state = await RunState.from_json(agent, state.to_json()) + + compact_enabled = True + result = await _run_session_resume(agent, state, session, False) + assert result.final_output == "done" + assert effects == [7] + assert len(compaction_inputs) == 1 + finally: + await backend.clear_session() + + @pytest.mark.asyncio @pytest.mark.parametrize("mode", ["input", "auto"]) async def test_compaction_reload_preserves_session_retrieval_window( @@ -499,7 +539,7 @@ async def test_failed_streamed_result_checkpoint_retains_detached_pending_write( @pytest.mark.asyncio -@pytest.mark.parametrize("invalid", ["old-schema", "batch-shape"]) +@pytest.mark.parametrize("invalid", ["old-schema", "batch-shape", "pending-input-null"]) async def test_pending_session_write_rejects_invalid_serialized_checkpoint(invalid: str) -> None: agent, _, session, state, _ = await _approved_session_state(False) session.failure = "before" @@ -508,8 +548,10 @@ async def test_pending_session_write_rejects_invalid_serialized_checkpoint(inval payload = state.to_json() if invalid == "old-schema": payload["$schemaVersion"] = "1.16" - else: + elif invalid == "batch-shape": payload["pending_session_write"]["items"] = "not an item batch" + else: + payload["pending_session_write"]["pending_input"] = None with pytest.raises(UserError, match="pending Session write is invalid"): await RunState.from_json(agent, payload) From 4a9e9f5f7b2271f11ddcf5f9964c303b2168aac5 Mon Sep 17 00:00:00 2001 From: Sean Date: Tue, 8 Sep 2026 11:57:45 +0800 Subject: [PATCH 4/8] fix(run-state): materialize computer output iterables --- src/agents/run_state.py | 6 ++++-- tests/test_run_impl_resume_paths.py | 26 ++++++++++++++++++++++++++ 2 files changed, 30 insertions(+), 2 deletions(-) diff --git a/src/agents/run_state.py b/src/agents/run_state.py index c4d8be878e..4d65df75d4 100644 --- a/src/agents/run_state.py +++ b/src/agents/run_state.py @@ -285,10 +285,12 @@ class _LocalShellCallOutputPayload(TypedDict): def _validate_pending_session_write_item(item: Any) -> TResponseInputItem: """Validate a checkpoint item, including SDK-generated local-shell replay outputs.""" + validated: Any try: - return _HANDOFF_OUTPUT_ADAPTER.validate_python(item) + validated = _HANDOFF_OUTPUT_ADAPTER.validate_python(item) except ValidationError: - return cast(TResponseInputItem, _LOCAL_SHELL_OUTPUT_ADAPTER.validate_python(item)) + validated = _LOCAL_SHELL_OUTPUT_ADAPTER.validate_python(item) + return cast(TResponseInputItem, _to_dump_compatible(validated)) _ALLOWED_MISSING_MESSAGE_FIELDS = frozenset({"status"}) diff --git a/tests/test_run_impl_resume_paths.py b/tests/test_run_impl_resume_paths.py index 27c2f52ae9..299548a462 100644 --- a/tests/test_run_impl_resume_paths.py +++ b/tests/test_run_impl_resume_paths.py @@ -582,6 +582,32 @@ async def test_pending_session_write_accepts_local_shell_replay_output() -> None ] +@pytest.mark.asyncio +async def test_pending_session_write_materializes_computer_safety_checks() -> None: + agent, _, session, state, _ = await _approved_session_state(False) + session.failure = "before" + with pytest.raises(RuntimeError): + await _run_session_resume(agent, state, session, False) + + payload = state.to_json() + payload["pending_session_write"]["items"] = [ + { + "type": "computer_call_output", + "call_id": "computer-1", + "output": {"type": "computer_screenshot", "image_url": "img"}, + "acknowledged_safety_checks": [ + {"id": "check-1", "code": "confirm", "message": "approved"} + ], + } + ] + + restored = await RunState.from_json(agent, payload) + expected = payload["pending_session_write"]["items"] + assert restored.to_json()["pending_session_write"]["items"] == expected + roundtripped = await RunState.from_string(agent, restored.to_string()) + assert roundtripped.to_json()["pending_session_write"]["items"] == expected + + @pytest.mark.asyncio async def test_resumed_session_append_partial_commit_fails_closed() -> None: agent, model, session, state, effects = await _approved_session_state(False) From 315b4bb3c1f92654e3da2e4b1cbb02d4af1e7319 Mon Sep 17 00:00:00 2001 From: Sean Date: Tue, 8 Sep 2026 12:03:51 +0800 Subject: [PATCH 5/8] fix(run): admit input added during session writes --- src/agents/run.py | 74 ++++++++++-------- src/agents/run_internal/run_loop.py | 103 ++++++++++++++------------ tests/test_run_state_pending_input.py | 26 +++++++ 3 files changed, 122 insertions(+), 81 deletions(-) diff --git a/src/agents/run.py b/src/agents/run.py index 297895a347..ebe97134da 100644 --- a/src/agents/run.py +++ b/src/agents/run.py @@ -1447,40 +1447,50 @@ def _mark_response_hooks_started() -> None: if run_state._current_step is None: run_state._current_step = NextStepRunAgain() - pending_input = run_state.pending_input - if pending_input: - pending_guardrails = current_agent.input_guardrails + ( - run_config.input_guardrails or [] - ) - try: - await run_input_guardrails( - current_agent, - pending_guardrails, - pending_input, - context_wrapper, - input_guardrail_results, + if run_state._pending_input: + while True: + pending_input = run_state.pending_input + if not pending_input: + break + pending_guardrails = current_agent.input_guardrails + ( + run_config.input_guardrails or [] ) - finally: - run_state._input_guardrail_results = list(input_guardrail_results) + try: + await run_input_guardrails( + current_agent, + pending_guardrails, + pending_input, + context_wrapper, + input_guardrail_results, + ) + finally: + run_state._input_guardrail_results = list( + input_guardrail_results + ) - admission_items = await admit_pending_input( - run_state=run_state, - agent=current_agent, - session=session, - server_conversation_tracker=server_conversation_tracker, - store=store_setting, - wrapper=context_wrapper, - ) - generated_items.extend(admission_items) - session_items.extend(admission_items) - if pending_server_items is not None: - pending_server_items.extend(admission_items) - pending_input_admission_items = [ - item for item in admission_items if isinstance(item, InputItem) - ] - if not run_state._pending_input: - run_state._generated_items = list(generated_items) - run_state._session_items = list(session_items) + admission_items = await admit_pending_input( + run_state=run_state, + agent=current_agent, + session=session, + server_conversation_tracker=server_conversation_tracker, + store=store_setting, + wrapper=context_wrapper, + ) + generated_items.extend(admission_items) + session_items.extend(admission_items) + if pending_server_items is not None: + pending_server_items.extend(admission_items) + pending_input_admission_items.extend( + item for item in admission_items if isinstance(item, InputItem) + ) + if not run_state._pending_input: + run_state._generated_items = list(generated_items) + run_state._session_items = list(session_items) + if ( + server_conversation_tracker is not None + or not run_state._pending_input + ): + break all_tools = await get_all_tools(execution_agent, context_wrapper) all_tools = await initialize_computer_tools( tools=all_tools, context_wrapper=context_wrapper diff --git a/src/agents/run_internal/run_loop.py b/src/agents/run_internal/run_loop.py index 9871a54041..2c7ed5484a 100644 --- a/src/agents/run_internal/run_loop.py +++ b/src/agents/run_internal/run_loop.py @@ -1496,58 +1496,63 @@ async def _save_max_turns_items( if run_state is not None and run_state._pending_input: if run_state._current_step is None: run_state._current_step = NextStepRunAgain() - pending_input = run_state.pending_input - pending_guardrails = current_agent.input_guardrails + ( - run_config.input_guardrails or [] - ) - previous_result_count = len(streamed_result.input_guardrail_results) - try: - await run_input_guardrails_with_queue( - current_agent, - pending_guardrails, - pending_input, - context_wrapper, - streamed_result, - current_span, + while True: + pending_input = run_state.pending_input + if not pending_input: + break + pending_guardrails = current_agent.input_guardrails + ( + run_config.input_guardrails or [] ) - finally: - run_state._input_guardrail_results = list( - streamed_result.input_guardrail_results + previous_result_count = len(streamed_result.input_guardrail_results) + try: + await run_input_guardrails_with_queue( + current_agent, + pending_guardrails, + pending_input, + context_wrapper, + streamed_result, + current_span, + ) + finally: + run_state._input_guardrail_results = list( + streamed_result.input_guardrail_results + ) + tripping_result = next( + ( + result + for result in streamed_result.input_guardrail_results[ + previous_result_count: + ] + if result.output.tripwire_triggered + ), + None, ) - tripping_result = next( - ( - result - for result in streamed_result.input_guardrail_results[ - previous_result_count: - ] - if result.output.tripwire_triggered - ), - None, - ) - if tripping_result is not None: - raise InputGuardrailTripwireTriggered(tripping_result) + if tripping_result is not None: + raise InputGuardrailTripwireTriggered(tripping_result) - store_setting = current_agent.model_settings.resolve( - run_config.model_settings - ).store - admission_items = await admit_pending_input( - run_state=run_state, - agent=current_agent, - session=session, - server_conversation_tracker=server_conversation_tracker, - store=store_setting, - wrapper=context_wrapper, - ) - streamed_result._model_input_items.extend(admission_items) - streamed_result.new_items.extend(admission_items) - if pending_server_items is not None: - pending_server_items.extend(admission_items) - pending_input_admission_items = [ - item for item in admission_items if isinstance(item, InputItem) - ] - if not run_state._pending_input: - run_state._generated_items = list(streamed_result._model_input_items) - run_state._session_items = list(streamed_result.new_items) + store_setting = current_agent.model_settings.resolve( + run_config.model_settings + ).store + admission_items = await admit_pending_input( + run_state=run_state, + agent=current_agent, + session=session, + server_conversation_tracker=server_conversation_tracker, + store=store_setting, + wrapper=context_wrapper, + ) + streamed_result._model_input_items.extend(admission_items) + streamed_result.new_items.extend(admission_items) + if pending_server_items is not None: + pending_server_items.extend(admission_items) + pending_input_admission_items.extend( + item for item in admission_items if isinstance(item, InputItem) + ) + if not run_state._pending_input: + run_state._generated_items = list(streamed_result._model_input_items) + run_state._session_items = list(streamed_result.new_items) + if server_conversation_tracker is not None or not run_state._pending_input: + break all_tools = await get_all_tools(execution_agent, context_wrapper) all_tools = await initialize_computer_tools( diff --git a/tests/test_run_state_pending_input.py b/tests/test_run_state_pending_input.py index f447d33b22..f32485e1e0 100644 --- a/tests/test_run_state_pending_input.py +++ b/tests/test_run_state_pending_input.py @@ -401,6 +401,32 @@ async def consume_stream() -> None: assert [_message_text(item) for item in await session.get_items()].count("During write") == 1 +@pytest.mark.asyncio +@pytest.mark.parametrize("streamed", [False, True]) +async def test_pending_input_added_during_successful_session_write_is_admitted( + streamed: bool, +) -> None: + session = _PendingInputWriteFailureSession() + model, agent, state, _calls = await _make_after_turn_state(session=session) + state.add_input("Before write") + model.enqueue([get_text_message("Recovered")]) + session.block_next_add = True + + result_task = asyncio.create_task( + _resume_pending_input_state(agent, state, session, streamed=streamed) + ) + await session.add_started.wait() + state.add_input("During write") + session.release_add.set() + result = await result_task + + assert result.final_output == "Recovered" + assert state.pending_input == [] + session_items = await session.get_items() + assert [_message_text(item) for item in session_items].count("Before write") == 1 + assert [_message_text(item) for item in session_items].count("During write") == 1 + + @pytest.mark.asyncio async def test_streamed_resume_matches_pending_input_ordering() -> None: model, agent, state, calls = await _make_after_turn_state() From 5c7106f2314ef2e76fcd747e3b01400aee698548 Mon Sep 17 00:00:00 2001 From: Sean Date: Tue, 8 Sep 2026 12:14:16 +0800 Subject: [PATCH 6/8] fix(run-state): restore staged local-shell input --- src/agents/run_state.py | 2 +- tests/test_run_impl_resume_paths.py | 22 ++++++++++++++++++++++ 2 files changed, 23 insertions(+), 1 deletion(-) diff --git a/src/agents/run_state.py b/src/agents/run_state.py index 4d65df75d4..88d19b4ad2 100644 --- a/src/agents/run_state.py +++ b/src/agents/run_state.py @@ -4151,7 +4151,7 @@ async def _build_run_state_from_json( raise validation_error_factory("Run state pending_input must be a list", UserError) try: state._pending_input = [ - _HANDOFF_OUTPUT_ADAPTER.validate_python(item) for item in pending_input_raw + _validate_pending_session_write_item(item) for item in pending_input_raw ] except ValidationError: raise validation_error_factory("Run state pending_input is invalid", UserError) from None diff --git a/tests/test_run_impl_resume_paths.py b/tests/test_run_impl_resume_paths.py index 299548a462..694d07687b 100644 --- a/tests/test_run_impl_resume_paths.py +++ b/tests/test_run_impl_resume_paths.py @@ -582,6 +582,28 @@ async def test_pending_session_write_accepts_local_shell_replay_output() -> None ] +@pytest.mark.asyncio +async def test_top_level_pending_input_accepts_local_shell_replay_output() -> None: + agent, _, session, state, _ = await _approved_session_state(False) + session.failure = "before" + with pytest.raises(RuntimeError): + await _run_session_resume(agent, state, session, False) + + payload = state.to_json() + payload["pending_input"] = [ + { + "type": "local_shell_call_output", + "call_id": "shell-1", + "output": "replayed", + } + ] + + restored = await RunState.from_json(agent, payload) + assert restored.to_json()["pending_input"] == payload["pending_input"] + roundtripped = await RunState.from_string(agent, restored.to_string()) + assert roundtripped.to_json()["pending_input"] == payload["pending_input"] + + @pytest.mark.asyncio async def test_pending_session_write_materializes_computer_safety_checks() -> None: agent, _, session, state, _ = await _approved_session_state(False) From c289a5db6b804fc6ed4e22841dfb521f282e198b Mon Sep 17 00:00:00 2001 From: Sean Date: Tue, 8 Sep 2026 12:25:51 +0800 Subject: [PATCH 7/8] fix(run-state): preserve replay provider metadata --- src/agents/run_state.py | 7 ++++++- tests/test_run_impl_resume_paths.py | 29 +++++++++++++++++++++++++++++ 2 files changed, 35 insertions(+), 1 deletion(-) diff --git a/src/agents/run_state.py b/src/agents/run_state.py index 88d19b4ad2..05dd787079 100644 --- a/src/agents/run_state.py +++ b/src/agents/run_state.py @@ -285,12 +285,17 @@ class _LocalShellCallOutputPayload(TypedDict): def _validate_pending_session_write_item(item: Any) -> TResponseInputItem: """Validate a checkpoint item, including SDK-generated local-shell replay outputs.""" + has_provider_data = isinstance(item, Mapping) and "provider_data" in item + provider_data = item.get("provider_data") if has_provider_data else None validated: Any try: validated = _HANDOFF_OUTPUT_ADAPTER.validate_python(item) except ValidationError: validated = _LOCAL_SHELL_OUTPUT_ADAPTER.validate_python(item) - return cast(TResponseInputItem, _to_dump_compatible(validated)) + dumped = _to_dump_compatible(validated) + if has_provider_data and isinstance(dumped, dict): + dumped["provider_data"] = _to_dump_compatible(provider_data) + return cast(TResponseInputItem, dumped) _ALLOWED_MISSING_MESSAGE_FIELDS = frozenset({"status"}) diff --git a/tests/test_run_impl_resume_paths.py b/tests/test_run_impl_resume_paths.py index 694d07687b..bba1c62d60 100644 --- a/tests/test_run_impl_resume_paths.py +++ b/tests/test_run_impl_resume_paths.py @@ -604,6 +604,35 @@ async def test_top_level_pending_input_accepts_local_shell_replay_output() -> No assert roundtripped.to_json()["pending_input"] == payload["pending_input"] +@pytest.mark.asyncio +async def test_pending_input_preserves_provider_data_during_restore() -> None: + agent, _, session, state, _ = await _approved_session_state(False) + session.failure = "before" + with pytest.raises(RuntimeError): + await _run_session_resume(agent, state, session, False) + + provider_item = { + "type": "function_call_output", + "call_id": "call-provider-data", + "output": "answer", + "provider_data": { + "thinking_blocks": [{"type": "thinking", "thinking": "hidden", "signature": "sig-1"}] + }, + } + payload = state.to_json() + payload["pending_input"] = [provider_item] + pending_write = cast(dict[str, Any], payload["pending_session_write"]) + pending_write["pending_input"] = [provider_item] + pending_write["items"] = [provider_item] + + restored = await RunState.from_json(agent, payload) + assert restored.to_json()["pending_input"] == [provider_item] + assert restored.to_json()["pending_session_write"]["pending_input"] == [provider_item] + roundtripped = await RunState.from_string(agent, restored.to_string()) + assert roundtripped.to_json()["pending_input"] == [provider_item] + assert roundtripped.to_json()["pending_session_write"]["items"] == [provider_item] + + @pytest.mark.asyncio async def test_pending_session_write_materializes_computer_safety_checks() -> None: agent, _, session, state, _ = await _approved_session_state(False) From b377972dfc7d3154e6a11596b1bb0397decfc1ab Mon Sep 17 00:00:00 2001 From: Sean Date: Tue, 8 Sep 2026 12:34:09 +0800 Subject: [PATCH 8/8] test(run-state): cover provider metadata lost ack --- src/agents/run_state.py | 13 ++++---- tests/test_run_state_pending_input.py | 47 +++++++++++++++++++++++++++ 2 files changed, 54 insertions(+), 6 deletions(-) diff --git a/src/agents/run_state.py b/src/agents/run_state.py index 05dd787079..872a4dc520 100644 --- a/src/agents/run_state.py +++ b/src/agents/run_state.py @@ -285,17 +285,18 @@ class _LocalShellCallOutputPayload(TypedDict): def _validate_pending_session_write_item(item: Any) -> TResponseInputItem: """Validate a checkpoint item, including SDK-generated local-shell replay outputs.""" - has_provider_data = isinstance(item, Mapping) and "provider_data" in item - provider_data = item.get("provider_data") if has_provider_data else None + provider_data: Any = None + if isinstance(item, Mapping) and isinstance(item.get("provider_data"), Mapping): + provider_data = _copy_json_compatible_value(item["provider_data"], set()) validated: Any try: validated = _HANDOFF_OUTPUT_ADAPTER.validate_python(item) except ValidationError: validated = _LOCAL_SHELL_OUTPUT_ADAPTER.validate_python(item) - dumped = _to_dump_compatible(validated) - if has_provider_data and isinstance(dumped, dict): - dumped["provider_data"] = _to_dump_compatible(provider_data) - return cast(TResponseInputItem, dumped) + materialized = _to_dump_compatible(validated) + if provider_data is not None and isinstance(materialized, dict): + materialized["provider_data"] = provider_data + return cast(TResponseInputItem, materialized) _ALLOWED_MISSING_MESSAGE_FIELDS = frozenset({"status"}) diff --git a/tests/test_run_state_pending_input.py b/tests/test_run_state_pending_input.py index f32485e1e0..afbbeb8c81 100644 --- a/tests/test_run_state_pending_input.py +++ b/tests/test_run_state_pending_input.py @@ -365,6 +365,53 @@ async def test_pending_input_conversations_session_reconciles_sanitized_message_ session.create_conversation.assert_awaited_once_with(items=[]) +@pytest.mark.asyncio +async def test_pending_input_provider_data_survives_lost_ack_reconciliation() -> None: + session = _PendingInputWriteFailureSession() + model, agent, state, _calls = await _make_after_turn_state(session=session) + provider_data = { + "model": "litellm/test", + "thinking_blocks": [{"signature": "signed-block"}], + } + state.add_input( + [ + { + "type": "function_call_output", + "call_id": "provider-replay", + "output": "provider-backed replay", + "provider_data": provider_data, + } + ] + ) + model.enqueue([get_text_message("Recovered")]) + session.failure = "after" + + with pytest.raises( + RuntimeError, match="conversation append failed|pending input Session append failed" + ): + await Runner.run(agent, state, session=session, run_config=RunConfig(tracing_disabled=True)) + + state = await RunState.from_json(agent, state.to_json()) + assert state._pending_session_write is not None + pending_item = state._pending_session_write["items"][0] + assert isinstance(pending_item, dict) + assert pending_item["provider_data"] == provider_data + + result = await Runner.run( + agent, state, session=session, run_config=RunConfig(tracing_disabled=True) + ) + + assert result.final_output == "Recovered" + stored = await session.get_items() + matches = [ + item + for item in stored + if isinstance(item, dict) and item.get("call_id") == "provider-replay" + ] + assert len(matches) == 1 + assert matches[0].get("provider_data") == provider_data + + @pytest.mark.asyncio async def test_pending_input_added_during_session_write_survives_stream_checkpoint() -> None: session = _PendingInputWriteFailureSession()