From 7bb9415c5b3f695b97f6409930b92e8a9f2b9e1e Mon Sep 17 00:00:00 2001 From: Kazuhiro Sera Date: Tue, 11 Aug 2026 19:21:22 +0900 Subject: [PATCH] feat: add approval argument overrides for RunState resumes Co-authored-by: Kazuhiro Sera --- integration_tests/_contract_support.py | 21 +- src/agents/memory/__init__.py | 16 + .../memory/openai_conversations_session.py | 1 + .../openai_responses_compaction_session.py | 298 ++-- src/agents/memory/session.py | 128 +- src/agents/memory/sqlite_session.py | 70 +- src/agents/run.py | 59 +- .../run_internal/agent_runner_helpers.py | 55 +- src/agents/run_internal/run_loop.py | 34 +- .../run_internal/session_persistence.py | 182 +- src/agents/run_internal/turn_resolution.py | 2 +- src/agents/run_state.py | 821 ++++++++- tests/fixtures/run_state/README.md | 6 +- .../v1_16_approval_argument_overrides.json | 280 +++ ...6_durable_approval_argument_overrides.json | 299 ++++ tests/fixtures/run_state/generate_corpus.py | 108 +- tests/fixtures/run_state/minimal/v1_16.json | 62 + tests/fixtures/run_state/sources.json | 19 + ...est_openai_responses_compaction_session.py | 821 ++++++++- tests/memory/test_session.py | 334 ++++ tests/test_agent_runner.py | 1497 +++++++++++++++-- tests/test_agent_runner_streamed.py | 174 +- tests/test_run_state.py | 618 +++++++ tests/test_run_state_compatibility_corpus.py | 54 +- tests/utils/simple_session.py | 25 +- 25 files changed, 5664 insertions(+), 320 deletions(-) create mode 100644 tests/fixtures/run_state/features/v1_16_approval_argument_overrides.json create mode 100644 tests/fixtures/run_state/features/v1_16_durable_approval_argument_overrides.json create mode 100644 tests/fixtures/run_state/minimal/v1_16.json diff --git a/integration_tests/_contract_support.py b/integration_tests/_contract_support.py index dcc437a3c2..9f619fcf6d 100644 --- a/integration_tests/_contract_support.py +++ b/integration_tests/_contract_support.py @@ -1649,7 +1649,7 @@ def _find_subset_errors(expected: object, actual: object, path: str = "state") - def _restore_agent(payload: dict[str, Any]) -> Any: - from agents import Agent, handoff + from agents import Agent, function_tool, handoff current_agent = payload.get("current_agent") name = ( @@ -1661,6 +1661,25 @@ def _restore_agent(payload: dict[str, Any]) -> Any: if identity == f"{name}#2": duplicate = Agent(name=name) return Agent(name=name, handoffs=[handoff(duplicate)]) + + processed_response = payload.get("last_processed_response") + serialized_functions = ( + processed_response.get("functions", []) if isinstance(processed_response, dict) else [] + ) + has_send_email_tool = any( + isinstance(entry, dict) + and isinstance(entry.get("tool"), dict) + and entry["tool"].get("name") == "send_email" + for entry in serialized_functions + ) + if has_send_email_tool: + + @function_tool(name_override="send_email") + def send_email(recipient: str) -> str: + return f"sent:{recipient}" + + return Agent(name=name, tools=[send_email]) + return Agent(name=name) diff --git a/src/agents/memory/__init__.py b/src/agents/memory/__init__.py index bb5c7356f7..f71ce41206 100644 --- a/src/agents/memory/__init__.py +++ b/src/agents/memory/__init__.py @@ -5,11 +5,19 @@ from .openai_conversations_session import OpenAIConversationsSession from .openai_responses_compaction_session import OpenAIResponsesCompactionSession from .session import ( + SERVER_MANAGED_CONVERSATION_SESSION_ATTR, OpenAIResponsesCompactionArgs, OpenAIResponsesCompactionAwareSession, + ServerManagedConversationSession, Session, SessionABC, + SessionHistoryMutation, + SessionHistoryRewriteArgs, + SessionHistoryRewriteAwareSession, + apply_session_history_mutations, is_openai_responses_compaction_aware_session, + is_server_managed_conversation_session, + is_session_history_rewrite_aware_session, ) from .session_settings import SessionSettings from .util import SessionInputCallback @@ -27,7 +35,15 @@ "OpenAIResponsesCompactionSession", "OpenAIResponsesCompactionArgs", "OpenAIResponsesCompactionAwareSession", + "SERVER_MANAGED_CONVERSATION_SESSION_ATTR", + "SessionHistoryMutation", + "SessionHistoryRewriteArgs", + "SessionHistoryRewriteAwareSession", + "ServerManagedConversationSession", + "apply_session_history_mutations", + "is_server_managed_conversation_session", "is_openai_responses_compaction_aware_session", + "is_session_history_rewrite_aware_session", ] diff --git a/src/agents/memory/openai_conversations_session.py b/src/agents/memory/openai_conversations_session.py index 8e0641067c..612c3cddb6 100644 --- a/src/agents/memory/openai_conversations_session.py +++ b/src/agents/memory/openai_conversations_session.py @@ -25,6 +25,7 @@ async def start_openai_conversations_session(openai_client: AsyncOpenAI | None = class OpenAIConversationsSession(SessionABC): + _server_managed_conversation_session = True session_settings: SessionSettings | None = None def __init__( diff --git a/src/agents/memory/openai_responses_compaction_session.py b/src/agents/memory/openai_responses_compaction_session.py index 1e8456bb46..91324118bf 100644 --- a/src/agents/memory/openai_responses_compaction_session.py +++ b/src/agents/memory/openai_responses_compaction_session.py @@ -11,11 +11,14 @@ from ..logger import log_model_and_tool_action_warning from ..models._openai_shared import get_default_openai_client from ..run_internal.items import normalize_input_items_for_api -from .openai_conversations_session import OpenAIConversationsSession from .session import ( OpenAIResponsesCompactionArgs, OpenAIResponsesCompactionAwareSession, SessionABC, + SessionHistoryRewriteArgs, + apply_session_history_mutations, + is_server_managed_conversation_session, + is_session_history_rewrite_aware_session, ) if TYPE_CHECKING: @@ -80,11 +83,13 @@ def is_openai_model_name(model: str) -> bool: class OpenAIResponsesCompactionSession(SessionABC, OpenAIResponsesCompactionAwareSession): """Session decorator that triggers responses.compact when stored history grows. - Works with OpenAI Responses API models only. Wraps any Session (except - OpenAIConversationsSession) and automatically calls the OpenAI responses.compact - API after each turn when the decision hook returns True. + Works with OpenAI Responses API models only. Wraps any client-managed Session and + automatically calls the OpenAI responses.compact API after each turn when the decision + hook returns True. """ + supports_expected_history_mutations: Literal[True] = True + def __init__( self, session_id: str, @@ -99,8 +104,7 @@ def __init__( Args: session_id: Identifier for this session. - underlying_session: Session store that holds the compacted history. Cannot be - OpenAIConversationsSession. + underlying_session: Client-managed session store that holds the compacted history. client: OpenAI client for responses.compact API calls. Defaults to get_default_openai_client() or new AsyncOpenAI(). model: Model to use for responses.compact. Defaults to "gpt-4.1". Must be an @@ -111,10 +115,10 @@ def __init__( should_trigger_compaction: Custom decision hook. Defaults to triggering when 10+ compaction candidates exist. """ - if isinstance(underlying_session, OpenAIConversationsSession): + if is_server_managed_conversation_session(underlying_session): raise ValueError( - "OpenAIResponsesCompactionSession cannot wrap OpenAIConversationsSession " - "because it manages its own history on the server." + "OpenAIResponsesCompactionSession cannot wrap a server-managed conversation " + "session because it manages its own history on the server." ) if not is_openai_model_name(model): @@ -137,6 +141,8 @@ def __init__( self._response_id: str | None = None self._deferred_response_id: str | None = None self._last_unstored_response_id: str | None = None + self._has_pending_local_history_rewrite = False + self._history_rewrite_generation = 0 # Serialize wrapper mutations against compaction snapshot/replace/restore so a # cancellation rollback cannot rewrite past a newer concurrent write. self._mutation_lock = asyncio.Lock() @@ -162,67 +168,75 @@ def _resolve_compaction_mode_for_response( and response_id is not None and response_id == self._last_unstored_response_id ): + resolved_mode: _ResolvedCompactionMode = "input" + else: + resolved_mode = _resolve_compaction_mode(mode, response_id=response_id, store=store) + if self._has_pending_local_history_rewrite and resolved_mode == "previous_response_id": + logger.debug("compact: forcing input mode after local history rewrite") return "input" - return _resolve_compaction_mode(mode, response_id=response_id, store=store) + return resolved_mode async def run_compaction(self, args: OpenAIResponsesCompactionArgs | None = None) -> None: """Run compaction using responses.compact API.""" - if args and args.get("response_id"): - self._response_id = args["response_id"] - requested_mode = args.get("compaction_mode") if args else None - if args and "store" in args: - store = args["store"] - if store is False and self._response_id: - self._last_unstored_response_id = self._response_id - elif store is True and self._response_id == self._last_unstored_response_id: - self._last_unstored_response_id = None - else: - store = None - resolved_mode = self._resolve_compaction_mode_for_response( - response_id=self._response_id, - store=store, - requested_mode=requested_mode, - ) - - if resolved_mode == "previous_response_id" and not self._response_id: - raise ValueError( - "OpenAIResponsesCompactionSession.run_compaction requires a response_id " - "when using previous_response_id compaction." + async with self._mutation_lock: + if args and args.get("response_id"): + self._response_id = args["response_id"] + requested_mode = args.get("compaction_mode") if args else None + if args and "store" in args: + store = args["store"] + if store is False and self._response_id: + self._last_unstored_response_id = self._response_id + elif store is True and self._response_id == self._last_unstored_response_id: + self._last_unstored_response_id = None + else: + store = None + resolved_mode = self._resolve_compaction_mode_for_response( + response_id=self._response_id, + store=store, + requested_mode=requested_mode, ) - compaction_candidate_items, session_items = await self._ensure_compaction_candidates() + if resolved_mode == "previous_response_id" and not self._response_id: + raise ValueError( + "OpenAIResponsesCompactionSession.run_compaction requires a response_id " + "when using previous_response_id compaction." + ) + + compaction_candidate_items, session_items = await self._ensure_compaction_candidates() + + force = args.get("force", False) if args else False + should_compact = force or self.should_trigger_compaction( + { + "response_id": self._response_id, + "compaction_mode": resolved_mode, + "compaction_candidate_items": compaction_candidate_items, + "session_items": session_items, + } + ) - force = args.get("force", False) if args else False - should_compact = force or self.should_trigger_compaction( - { - "response_id": self._response_id, - "compaction_mode": resolved_mode, - "compaction_candidate_items": compaction_candidate_items, - "session_items": session_items, - } - ) + if not should_compact: + logger.debug( + "skip: decision hook declined compaction for %s (mode=%s)", + self._response_id, + resolved_mode, + ) + return - if not should_compact: + self._deferred_response_id = None + response_id = self._response_id + rewrite_generation = self._history_rewrite_generation logger.debug( - "skip: decision hook declined compaction for %s (mode=%s)", - self._response_id, + "compact: start for %s using %s (mode=%s)", + response_id, + self.model, resolved_mode, ) - return - self._deferred_response_id = None - logger.debug( - "compact: start for %s using %s (mode=%s)", - self._response_id, - self.model, - resolved_mode, - ) - - compact_kwargs: dict[str, Any] = {"model": self.model} - if resolved_mode == "previous_response_id": - compact_kwargs["previous_response_id"] = self._response_id - else: - compact_kwargs["input"] = session_items + compact_kwargs: dict[str, Any] = {"model": self.model} + if resolved_mode == "previous_response_id": + compact_kwargs["previous_response_id"] = response_id + else: + compact_kwargs["input"] = session_items compacted = await self.client.responses.compact(**compact_kwargs) @@ -231,6 +245,12 @@ async def run_compaction(self, args: OpenAIResponsesCompactionArgs | None = None ) async with self._mutation_lock: + if rewrite_generation != self._history_rewrite_generation: + logger.debug( + "compact: discarding stale result for %s after local history mutation", + response_id, + ) + return previous_items = await self._get_all_underlying_session_items() await self._replace_underlying_session_items( output_items=output_items, @@ -238,10 +258,12 @@ async def run_compaction(self, args: OpenAIResponsesCompactionArgs | None = None ) self._compaction_candidate_items = select_compaction_candidate_items(output_items) self._session_items = output_items + if resolved_mode == "input": + self._has_pending_local_history_rewrite = False logger.debug( "compact: done for %s (mode=%s, output=%s, candidates=%s)", - self._response_id, + response_id, resolved_mode, len(output_items), len(self._compaction_candidate_items or []), @@ -250,29 +272,77 @@ async def run_compaction(self, args: OpenAIResponsesCompactionArgs | None = None async def get_items(self, limit: int | None = None) -> list[TResponseInputItem]: return await self.underlying_session.get_items(limit) - async def _get_all_underlying_session_items(self) -> list[TResponseInputItem]: - return await self.underlying_session.get_items(limit=_ALL_SESSION_ITEMS_LIMIT) + async def apply_history_mutations(self, args: SessionHistoryRewriteArgs) -> bool: + """Rewrite persisted history and keep compaction caches aligned.""" + mutations = list(args.get("mutations", [])) + if not mutations: + return True + + async with self._mutation_lock: + underlying_session = self.underlying_session + if is_server_managed_conversation_session(underlying_session): + raise ValueError( + "OpenAIResponsesCompactionSession cannot rewrite a server-managed " + "conversation session because it manages its own history on the server." + ) + # Invalidate any compaction that started before this potentially committing + # rewrite. This evidence must precede the first cancellable storage await. + self._history_rewrite_generation += 1 + self._has_pending_local_history_rewrite = True + self._session_items = None + self._compaction_candidate_items = None + if is_session_history_rewrite_aware_session(underlying_session): + applied = await underlying_session.apply_history_mutations({"mutations": mutations}) + if applied is not True: + raise ValueError("Underlying session did not confirm its history rewrite.") + rewritten_items = await self._get_all_underlying_session_items(underlying_session) + else: + previous_items = await self._get_all_underlying_session_items(underlying_session) + rewritten_items = apply_session_history_mutations(previous_items, mutations) + await self._replace_underlying_session_items( + output_items=rewritten_items, + previous_items=previous_items, + underlying_session=underlying_session, + ) + if self.underlying_session is not underlying_session: + raise ValueError( + "OpenAIResponsesCompactionSession underlying session changed during history " + "rewrite." + ) + self._session_items = rewritten_items + self._compaction_candidate_items = select_compaction_candidate_items(rewritten_items) + return True + + async def _get_all_underlying_session_items( + self, + underlying_session: Session | None = None, + ) -> list[TResponseInputItem]: + session = underlying_session if underlying_session is not None else self.underlying_session + return await session.get_items(limit=_ALL_SESSION_ITEMS_LIMIT) async def _replace_underlying_session_items( self, *, output_items: list[TResponseInputItem], previous_items: list[TResponseInputItem], + underlying_session: Session | None = None, ) -> None: # Treat clear → add as one replacement transaction. Exception and CancelledError # both restore previous history, and restore settlement is always drained so a # cancel during restore cannot leave an empty session. + session = underlying_session if underlying_session is not None else self.underlying_session cleared = False try: - await self.underlying_session.clear_session() + await session.clear_session() cleared = True if output_items: - await self.underlying_session.add_items(output_items) + await session.add_items(output_items) except Exception as error: await self._recover_from_failed_replacement( previous_items=previous_items, error=error, cleared=cleared, + underlying_session=session, ) raise except asyncio.CancelledError as error: @@ -280,6 +350,7 @@ async def _replace_underlying_session_items( previous_items=previous_items, error=error, cleared=cleared, + underlying_session=session, ) raise @@ -289,13 +360,20 @@ async def _recover_from_failed_replacement( previous_items: list[TResponseInputItem], error: BaseException, cleared: bool, + underlying_session: Session, ) -> None: if not cleared: restore = self._restore_underlying_session_items_after_failed_clear( - previous_items, error + previous_items, + error, + underlying_session=underlying_session, ) else: - restore = self._restore_underlying_session_items(previous_items, error) + restore = self._restore_underlying_session_items( + previous_items, + error, + underlying_session=underlying_session, + ) await self._await_restore_despite_cancellation(restore) async def _await_restore_despite_cancellation(self, restore: Awaitable[None]) -> None: @@ -324,9 +402,11 @@ async def _restore_underlying_session_items_after_failed_clear( self, previous_items: list[TResponseInputItem], clear_error: BaseException, + *, + underlying_session: Session, ) -> None: try: - current_items = await self._get_all_underlying_session_items() + current_items = await self._get_all_underlying_session_items(underlying_session) except Exception as inspection_error: log_model_and_tool_action_warning( logger, @@ -339,7 +419,10 @@ async def _restore_underlying_session_items_after_failed_clear( return await self._restore_underlying_session_items( - previous_items, clear_error, clear_existing_items=False + previous_items, + clear_error, + clear_existing_items=False, + underlying_session=underlying_session, ) async def _restore_underlying_session_items( @@ -348,12 +431,13 @@ async def _restore_underlying_session_items( replacement_error: BaseException, *, clear_existing_items: bool = True, + underlying_session: Session, ) -> None: try: if clear_existing_items: - await self.underlying_session.clear_session() + await underlying_session.clear_session() if previous_items: - await self.underlying_session.add_items(list(previous_items)) + await underlying_session.add_items(list(previous_items)) except Exception as restore_error: log_model_and_tool_action_warning( logger, @@ -369,24 +453,25 @@ async def _restore_underlying_session_items( ) async def _defer_compaction(self, response_id: str, store: bool | None = None) -> None: - if self._deferred_response_id is not None: - return - compaction_candidate_items, session_items = await self._ensure_compaction_candidates() - resolved_mode = self._resolve_compaction_mode_for_response( - response_id=response_id, - store=store, - requested_mode=None, - ) - should_compact = self.should_trigger_compaction( - { - "response_id": response_id, - "compaction_mode": resolved_mode, - "compaction_candidate_items": compaction_candidate_items, - "session_items": session_items, - } - ) - if should_compact: - self._deferred_response_id = response_id + async with self._mutation_lock: + if self._deferred_response_id is not None: + return + compaction_candidate_items, session_items = await self._ensure_compaction_candidates() + resolved_mode = self._resolve_compaction_mode_for_response( + response_id=response_id, + store=store, + requested_mode=None, + ) + should_compact = self.should_trigger_compaction( + { + "response_id": response_id, + "compaction_mode": resolved_mode, + "compaction_candidate_items": compaction_candidate_items, + "session_items": session_items, + } + ) + if should_compact: + self._deferred_response_id = response_id def _get_deferred_compaction_response_id(self) -> str | None: return self._deferred_response_id @@ -396,29 +481,54 @@ def _clear_deferred_compaction(self) -> None: async def add_items(self, items: list[TResponseInputItem]) -> None: async with self._mutation_lock: + cached_candidates = ( + list(self._compaction_candidate_items) + if self._compaction_candidate_items is not None + else None + ) + cached_items = list(self._session_items) if self._session_items is not None else None + # Invalidate any compaction and cached history before the append can commit and + # propagate cancellation. + self._history_rewrite_generation += 1 + self._compaction_candidate_items = None + self._session_items = None await self.underlying_session.add_items(items) - if self._compaction_candidate_items is not None: + if cached_candidates is not None: new_items = _normalize_compaction_session_items(items) new_candidates = select_compaction_candidate_items(new_items) if new_candidates: - self._compaction_candidate_items.extend(new_candidates) - if self._session_items is not None: - self._session_items.extend(_normalize_compaction_session_items(items)) + cached_candidates.extend(new_candidates) + self._compaction_candidate_items = cached_candidates + if cached_items is not None: + cached_items.extend(_normalize_compaction_session_items(items)) + self._session_items = cached_items async def pop_item(self) -> TResponseInputItem | None: async with self._mutation_lock: + cached_candidates = self._compaction_candidate_items + cached_items = self._session_items + # Invalidate any compaction and cached history before the pop can commit and + # propagate cancellation. + self._history_rewrite_generation += 1 + self._compaction_candidate_items = None + self._session_items = None popped = await self.underlying_session.pop_item() - if popped: - self._compaction_candidate_items = None - self._session_items = None + if popped is None: + self._compaction_candidate_items = cached_candidates + self._session_items = cached_items return popped async def clear_session(self) -> None: async with self._mutation_lock: + # Invalidate compaction before the clear can commit and propagate cancellation. + self._history_rewrite_generation += 1 + self._compaction_candidate_items = None + self._session_items = None await self.underlying_session.clear_session() self._compaction_candidate_items = [] self._session_items = [] self._deferred_response_id = None + self._has_pending_local_history_rewrite = False async def _ensure_compaction_candidates( self, diff --git a/src/agents/memory/session.py b/src/agents/memory/session.py index 26690c2c71..2eb3e85a23 100644 --- a/src/agents/memory/session.py +++ b/src/agents/memory/session.py @@ -1,8 +1,9 @@ from __future__ import annotations +import copy import inspect from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Literal, Protocol, TypeGuard, runtime_checkable +from typing import TYPE_CHECKING, Any, Literal, Protocol, TypeGuard, cast, runtime_checkable from typing_extensions import TypedDict @@ -11,6 +12,8 @@ from ..run_context import RunContextWrapper from .session_settings import SessionSettings +SERVER_MANAGED_CONVERSATION_SESSION_ATTR = "_server_managed_conversation_session" + @runtime_checkable class Session(Protocol): @@ -106,6 +109,129 @@ async def clear_session(self) -> None: ... +@runtime_checkable +class ServerManagedConversationSession(Session, Protocol): + """Protocol for sessions whose canonical history is managed by a remote service.""" + + _server_managed_conversation_session: Literal[True] + + +def is_server_managed_conversation_session( + session: Session | None, +) -> TypeGuard[ServerManagedConversationSession]: + """Check whether a session advertises server-managed history semantics.""" + if session is None: + return False + try: + marker = inspect.getattr_static(session, SERVER_MANAGED_CONVERSATION_SESSION_ATTR, False) + except Exception: + return False + return marker is True + + +class ReplaceFunctionCallSessionHistoryMutation(TypedDict): + """Replace the canonical persisted function call for a tool call.""" + + type: Literal["replace_function_call"] + call_id: str + expected: TResponseInputItem + replacement: TResponseInputItem + + +SessionHistoryMutation = ReplaceFunctionCallSessionHistoryMutation + + +class SessionHistoryRewriteArgs(TypedDict): + """Arguments for persisted-history rewrites.""" + + mutations: list[SessionHistoryMutation] + + +@runtime_checkable +class SessionHistoryRewriteAwareSession(Session, Protocol): + """Protocol for sessions that can compare and rewrite persisted function calls.""" + + supports_expected_history_mutations: Literal[True] + + async def apply_history_mutations(self, args: SessionHistoryRewriteArgs) -> bool: + """Apply every expected mutation and confirm all targets were reconciled.""" + ... + + +def is_session_history_rewrite_aware_session( + session: Session | None, +) -> TypeGuard[SessionHistoryRewriteAwareSession]: + """Check whether a session supports expected persisted-history rewrites.""" + if session is None: + return False + try: + apply_history_mutations = inspect.getattr_static(session, "apply_history_mutations", None) + supports_expected_mutations = inspect.getattr_static( + session, "supports_expected_history_mutations", False + ) + except Exception: + return False + return callable(apply_history_mutations) and supports_expected_mutations is True + + +def apply_session_history_mutations( + items: list[TResponseInputItem], + mutations: list[SessionHistoryMutation], +) -> list[TResponseInputItem]: + """Apply structured history mutations to a list of persisted session items.""" + next_items = list(items) + for mutation in mutations: + if mutation["type"] == "replace_function_call": + next_items = _apply_replace_function_call_mutation(next_items, mutation) + return next_items + + +def _apply_replace_function_call_mutation( + items: list[TResponseInputItem], + mutation: ReplaceFunctionCallSessionHistoryMutation, +) -> list[TResponseInputItem]: + """Replace the latest expected call or accept an already-applied replacement.""" + call_id = mutation["call_id"] + expected = _snapshot_matching_function_call(mutation["expected"], call_id) + replacement = _snapshot_matching_function_call(mutation["replacement"], call_id) + if expected is None or replacement is None: + raise ValueError("Session history mutation contains an invalid function call.") + + for index in range(len(items) - 1, -1, -1): + candidate = _snapshot_matching_function_call(items[index], call_id) + if candidate is None: + continue + if candidate == replacement: + return list(items) + if candidate == expected: + next_items = list(items) + next_items[index] = cast(Any, replacement) + return next_items + break + + raise ValueError("Session history mutation target did not match the expected function call.") + + +def _snapshot_matching_function_call(item: Any, call_id: str) -> dict[str, Any] | None: + """Snapshot a matching function call without inspecting unrelated history values.""" + if isinstance(item, dict): + payload = item + elif hasattr(item, "model_dump"): + try: + payload = item.model_dump(exclude_unset=True) + except TypeError: + payload = item.model_dump() + else: + return None + if ( + not isinstance(payload, dict) + or payload.get("type") != "function_call" + or payload.get("call_id") != call_id + ): + return None + return copy.deepcopy(payload) + + class OpenAIResponsesCompactionArgs(TypedDict, total=False): """Arguments for the run_compaction method.""" diff --git a/src/agents/memory/sqlite_session.py b/src/agents/memory/sqlite_session.py index fc9f3fdb8f..0fa9e05776 100644 --- a/src/agents/memory/sqlite_session.py +++ b/src/agents/memory/sqlite_session.py @@ -7,10 +7,10 @@ from collections.abc import Awaitable, Iterator from contextlib import contextmanager from pathlib import Path -from typing import Any, ClassVar, TypeVar +from typing import Any, ClassVar, Literal, TypeVar from ..items import TResponseInputItem -from .session import SessionABC +from .session import SessionABC, SessionHistoryRewriteArgs, apply_session_history_mutations from .session_settings import SessionSettings, coerce_session_settings, resolve_session_limit _T = TypeVar("_T") @@ -47,6 +47,7 @@ class SQLiteSession(SessionABC): """ session_settings: SessionSettings | None = None + supports_expected_history_mutations: Literal[True] = True _file_locks: ClassVar[dict[Path, threading.RLock]] = {} _file_lock_counts: ClassVar[dict[Path, int]] = {} _file_locks_guard: ClassVar[threading.Lock] = threading.Lock() @@ -429,6 +430,71 @@ def _clear_session_sync(): await _await_mutation(asyncio.to_thread(_clear_session_sync)) + async def apply_history_mutations(self, args: SessionHistoryRewriteArgs) -> bool: + """Rewrite persisted session history using structured mutations.""" + mutations = list(args.get("mutations", [])) + if not mutations: + return True + + def _apply_history_mutations_sync() -> None: + with self._write_connection() as conn: + cursor = conn.execute( + f""" + SELECT id, message_data FROM {self.messages_table} + WHERE session_id = ? + ORDER BY id ASC + """, + (self.session_id,), + ) + rows = cursor.fetchall() + + decoded_rows: list[tuple[int, TResponseInputItem]] = [] + for row_id, message_data in rows: + try: + item = json.loads(message_data) + except (json.JSONDecodeError, TypeError): + continue + decoded_rows.append((row_id, item)) + + for mutation in mutations: + previous_items = [item for _, item in decoded_rows] + rewritten_items = apply_session_history_mutations( + previous_items, + [mutation], + ) + for (row_id, previous_item), rewritten_item in zip( + decoded_rows, rewritten_items, strict=True + ): + if previous_item == rewritten_item: + continue + conn.execute( + f""" + UPDATE {self.messages_table} + SET message_data = ? + WHERE id = ? AND session_id = ? + """, + (json.dumps(rewritten_item), row_id, self.session_id), + ) + decoded_rows = [ + (row_id, rewritten_item) + for (row_id, _), rewritten_item in zip( + decoded_rows, rewritten_items, strict=True + ) + ] + + conn.execute( + f""" + UPDATE {self.sessions_table} + SET updated_at = CURRENT_TIMESTAMP + WHERE session_id = ? + """, + (self.session_id,), + ) + conn.commit() + + await _await_mutation(asyncio.to_thread(_apply_history_mutations_sync)) + return True + def close(self) -> None: """Close the database connection.""" with self._lock: diff --git a/src/agents/run.py b/src/agents/run.py index c2c4839fb3..f97b995e50 100644 --- a/src/agents/run.py +++ b/src/agents/run.py @@ -74,6 +74,7 @@ snapshot_usage, update_run_state_for_interruption, usage_delta, + validate_override_history_persistence_support, validate_session_conversation_settings, ) from .run_internal.approvals import approvals_from_step @@ -117,6 +118,7 @@ from .run_internal.session_persistence import ( _session_get_items, admit_pending_input, + apply_pending_session_history_mutations, commit_server_pending_input, persist_session_items_for_guardrail_trip, prepare_input_with_session, @@ -543,6 +545,7 @@ async def run( session_input_items_for_persistence: list[TResponseInputItem] | None = ( [] if (session is not None and is_resumed_state) else None ) + server_manages_conversation = False # Track the most recent input batch we persisted so conversation-lock retries can rewind # exactly those items (and not the full history). last_saved_input_snapshot_for_rewind: list[TResponseInputItem] | None = None @@ -622,6 +625,16 @@ async def run( ) original_input_for_state = prepared_input + server_manages_conversation = ( + conversation_id is not None + or previous_response_id is not None + or auto_previous_response_id + ) + validate_override_history_persistence_support( + input=input, + session=session, + response_history_is_server_managed=server_manages_conversation, + ) # Check whether to enable OpenAI server-managed conversation if ( conversation_id is not None @@ -998,20 +1011,31 @@ def _mark_response_hooks_started() -> None: ) raise UserError("No processed response found in previous state") - turn_result = await resolve_interrupted_turn( - bindings=current_bindings, - original_input=original_input, - original_pre_step_items=generated_items, - new_response=run_state._model_responses[-1], - processed_response=run_state._last_processed_response, - hooks=hooks, - context_wrapper=context_wrapper, - run_config=run_config, - server_manages_conversation=server_conversation_tracker is not None, - run_state=run_state, - error_handlers=error_handlers, + await apply_pending_session_history_mutations( + session, + run_state, + wrapper=context_wrapper, ) + try: + turn_result = await resolve_interrupted_turn( + bindings=current_bindings, + original_input=original_input, + original_pre_step_items=generated_items, + new_response=run_state._model_responses[-1], + processed_response=run_state._last_processed_response, + hooks=hooks, + context_wrapper=context_wrapper, + run_config=run_config, + server_manages_conversation=( + server_conversation_tracker is not None + ), + run_state=run_state, + error_handlers=error_handlers, + ) + finally: + run_state._clear_executed_approval_argument_overrides() + if run_state._last_processed_response is not None: tool_use_tracker.record_processed_response( current_agent, @@ -1981,6 +2005,7 @@ def run_streamed( run_state: RunState[TContext] | None = None input_for_result: str | list[TResponseInputItem] starting_input = input if not is_resumed_state else None + server_manages_conversation = False if is_resumed_state: run_state = cast(RunState[TContext], input) @@ -2058,6 +2083,16 @@ def run_streamed( auto_previous_response_id=auto_previous_response_id, ) + server_manages_conversation = ( + conversation_id is not None + or previous_response_id is not None + or auto_previous_response_id + ) + validate_override_history_persistence_support( + input=input, + session=session, + response_history_is_server_managed=server_manages_conversation, + ) resolved_reasoning_item_id_policy: ReasoningItemIdPolicy | None = ( run_config.reasoning_item_id_policy if run_config.reasoning_item_id_policy is not None diff --git a/src/agents/run_internal/agent_runner_helpers.py b/src/agents/run_internal/agent_runner_helpers.py index f4dca6b1ac..0c60b5d1ed 100644 --- a/src/agents/run_internal/agent_runner_helpers.py +++ b/src/agents/run_internal/agent_runner_helpers.py @@ -12,7 +12,11 @@ from ..exceptions import UserError from ..guardrail import InputGuardrailResult from ..items import ModelResponse, RunItem, ToolApprovalItem, TResponseInputItem -from ..memory import Session +from ..memory import ( + Session, + is_server_managed_conversation_session, + is_session_history_rewrite_aware_session, +) from ..models.openai_agent_registration import add_openai_harness_id_to_metadata from ..result import RunResult from ..run_config import RunConfig @@ -63,6 +67,7 @@ "save_turn_items_if_needed", "should_cancel_parallel_model_task_on_input_guardrail_trip", "update_run_state_for_interruption", + "validate_override_history_persistence_support", ] _PARALLEL_INPUT_GUARDRAIL_CANCEL_PATCH_ID = ( @@ -257,6 +262,54 @@ def validate_session_conversation_settings( ) +def validate_override_history_persistence_support( + *, + input: str | list[TResponseInputItem] | RunState[Any], + session: Session | None, + response_history_is_server_managed: bool, +) -> None: + """Fail fast when approval override persistence requirements are not satisfied.""" + if not isinstance(input, RunState): + return + + session_history_is_server_managed = is_server_managed_conversation_session(session) + canonical_history_is_server_managed = ( + response_history_is_server_managed or session_history_is_server_managed + ) + + if ( + input._has_pending_execution_only_approval_overrides() + and not canonical_history_is_server_managed + ): + raise UserError( + "save_override_arguments=False is only supported when canonical history is managed " + "by a server, such as with OpenAIConversationsSession, conversation_id, " + "previous_response_id, or auto_previous_response_id." + ) + + mutations = input._get_session_history_mutations() + if not mutations: + return + + if canonical_history_is_server_managed: + raise UserError( + "save_override_arguments requires local canonical history. " + "Server-managed conversations cannot persist corrected function_call arguments. " + "Pass save_override_arguments=False to apply the override only to the current " + "execution." + ) + + if session is None or is_session_history_rewrite_aware_session(session): + return + + raise UserError( + "save_override_arguments requires a session that supports expected history rewrites. " + "Use SQLiteSession, OpenAIResponsesCompactionSession, or another " + "SessionHistoryRewriteAwareSession. The configured session has no safe execution-only " + "alternative because it owns the local canonical history." + ) + + def resolve_trace_settings( *, run_state: RunState[TContext] | None, diff --git a/src/agents/run_internal/run_loop.py b/src/agents/run_internal/run_loop.py index 4a2f9141d3..ba32f26e3b 100644 --- a/src/agents/run_internal/run_loop.py +++ b/src/agents/run_internal/run_loop.py @@ -152,6 +152,7 @@ from .session_persistence import ( _session_get_items, admit_pending_input, + apply_pending_session_history_mutations, commit_server_pending_input, persist_session_items_for_guardrail_trip, prepare_input_with_session, @@ -1017,20 +1018,29 @@ async def _save_stream_items_without_count( last_model_response = run_state._model_responses[-1] - turn_result = await resolve_interrupted_turn( - bindings=current_bindings, - original_input=run_state._original_input, - original_pre_step_items=run_state._generated_items, - new_response=last_model_response, - processed_response=run_state._last_processed_response, - hooks=hooks, - context_wrapper=context_wrapper, - run_config=run_config, - server_manages_conversation=server_conversation_tracker is not None, - run_state=run_state, - error_handlers=error_handlers, + await apply_pending_session_history_mutations( + session, + run_state, + wrapper=context_wrapper, ) + try: + turn_result = await resolve_interrupted_turn( + bindings=current_bindings, + original_input=run_state._original_input, + original_pre_step_items=run_state._generated_items, + new_response=last_model_response, + processed_response=run_state._last_processed_response, + hooks=hooks, + context_wrapper=context_wrapper, + run_config=run_config, + server_manages_conversation=server_conversation_tracker is not None, + run_state=run_state, + error_handlers=error_handlers, + ) + finally: + run_state._clear_executed_approval_argument_overrides() + tool_use_tracker.record_processed_response( current_agent, run_state._last_processed_response ) diff --git a/src/agents/run_internal/session_persistence.py b/src/agents/run_internal/session_persistence.py index 44c39d5cef..1a091a3ec2 100644 --- a/src/agents/run_internal/session_persistence.py +++ b/src/agents/run_internal/session_persistence.py @@ -32,9 +32,11 @@ from ..memory import ( OpenAIResponsesCompactionArgs, Session, + SessionHistoryMutation, SessionInputCallback, SessionSettings, is_openai_responses_compaction_aware_session, + is_session_history_rewrite_aware_session, ) from ..memory.openai_conversations_session import OpenAIConversationsSession from ..memory.session import _call_session_method, _get_session_wrapper @@ -63,6 +65,7 @@ from .run_steps import NextStepInterruption, ProcessedResponse, SingleStepResult __all__ = [ + "apply_pending_session_history_mutations", "admit_pending_input", "commit_server_pending_input", "prepare_input_with_session", @@ -540,6 +543,7 @@ def update_run_state_after_resume( if session_items is not None: run_state._session_items = list(session_items) run_state._current_step = turn_result.next_step # type: ignore[assignment] + run_state._clear_executed_approval_argument_overrides() async def save_result_to_session( @@ -563,6 +567,8 @@ async def save_result_to_session( already_persisted = run_state._current_turn_persisted_item_count if run_state is not None else 0 if session is None: + if run_state is not None: + run_state._clear_session_history_mutations() return 0 wrapper = _get_session_wrapper(session, wrapper) @@ -643,58 +649,135 @@ async def save_result_to_session( ] if len(items_to_save) == 0: + has_pending_history_mutations = bool( + run_state is not None and run_state._get_session_history_mutations() + ) + if has_pending_history_mutations: + await apply_pending_session_history_mutations( + session, + run_state, + wrapper=wrapper, + ) + await _run_compaction_on_session( + session=session, + response_id=response_id, + new_items=new_items, + store=store, + wrapper=wrapper, + ) if run_state is not None: run_state._current_turn_persisted_item_count = already_persisted + saved_run_items_count return saved_run_items_count await _session_add_items(session, items_to_save, wrapper=wrapper) + await apply_pending_session_history_mutations( + session, + run_state, + wrapper=wrapper, + ) if run_state is not None: run_state._current_turn_persisted_item_count = already_persisted + saved_run_items_count - if response_id and is_openai_responses_compaction_aware_session(session): - has_local_tool_outputs = any( - isinstance(item, ToolCallOutputItem | HandoffOutputItem) for item in new_items + await _run_compaction_on_session( + session=session, + response_id=response_id, + new_items=new_items, + store=store, + wrapper=wrapper, + ) + + return saved_run_items_count + + +async def apply_pending_session_history_mutations( + session: Session | None, + run_state: RunState | None, + *, + wrapper: RunContextWrapper[Any] | None = None, +) -> None: + """Apply pending history rewrites before approved tools execute.""" + if run_state is None: + return + + mutations = run_state._get_session_history_mutations() + if not mutations: + return + + if session is None: + return + + normalized_mutations = _normalize_history_mutations_for_session_persistence(session, mutations) + + if not is_session_history_rewrite_aware_session(session): + raise UserError( + "Cannot persist approval argument overrides because the session does not support " + "expected history rewrites." ) - if has_local_tool_outputs: - defer_compaction = getattr(session, "_defer_compaction", None) - if callable(defer_compaction): - await _call_session_method( - defer_compaction, - response_id, - store=store, - wrapper=wrapper, - ) - logger.debug( - "skip: deferring compaction for response %s due to local tool outputs", - response_id, - ) - return saved_run_items_count - - deferred_response_id = None - get_deferred = getattr(session, "_get_deferred_compaction_response_id", None) - if callable(get_deferred): - deferred_response_id = get_deferred() - force_compaction = deferred_response_id is not None - if force_compaction: - logger.debug( - "compact: forcing for response %s after deferred %s", + + applied = await _call_session_method( + session.apply_history_mutations, + {"mutations": normalized_mutations}, + wrapper=wrapper, + ) + if applied is not True: + raise UserError( + "Cannot persist approval argument overrides because the session did not confirm " + "that every target call was rewritten." + ) + + +async def _run_compaction_on_session( + *, + session: Session, + response_id: str | None, + new_items: list[RunItem], + store: bool | None, + wrapper: RunContextWrapper[Any] | None = None, +) -> None: + """Run session compaction hooks after persistence or approval-only mutation cycles.""" + if not response_id or not is_openai_responses_compaction_aware_session(session): + return + + has_local_tool_outputs = any( + isinstance(item, ToolCallOutputItem | HandoffOutputItem) for item in new_items + ) + if has_local_tool_outputs: + defer_compaction = getattr(session, "_defer_compaction", None) + if callable(defer_compaction): + await _call_session_method( + defer_compaction, response_id, - deferred_response_id, + store=store, + wrapper=wrapper, ) - compaction_args: OpenAIResponsesCompactionArgs = { - "response_id": response_id, - "force": force_compaction, - } - if store is not None: - compaction_args["store"] = store - await _call_session_method( - session.run_compaction, - compaction_args, - wrapper=wrapper, + logger.debug( + "skip: deferring compaction for response %s due to local tool outputs", response_id ) + return - return saved_run_items_count + deferred_response_id = None + get_deferred = getattr(session, "_get_deferred_compaction_response_id", None) + if callable(get_deferred): + deferred_response_id = get_deferred() + force_compaction = deferred_response_id is not None + if force_compaction: + logger.debug( + "compact: forcing for response %s after deferred %s", + response_id, + deferred_response_id, + ) + compaction_args: OpenAIResponsesCompactionArgs = { + "response_id": response_id, + "force": force_compaction, + } + if store is not None: + compaction_args["store"] = store + await _call_session_method( + session.run_compaction, + compaction_args, + wrapper=wrapper, + ) async def save_resumed_turn_items( @@ -1097,6 +1180,29 @@ def _collect_retry_owned_tail_serializations( return [] +def _normalize_history_mutations_for_session_persistence( + session: Session, + mutations: list[SessionHistoryMutation], +) -> list[SessionHistoryMutation]: + """Normalize persisted-history mutations to the same session-safe item shape used on writes.""" + normalized: list[SessionHistoryMutation] = [] + for mutation in mutations: + if mutation["type"] != "replace_function_call": + continue + replacement = ensure_input_item_format(mutation["replacement"]) + if isinstance(session, OpenAIConversationsSession): + replacement = _sanitize_openai_conversation_item(replacement) + normalized.append( + { + "type": "replace_function_call", + "call_id": mutation["call_id"], + "expected": ensure_input_item_format(mutation["expected"]), + "replacement": replacement, + } + ) + return normalized + + def _session_item_key(item: Any, *, ignore_openai_conversation_item_ids: bool = False) -> str: """Return a stable representation of a session item for comparison.""" try: diff --git a/src/agents/run_internal/turn_resolution.py b/src/agents/run_internal/turn_resolution.py index 8c756b0bc9..5a649aa26f 100644 --- a/src/agents/run_internal/turn_resolution.py +++ b/src/agents/run_internal/turn_resolution.py @@ -2126,7 +2126,7 @@ def _rebind_function_run( continue _add_pending_interruption(original_approval) continue - nested_state, nested_item = nested_approval + nested_state, nested_item, _outer_tool_call, _pending_result = nested_approval nested_context = nested_state._context nested_call_id = get_tool_approval_item_call_id(nested_item) nested_status = ( diff --git a/src/agents/run_state.py b/src/agents/run_state.py index adc44c3228..6243c21273 100644 --- a/src/agents/run_state.py +++ b/src/agents/run_state.py @@ -6,6 +6,7 @@ import copy import dataclasses import json +import math import threading from collections import deque from collections.abc import Callable, Collection, Iterator, Mapping, Sequence @@ -100,6 +101,7 @@ log_model_and_tool_data_warning, logger, ) +from .memory import SessionHistoryMutation from .run_context import RunContextWrapper from .run_internal.items import ( NestedHistoryOwnedItemRef, @@ -150,6 +152,8 @@ ToolRunFunction, ) +from .run_internal.items import ensure_input_item_format + TContext = TypeVar("TContext", default=Any) TAgent = TypeVar("TAgent", bound="Agent[Any]", default="Agent[Any]") TAction = TypeVar("TAction") @@ -177,7 +181,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.15" +CURRENT_SCHEMA_VERSION = "1.16" _PROGRAMMATIC_TOOL_CALLING_MIN_SCHEMA_VERSION = "1.13" _HOSTED_MCP_APPROVALS_MIN_SCHEMA_VERSION = "1.14" # Keep this mapping in chronological order. Every schema bump must add a one-line summary here. @@ -207,6 +211,7 @@ def _default_run_state_validation_error( "Persists canonical tool invocation identity plus sanitized mount authority and trusted " "rebind metadata, durable pending input, and resumable next-model-call state." ), + "1.16": "Persists approval argument overrides and pending session-history rewrites.", } SUPPORTED_SCHEMA_VERSIONS = frozenset(SCHEMA_VERSION_SUMMARIES) @@ -346,6 +351,14 @@ class RunState(Generic[TContext, TAgent]): _tool_use_tracker_snapshot: dict[str, list[str]] = field(default_factory=dict) """Serialized snapshot of the AgentToolUseTracker (agent name -> tools used).""" + _session_history_mutations: list[SessionHistoryMutation] = field(default_factory=list) + """Pending session history rewrites that must be applied after persistence.""" + + _approval_argument_override_modes: dict[str, Literal["durable", "execution_only"]] = field( + default_factory=dict + ) + """Persistence mode for each pending approved argument override.""" + _trace_state: TraceState | None = field(default=None, repr=False) """Serialized trace metadata for resuming tracing context.""" @@ -395,6 +408,8 @@ def __init__( self._generated_items_last_processed_marker = None self._current_turn_persisted_item_count = 0 self._tool_use_tracker_snapshot = {} + self._session_history_mutations = [] + self._approval_argument_override_modes = {} self._trace_state = None self._sandbox = None self._schema_version = CURRENT_SCHEMA_VERSION @@ -491,8 +506,8 @@ def _approval_items_match( def _find_nested_approval_state( self, approval_item: ToolApprovalItem, - ) -> tuple[RunState[Any, Agent[Any]], ToolApprovalItem] | None: - """Find the nested agent-tool state that owns an approval interruption.""" + ) -> tuple[RunState[Any, Agent[Any]], ToolApprovalItem, Any, Any] | None: + """Find the nested state, approval, owning tool call, and cached result.""" if self._last_processed_response is None: return None @@ -579,8 +594,8 @@ def _find_nested_approval_state( current_state_owns_approval and approval_identity in current_response_identities ) - exact_match: tuple[RunState[Any, Agent[Any]], ToolApprovalItem] | None = None - canonical_matches: list[tuple[RunState[Any, Agent[Any]], ToolApprovalItem]] = [] + exact_match: tuple[RunState[Any, Agent[Any]], ToolApprovalItem, Any, Any] | None = None + canonical_matches: list[tuple[RunState[Any, Agent[Any]], ToolApprovalItem, Any, Any]] = [] for function_run in self._last_processed_response.functions: pending_result = peek_agent_tool_run_result( function_run.tool_call, @@ -597,10 +612,12 @@ def _find_nested_approval_state( if not isinstance(candidate, ToolApprovalItem): continue if candidate is approval_item: - exact_match = (nested_state, candidate) + exact_match = (nested_state, candidate, function_run.tool_call, pending_result) break if self._approval_items_match(candidate, approval_item): - canonical_matches.append((nested_state, candidate)) + canonical_matches.append( + (nested_state, candidate, function_run.tool_call, pending_result) + ) if exact_match is not None: break @@ -620,17 +637,287 @@ def _find_nested_approval_state( ) return None - def approve(self, approval_item: ToolApprovalItem, always_approve: bool = False) -> None: + def approve( + self, + approval_item: ToolApprovalItem, + always_approve: bool = False, + *, + override_arguments: dict[str, Any] | None = None, + save_override_arguments: bool | None = None, + ) -> None: """Approve a tool call and rerun with this state to continue.""" if self._context is None: raise UserError("Cannot approve tool: RunState has no context") nested_approval = self._find_nested_approval_state(approval_item) if nested_approval is not None: - nested_state, nested_item = nested_approval - nested_state.approve(nested_item, always_approve=always_approve) + nested_state, nested_item, outer_tool_call, pending_result = nested_approval + if override_arguments is None: + nested_state.approve( + nested_item, + always_approve=always_approve, + save_override_arguments=save_override_arguments, + ) + return + + from .agent_tool_state import ( + record_agent_tool_resume_state, + record_agent_tool_run_result, + ) + + approval_items = getattr(pending_result, "interruptions", None) + record_agent_tool_resume_state( + outer_tool_call, + nested_state, + scope_id=self._agent_tool_state_scope_id, + approval_items=approval_items if isinstance(approval_items, list) else None, + ) + try: + nested_state.approve( + nested_item, + always_approve=always_approve, + override_arguments=override_arguments, + save_override_arguments=save_override_arguments, + ) + except Exception: + record_agent_tool_run_result( + outer_tool_call, + pending_result, + scope_id=self._agent_tool_state_scope_id, + ) + raise return + + if save_override_arguments is not None and override_arguments is None: + raise UserError( + "save_override_arguments can only be used together with override_arguments." + ) + + if override_arguments is not None: + self._apply_approval_argument_override( + approval_item=approval_item, + override_arguments=override_arguments, + always_approve=always_approve, + save_override_arguments=save_override_arguments, + ) self._context.approve_tool(approval_item, always_approve=always_approve) + def _get_session_history_mutations(self) -> list[SessionHistoryMutation]: + """Return a defensive copy of pending persisted-history mutations.""" + return copy.deepcopy(self._session_history_mutations) + + def _has_pending_execution_only_approval_overrides(self) -> bool: + """Return whether execution-only argument overrides still need a server-managed resume.""" + return "execution_only" in self._approval_argument_override_modes.values() + + def _clear_executed_approval_argument_overrides( + self, + *, + _visited_state_ids: set[int] | None = None, + ) -> None: + """Clear executed override bookkeeping in this state and owned checkpoints.""" + visited_state_ids = _visited_state_ids if _visited_state_ids is not None else set() + if id(self) in visited_state_ids: + return + visited_state_ids.add(id(self)) + + if self._context is None: + return + for interruption in self.get_interruptions(): + call_id = _get_raw_item_call_id(interruption.raw_item) + if call_id is None: + continue + mode = self._approval_argument_override_modes.get(call_id) + if mode is None: + continue + try: + invocation_status = self._context._approved_tool_invocation_status( + interruption.raw_item, + tool_lookup_key=interruption.tool_lookup_key, + tool_name=interruption.tool_name, + ) + except ModelBehaviorError: + continue + if invocation_status is not None and invocation_status[2]: + self._approval_argument_override_modes.pop(call_id, None) + if mode == "durable": + self._session_history_mutations = [ + mutation + for mutation in self._session_history_mutations + if mutation["call_id"] != call_id + ] + + if self._last_processed_response is None: + return + + from .agent_tool_state import get_agent_tool_resume_state, peek_agent_tool_run_result + + for function_run in getattr(self._last_processed_response, "functions", ()): + pending_result = peek_agent_tool_run_result( + function_run.tool_call, + scope_id=self._agent_tool_state_scope_id, + ) + nested_state = get_agent_tool_resume_state(pending_result) + if isinstance(nested_state, RunState): + nested_state._clear_executed_approval_argument_overrides( + _visited_state_ids=visited_state_ids + ) + + def _clear_session_history_mutations(self) -> None: + """Clear pending persisted-history mutations after they are applied.""" + self._session_history_mutations = [] + self._approval_argument_override_modes = { + call_id: mode + for call_id, mode in self._approval_argument_override_modes.items() + if mode != "durable" + } + + def _apply_approval_argument_override( + self, + *, + approval_item: ToolApprovalItem, + override_arguments: dict[str, Any], + always_approve: bool, + save_override_arguments: bool | None, + ) -> None: + context = self._context + assert context is not None + if always_approve: + raise UserError("override_arguments cannot be used together with always_approve.") + + if not _is_function_call_raw_item(approval_item.raw_item): + raise UserError("override_arguments is only supported for function_call approvals.") + + if type(override_arguments) is not dict: + raise UserError("override_arguments must be a plain JSON object.") + + serialized_arguments: str | None = None + try: + if not _is_plain_json_value(override_arguments): + raise TypeError("Approval arguments must contain only plain JSON values") + serialized_arguments = json.dumps(override_arguments, allow_nan=False) + except Exception: + pass + + if not isinstance(serialized_arguments, str): + raise UserError( + "override_arguments must contain only JSON-serializable values." + ) from None + + should_save_override_arguments = save_override_arguments is not False + has_server_managed_conversation = bool( + self._conversation_id or self._previous_response_id or self._auto_previous_response_id + ) + if should_save_override_arguments and has_server_managed_conversation: + raise UserError( + "save_override_arguments requires local canonical history. " + "Server-managed conversations cannot persist corrected function_call arguments. " + "Pass save_override_arguments=False to apply the override only to the current " + "execution." + ) + + call_id = _get_raw_item_call_id(approval_item.raw_item) + if not isinstance(call_id, str): + raise UserError("override_arguments requires a function_call with a call_id.") + previous_override_mode = self._approval_argument_override_modes.get(call_id) + if previous_override_mode is not None: + raise UserError( + "Cannot replace an existing argument override for the same function_call. " + "Rebuild the RunState from before the first argument override." + ) + + expected_tool_call = copy.deepcopy( + ensure_input_item_format(cast(TResponseInputItem, approval_item.raw_item)) + ) + updated_tool_call = _create_function_call_override( + approval_item.raw_item, + serialized_arguments, + ) + previous_identity = tool_invocation_identity( + approval_item.raw_item, + tool_lookup_key=approval_item.tool_lookup_key, + tool_name=approval_item.tool_name, + ) + if previous_identity is None: + raise UserError("override_arguments requires a canonical function_call identity.") + rebound = context._rebind_tool_invocation( + updated_tool_call, + previous_identity=previous_identity, + tool_lookup_key=approval_item.tool_lookup_key, + tool_name=approval_item.tool_name, + ) + if rebound is None: + raise UserError("override_arguments could not rebind the function_call identity.") + approval_item.raw_item = updated_tool_call + self._replace_function_call_in_interruptions(call_id, updated_tool_call) + if self._last_processed_response is not None: + for interruption in self._last_processed_response.interruptions: + if not _is_function_call_raw_item(interruption.raw_item): + continue + if _get_raw_item_call_id(interruption.raw_item) != call_id: + continue + interruption.raw_item = updated_tool_call + + self._approval_argument_override_modes[call_id] = ( + "durable" if should_save_override_arguments else "execution_only" + ) + + if self._last_processed_response is not None: + for function_run in self._last_processed_response.functions: + if _get_raw_item_call_id(function_run.tool_call) != call_id: + continue + function_run.tool_call = updated_tool_call + + if should_save_override_arguments: + self._replace_function_call_in_run_items( + self._last_processed_response.new_items, + call_id, + updated_tool_call, + ) + + if should_save_override_arguments: + self._replace_function_call_in_run_items( + self._generated_items, call_id, updated_tool_call + ) + self._replace_function_call_in_run_items( + self._session_items, call_id, updated_tool_call + ) + self._record_session_history_mutation( + { + "type": "replace_function_call", + "call_id": call_id, + "expected": expected_tool_call, + "replacement": ensure_input_item_format(updated_tool_call), + } + ) + self._mark_generated_items_merged_with_last_processed() + + def _replace_function_call_in_interruptions(self, call_id: str, tool_call: Any) -> None: + """Replace a function call inside pending approval interruptions.""" + for interruption in self.get_interruptions(): + if not _is_function_call_raw_item(interruption.raw_item): + continue + if _get_raw_item_call_id(interruption.raw_item) != call_id: + continue + interruption.raw_item = tool_call + + def _replace_function_call_in_run_items( + self, + items: list[RunItem], + call_id: str, + tool_call: Any, + ) -> None: + """Replace matching function-call raw items inside serialized run item history.""" + for item in items: + if not _is_function_call_raw_item(getattr(item, "raw_item", None)): + continue + if _get_raw_item_call_id(item.raw_item) != call_id: + continue + item.raw_item = tool_call + + def _record_session_history_mutation(self, mutation: SessionHistoryMutation) -> None: + """Record one pending persisted-history mutation for an approved function call.""" + self._session_history_mutations.append(copy.deepcopy(mutation)) + def reject( self, approval_item: ToolApprovalItem, @@ -648,7 +935,7 @@ def reject( raise UserError("Cannot reject tool: RunState has no context") nested_approval = self._find_nested_approval_state(approval_item) if nested_approval is not None: - nested_state, nested_item = nested_approval + nested_state, nested_item, _outer_tool_call, _pending_result = nested_approval nested_state.reject( nested_item, always_reject=always_reject, @@ -1167,6 +1454,14 @@ def to_json( for item_ref in self._nested_history_owned_session_item_refs ], "generated_session_item_indexes": self._generated_session_item_indexes(generated_items), + "approval_argument_override_modes": [ + {"call_id": call_id, "mode": mode} + for call_id, mode in self._approval_argument_override_modes.items() + ], + "session_history_mutations": [ + _serialize_session_history_mutation(mutation) + for mutation in self._session_history_mutations + ], } result["generated_items"] = [ @@ -1776,6 +2071,474 @@ def _serialize_agent_reference( return entry +def _serialize_session_history_mutation(mutation: SessionHistoryMutation) -> dict[str, Any]: + """Serialize a session history mutation into a JSON-compatible dictionary.""" + return { + "type": mutation["type"], + "call_id": mutation["call_id"], + "expected": _serialize_raw_item_value(mutation["expected"]), + "replacement": _serialize_raw_item_value(mutation["replacement"]), + } + + +def _is_plain_json_value(value: Any, active_containers: set[int] | None = None) -> bool: + """Return whether a value is an acyclic tree of exact built-in JSON types.""" + value_type = type(value) + if value is None or value_type in {str, bool, int}: + return True + if value_type is float: + return math.isfinite(value) + if active_containers is None: + active_containers = set() + if value_type is list: + value_id = id(value) + if value_id in active_containers: + return False + active_containers.add(value_id) + try: + return all(_is_plain_json_value(item, active_containers) for item in value) + finally: + active_containers.remove(value_id) + if value_type is dict: + value_id = id(value) + if value_id in active_containers: + return False + active_containers.add(value_id) + try: + return all( + type(key) is str and _is_plain_json_value(item, active_containers) + for key, item in dict.items(value) + ) + finally: + active_containers.remove(value_id) + return False + + +def _deserialize_session_history_mutations( + serialized_mutations: Any, + *, + validation_error_factory: RunStateValidationErrorFactory = _default_run_state_validation_error, +) -> list[SessionHistoryMutation]: + """Deserialize persisted session history mutations from JSON data.""" + if not isinstance(serialized_mutations, Sequence) or isinstance( + serialized_mutations, str | bytes + ): + raise validation_error_factory( + "Run state session_history_mutations must be a list of valid function-call " + "replacements.", + UserError, + ) + + mutations: list[SessionHistoryMutation] = [] + for mutation in serialized_mutations: + call_id = mutation.get("call_id") if isinstance(mutation, Mapping) else None + expected = mutation.get("expected") if isinstance(mutation, Mapping) else None + replacement = mutation.get("replacement") if isinstance(mutation, Mapping) else None + normalized_expected = dict(expected) if isinstance(expected, Mapping) else {} + normalized_replacement = dict(replacement) if isinstance(replacement, Mapping) else {} + is_valid = ( + isinstance(mutation, Mapping) + and mutation.get("type") == "replace_function_call" + and isinstance(call_id, str) + and normalized_expected.get("type") == "function_call" + and normalized_expected.get("call_id") == call_id + and isinstance(normalized_expected.get("name"), str) + and isinstance(normalized_expected.get("arguments"), str) + and normalized_replacement.get("type") == "function_call" + and normalized_replacement.get("call_id") == call_id + and isinstance(normalized_replacement.get("name"), str) + and isinstance(normalized_replacement.get("arguments"), str) + ) + if not is_valid: + raise validation_error_factory( + "Run state session_history_mutations contains an invalid function-call " + "replacement.", + UserError, + ) + mutations.append( + { + "type": "replace_function_call", + "call_id": cast(str, call_id), + "expected": cast(TResponseInputItem, normalized_expected), + "replacement": cast(TResponseInputItem, normalized_replacement), + } + ) + return mutations + + +def _validate_session_history_mutations_match_interruptions( + state: RunState[Any, Agent[Any]], + *, + validation_error_factory: RunStateValidationErrorFactory, +) -> None: + """Validate that durable rewrites exactly match their pending execution calls.""" + interruptions = state.get_interruptions() + canonical_calls = [ + item.raw_item + for item in state._generated_items + if isinstance(item, ToolCallItem) and _is_function_call_raw_item(item.raw_item) + ] + session_calls = [ + item.raw_item + for item in state._session_items + if isinstance(item, ToolCallItem) and _is_function_call_raw_item(item.raw_item) + ] + processed_calls = ( + [run.tool_call for run in state._last_processed_response.functions] + if state._last_processed_response is not None + else [] + ) + processed_new_item_calls = ( + [ + item.raw_item + for item in state._last_processed_response.new_items + if isinstance(item, ToolCallItem) and _is_function_call_raw_item(item.raw_item) + ] + if state._last_processed_response is not None + else [] + ) + model_response_calls = [ + raw_item + for response in state._model_responses + for raw_item in response.output + if _is_function_call_raw_item(raw_item) + ] + for mutation in state._session_history_mutations: + call_id = mutation["call_id"] + matching_interruptions = [ + interruption + for interruption in interruptions + if _is_function_call_raw_item(interruption.raw_item) + and _get_raw_item_call_id(interruption.raw_item) == call_id + ] + matching_canonical_calls = [ + raw_item for raw_item in canonical_calls if _get_raw_item_call_id(raw_item) == call_id + ] + matching_session_calls = [ + raw_item for raw_item in session_calls if _get_raw_item_call_id(raw_item) == call_id + ] + matching_processed_calls = [ + raw_item for raw_item in processed_calls if _get_raw_item_call_id(raw_item) == call_id + ] + matching_processed_new_item_calls = [ + raw_item + for raw_item in processed_new_item_calls + if _get_raw_item_call_id(raw_item) == call_id + ] + matching_model_response_calls = [ + raw_item + for raw_item in model_response_calls + if _get_raw_item_call_id(raw_item) == call_id + ] + expected = ensure_input_item_format(mutation["expected"]) + replacement = mutation["replacement"] + normalized_replacement = ensure_input_item_format(replacement) + is_match = ( + isinstance(expected, dict) + and isinstance(normalized_replacement, dict) + and len(matching_interruptions) == 1 + and len(matching_canonical_calls) == 1 + and len(matching_session_calls) == 1 + and len(matching_processed_calls) == 1 + and len(matching_processed_new_item_calls) == 1 + and len(matching_model_response_calls) == 1 + and expected + == ensure_input_item_format(cast(TResponseInputItem, matching_model_response_calls[0])) + and normalized_replacement + == ensure_input_item_format( + cast(TResponseInputItem, matching_interruptions[0].raw_item) + ) + and normalized_replacement + == ensure_input_item_format(cast(TResponseInputItem, matching_canonical_calls[0])) + and normalized_replacement + == ensure_input_item_format(cast(TResponseInputItem, matching_session_calls[0])) + and normalized_replacement + == ensure_input_item_format(cast(TResponseInputItem, matching_processed_calls[0])) + and normalized_replacement + == ensure_input_item_format( + cast(TResponseInputItem, matching_processed_new_item_calls[0]) + ) + ) + if not is_match: + raise validation_error_factory( + "Run state session_history_mutations does not match the pending function call.", + UserError, + ) + + +def _deserialize_approval_argument_override_modes( + serialized_modes: Any, + *, + validation_error_factory: RunStateValidationErrorFactory, +) -> dict[str, Literal["durable", "execution_only"]]: + """Deserialize per-call override modes with strict current-schema validation.""" + if type(serialized_modes) is not list: + raise validation_error_factory( + "Run state approval_argument_override_modes must be a list of unique valid " + "per-call modes.", + UserError, + ) + result: dict[str, Literal["durable", "execution_only"]] = {} + for item in serialized_modes: + if type(item) is not dict or set(item) != {"call_id", "mode"}: + raise validation_error_factory( + "Run state approval_argument_override_modes must be a list of unique valid " + "per-call modes.", + UserError, + ) + call_id = item.get("call_id") + mode = item.get("mode") + if ( + type(call_id) is not str + or not call_id + or mode not in {"durable", "execution_only"} + or call_id in result + ): + raise validation_error_factory( + "Run state approval_argument_override_modes must be a list of unique valid " + "per-call modes.", + UserError, + ) + result[call_id] = cast(Literal["durable", "execution_only"], mode) + return result + + +def _validate_approval_argument_overrides( + state: RunState[Any, Agent[Any]], + *, + validation_error_factory: RunStateValidationErrorFactory, +) -> None: + """Validate override modes against pending execution, replay, and durable mutations.""" + context = state._context + assert context is not None + modes = state._approval_argument_override_modes + mutation_call_ids = [mutation["call_id"] for mutation in state._session_history_mutations] + durable_call_ids = {call_id for call_id, mode in modes.items() if mode == "durable"} + if len(mutation_call_ids) != len(set(mutation_call_ids)) or durable_call_ids != set( + mutation_call_ids + ): + raise validation_error_factory( + "Run state approval_argument_override_modes is inconsistent with pending function " + "calls and session history mutations.", + UserError, + ) + + interruptions = [ + interruption + for interruption in state.get_interruptions() + if _is_function_call_raw_item(interruption.raw_item) + ] + canonical_calls = [ + item.raw_item + for item in state._generated_items + if isinstance(item, ToolCallItem) and _is_function_call_raw_item(item.raw_item) + ] + session_calls = [ + item.raw_item + for item in state._session_items + if isinstance(item, ToolCallItem) and _is_function_call_raw_item(item.raw_item) + ] + model_response_calls = [ + raw_item + for response in state._model_responses + for raw_item in response.output + if _is_function_call_raw_item(raw_item) + ] + processed_calls = ( + [run.tool_call for run in state._last_processed_response.functions] + if state._last_processed_response is not None + else [] + ) + processed_new_item_calls = ( + [ + item.raw_item + for item in state._last_processed_response.new_items + if isinstance(item, ToolCallItem) and _is_function_call_raw_item(item.raw_item) + ] + if state._last_processed_response is not None + else [] + ) + + approved_pending_call_ids: set[str] = set() + for interruption in interruptions: + call_id = _get_raw_item_call_id(interruption.raw_item) + if call_id is None: + continue + try: + approval_status = context._approved_tool_invocation_status( + interruption.raw_item, + tool_lookup_key=interruption.tool_lookup_key, + tool_name=interruption.tool_name, + ) + except ModelBehaviorError: + continue + if approval_status is None or approval_status[1] or approval_status[2]: + continue + if call_id in approved_pending_call_ids: + raise validation_error_factory( + "Run state approval_argument_override_modes is inconsistent with pending " + "function calls and session history mutations.", + UserError, + ) + approved_pending_call_ids.add(call_id) + matching_canonical_calls = [ + raw_item for raw_item in canonical_calls if _get_raw_item_call_id(raw_item) == call_id + ] + matching_model_response_calls = [ + raw_item + for raw_item in model_response_calls + if _get_raw_item_call_id(raw_item) == call_id + ] + matching_session_calls = [ + raw_item for raw_item in session_calls if _get_raw_item_call_id(raw_item) == call_id + ] + matching_processed_calls = [ + raw_item for raw_item in processed_calls if _get_raw_item_call_id(raw_item) == call_id + ] + matching_processed_new_item_calls = [ + raw_item + for raw_item in processed_new_item_calls + if _get_raw_item_call_id(raw_item) == call_id + ] + execution_payload = ensure_input_item_format( + cast(TResponseInputItem, interruption.raw_item) + ) + canonical_payload = ( + ensure_input_item_format(cast(TResponseInputItem, matching_canonical_calls[0])) + if len(matching_canonical_calls) == 1 + else None + ) + session_payloads = [ + ensure_input_item_format(cast(TResponseInputItem, raw_item)) + for raw_item in matching_session_calls + ] + model_response_payloads = [ + ensure_input_item_format(cast(TResponseInputItem, raw_item)) + for raw_item in matching_model_response_calls + ] + processed_payloads = [ + ensure_input_item_format(cast(TResponseInputItem, raw_item)) + for raw_item in matching_processed_calls + ] + processed_new_item_payloads = [ + ensure_input_item_format(cast(TResponseInputItem, raw_item)) + for raw_item in matching_processed_new_item_calls + ] + mode = modes.get(call_id) + model_response_identity_payload = ( + _function_call_payload_without_arguments(model_response_payloads[0]) + if len(model_response_payloads) == 1 + else None + ) + execution_identity_payload = _function_call_payload_without_arguments(execution_payload) + audit_identity_is_consistent = mode not in {"durable", "execution_only"} or ( + model_response_identity_payload is not None + and execution_identity_payload == model_response_identity_payload + ) + execution_only_is_consistent = mode != "execution_only" or ( + canonical_payload is not None + and len(session_payloads) == 1 + and len(model_response_payloads) == 1 + and len(processed_payloads) == 1 + and len(processed_new_item_payloads) == 1 + and canonical_payload == session_payloads[0] + and canonical_payload == model_response_payloads[0] + and canonical_payload == processed_new_item_payloads[0] + and execution_payload == processed_payloads[0] + ) + unmarked_payload_diverges = mode is None and ( + (canonical_payload is not None and execution_payload != canonical_payload) + or any(execution_payload != payload for payload in model_response_payloads) + ) + if ( + not audit_identity_is_consistent + or not execution_only_is_consistent + or unmarked_payload_diverges + ): + raise validation_error_factory( + "Run state approval_argument_override_modes is inconsistent with pending " + "function calls and session history mutations.", + UserError, + ) + + if not set(modes).issubset(approved_pending_call_ids): + raise validation_error_factory( + "Run state approval_argument_override_modes is inconsistent with pending function " + "calls and session history mutations.", + UserError, + ) + + +def _function_call_payload_without_arguments( + payload: TResponseInputItem, +) -> dict[str, Any] | None: + """Return the stable function-call payload used to compare override audit identity.""" + if not isinstance(payload, dict): + return None + identity_payload = dict(payload) + identity_payload.pop("arguments", None) + return identity_payload + + +def _validate_last_model_response_matches_model_responses( + state_json: Mapping[str, Any], + *, + validation_error_factory: RunStateValidationErrorFactory, +) -> None: + """Validate the current writer's duplicate last-model-response audit record.""" + model_responses = state_json.get("model_responses") + expected_last_response = ( + model_responses[-1] if isinstance(model_responses, list) and model_responses else None + ) + if state_json.get("last_model_response") != expected_last_response: + raise validation_error_factory( + "Run state approval_argument_override_modes is inconsistent with pending function " + "calls and session history mutations.", + UserError, + ) + + +def _is_function_call_raw_item(raw_item: Any) -> bool: + """Return whether the raw item represents a function call.""" + if isinstance(raw_item, dict): + return raw_item.get("type") == "function_call" + return getattr(raw_item, "type", None) == "function_call" + + +def _get_raw_item_call_id(raw_item: Any) -> str | None: + """Return the call_id for a raw tool item when available.""" + if isinstance(raw_item, dict): + call_id = raw_item.get("call_id") or raw_item.get("callId") or raw_item.get("id") + return call_id if isinstance(call_id, str) else None + for attr in ("call_id", "callId", "id"): + value = getattr(raw_item, attr, None) + if isinstance(value, str): + return value + return None + + +def _create_function_call_override(raw_item: Any, serialized_arguments: str) -> Any: + """Return a copy of a function call raw item with corrected arguments.""" + if isinstance(raw_item, ResponseFunctionToolCall): + return raw_item.model_copy(update={"arguments": serialized_arguments}) + if isinstance(raw_item, dict): + updated = dict(raw_item) + updated["arguments"] = serialized_arguments + return updated + if hasattr(raw_item, "model_dump"): + try: + payload = raw_item.model_dump(exclude_unset=True) + except TypeError: + payload = raw_item.model_dump() + payload = dict(payload) + payload["arguments"] = serialized_arguments + try: + return ResponseFunctionToolCall(**payload) + except Exception: + return payload + raise UserError("override_arguments is only supported for function_call approvals.") + + def _ensure_json_compatible(value: Any) -> Any: try: return json.loads(json.dumps(value, default=str)) @@ -3631,6 +4394,31 @@ async def _build_run_state_from_json( state._generated_prompt_cache_key = ( serialized_prompt_cache_key if isinstance(serialized_prompt_cache_key, str) else None ) + if (schema_major, schema_minor) >= (1, 16): + state._approval_argument_override_modes = _deserialize_approval_argument_override_modes( + state_json.get("approval_argument_override_modes", []), + validation_error_factory=validation_error_factory, + ) + if state._approval_argument_override_modes: + _validate_last_model_response_matches_model_responses( + state_json, + validation_error_factory=validation_error_factory, + ) + state._session_history_mutations = _deserialize_session_history_mutations( + state_json.get("session_history_mutations", []), + validation_error_factory=validation_error_factory, + ) + _validate_session_history_mutations_match_interruptions( + state, + validation_error_factory=validation_error_factory, + ) + _validate_approval_argument_overrides( + state, + validation_error_factory=validation_error_factory, + ) + else: + state._approval_argument_override_modes = {} + state._session_history_mutations = [] state.set_tool_use_tracker_snapshot(state_json.get("tool_use_tracker", {})) trace_data = state_json.get("trace") if isinstance(trace_data, Mapping): @@ -4779,6 +5567,17 @@ def _clone_original_input(original_input: str | list[Any]) -> str | list[Any]: ), "Run state agent not found in agent map", "Run state pending_input must be a list", + ("Run state session_history_mutations must be a list of valid function-call replacements."), + ("Run state session_history_mutations contains an invalid function-call replacement."), + "Run state session_history_mutations does not match the pending function call.", + ( + "Run state approval_argument_override_modes must be a list of unique valid per-call " + "modes." + ), + ( + "Run state approval_argument_override_modes is inconsistent with pending function " + "calls and session history mutations." + ), "Run state references an agent identity that is not present in the restored graph", ( "RunState context was serialized from a custom type; provide context_deserializer " diff --git a/tests/fixtures/run_state/README.md b/tests/fixtures/run_state/README.md index f111591b0d..fff9b9a3ed 100644 --- a/tests/fixtures/run_state/README.md +++ b/tests/fixtures/run_state/README.md @@ -1,6 +1,6 @@ # RunState compatibility corpus -The `minimal/` fixtures cover every schema version accepted by the current reader. The `features/` fixtures cover the schema-bearing behavior introduced in versions 1.2 through 1.15. The `resume/` fixture records an actual pending function-tool approval emitted by the v0.19.4 writer, and the `security/` fixture records that writer's credential-bearing sandbox state. `sources.json` records the source commit and provenance for every fixture. +The `minimal/` fixtures cover every schema version accepted by the current reader. The `features/` fixtures cover the schema-bearing behavior introduced in versions 1.2 through 1.16. The `resume/` fixture records an actual pending function-tool approval emitted by the v0.19.4 writer, and the `security/` fixture records that writer's credential-bearing sandbox state. `sources.json` records the source commit and provenance for every fixture. Regenerate the feature corpus from the recorded historical source trees with: @@ -10,6 +10,8 @@ UV_DEFAULT_INDEX=https://pypi.org/simple uv run python tests/fixtures/run_state/ The generator extracts each recorded commit with `git archive` and runs that commit's writer in a fresh locked environment. It does not import the current checkout. -Versions 1.7 and 1.8 are explicit exceptions. Release-boundary schema renumbering assigned duplicate-agent/sandbox state to 1.7 and prompt-cache state to 1.8 without any writer commit that emitted those final version numbers. Their fixtures are therefore marked `canonical_compatibility`: the recorded 1.9 writer produces the payload, and the generator changes only the schema label so the corresponding reader branch remains covered. They must not be represented as historical-writer output. +Versions 1.7 and 1.8 are explicit exceptions. Release-boundary schema renumbering assigned duplicate-agent/sandbox state to 1.7 and prompt-cache state to 1.8 without any writer commit that emitted those final version numbers. Their fixtures are therefore marked `canonical_compatibility`: the recorded writer produces the payload, and the generator changes only the schema label so the corresponding reader branch remains covered. They must not be represented as historical-writer output. + +The current unreleased schema uses `current_writer` provenance because a fixture cannot name the commit that contains itself. The generator executes the current checkout for that fixture, and `lineage_commit` records the pre-takeover implementation commit from which the behavior originated. Once a later commit becomes the historical writer, replace the current-writer provenance with that stable commit rather than relabeling the fixture as historical output prematurely. Ordinary tests never run the generator. They read the frozen payloads, compare every durable field emitted by the historical writer across the upgrade, rewrite to the current schema, and verify that the rewritten form is idempotent. They also approve and reject the historical pending interruption through actual `Runner` resumes. The schema version itself is the only normalization for the ordinary corpus; fields added by newer writers may be absent from an older payload, but every field present in that payload must survive. The security fixture has one explicit migration normalization: persisted mount credentials and opaque driver options are removed and the trusted-rebind marker is added. All non-authority topology remains part of the comparison. diff --git a/tests/fixtures/run_state/features/v1_16_approval_argument_overrides.json b/tests/fixtures/run_state/features/v1_16_approval_argument_overrides.json new file mode 100644 index 0000000000..44652d5e6c --- /dev/null +++ b/tests/fixtures/run_state/features/v1_16_approval_argument_overrides.json @@ -0,0 +1,280 @@ +{ + "$schemaVersion": "1.16", + "approval_argument_override_modes": [ + { + "call_id": "approval-override-1", + "mode": "execution_only" + } + ], + "auto_previous_response_id": false, + "context": { + "approvals": { + "send_email": { + "approved": [ + "approval-override-1" + ], + "rejected": [] + } + }, + "context": {}, + "context_meta": { + "omitted": false, + "original_type": "mapping", + "requires_deserializer": false, + "serialized_via": "mapping" + }, + "tool_invocations": { + "approval-override-1": { + "approval_scope": "03c0b69e3d8062f4c5ad4a51305b5ac6f32321871b2433b01ea6025ed2565956", + "completed": false, + "executed": false, + "fingerprint": "05d7f7f0356c79296d60ce0ccc75dbb2d5f990eb09664cd7a7a459182cae969a", + "type": "function_call" + } + }, + "usage": { + "input_tokens": 0, + "input_tokens_details": [ + { + "cache_write_tokens": 0, + "cached_tokens": 0 + } + ], + "output_tokens": 0, + "output_tokens_details": [ + { + "reasoning_tokens": 0 + } + ], + "request_usage_entries": [], + "requests": 0, + "total_tokens": 0 + } + }, + "conversation_id": null, + "current_agent": { + "name": "compat-agent" + }, + "current_step": { + "data": { + "interruptions": [ + { + "agent": { + "name": "compat-agent" + }, + "raw_item": { + "arguments": "{\"recipient\": \"bob@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + }, + "tool_lookup_key": { + "kind": "bare", + "name": "send_email" + }, + "tool_name": "send_email", + "type": "tool_approval_item" + } + ], + "llm_end_hooks_started": true, + "response_accepted": false + }, + "type": "next_step_interruption" + }, + "current_turn": 0, + "current_turn_persisted_item_count": 0, + "generated_items": [ + { + "agent": { + "name": "compat-agent" + }, + "raw_item": { + "arguments": "{\"recipient\": \"alice@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + }, + "tool_name": "send_email", + "type": "tool_call_item" + } + ], + "generated_prompt_cache_key": null, + "generated_session_item_indexes": [ + null + ], + "input_guardrail_results": [], + "last_model_response": { + "output": [ + { + "arguments": "{\"recipient\": \"alice@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + } + ], + "request_id": null, + "response_id": "response-override", + "usage": { + "input_tokens": 0, + "input_tokens_details": [ + { + "cache_write_tokens": 0, + "cached_tokens": 0 + } + ], + "output_tokens": 0, + "output_tokens_details": [ + { + "reasoning_tokens": 0 + } + ], + "request_usage_entries": [], + "requests": 0, + "total_tokens": 0 + } + }, + "last_processed_response": { + "apply_patch_actions": [], + "computer_actions": [], + "custom_tool_actions": [], + "functions": [ + { + "tool": { + "description": "", + "lookupKey": { + "kind": "bare", + "name": "send_email" + }, + "name": "send_email", + "paramsJsonSchema": { + "additionalProperties": false, + "properties": { + "recipient": { + "title": "Recipient", + "type": "string" + } + }, + "required": [ + "recipient" + ], + "title": "send_email_args", + "type": "object" + } + }, + "tool_call": { + "arguments": "{\"recipient\": \"bob@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + } + } + ], + "handoffs": [], + "interruptions": [ + { + "agent": { + "name": "compat-agent" + }, + "raw_item": { + "arguments": "{\"recipient\": \"bob@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + }, + "tool_lookup_key": { + "kind": "bare", + "name": "send_email" + }, + "tool_name": "send_email", + "type": "tool_approval_item" + } + ], + "local_shell_actions": [], + "mcp_approval_requests": [], + "new_items": [ + { + "agent": { + "name": "compat-agent" + }, + "raw_item": { + "arguments": "{\"recipient\": \"alice@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + }, + "tool_name": "send_email", + "type": "tool_call_item" + } + ], + "shell_actions": [], + "tools_used": [] + }, + "max_turns": 10, + "model_responses": [ + { + "output": [ + { + "arguments": "{\"recipient\": \"alice@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + } + ], + "request_id": null, + "response_id": "response-override", + "usage": { + "input_tokens": 0, + "input_tokens_details": [ + { + "cache_write_tokens": 0, + "cached_tokens": 0 + } + ], + "output_tokens": 0, + "output_tokens_details": [ + { + "reasoning_tokens": 0 + } + ], + "request_usage_entries": [], + "requests": 0, + "total_tokens": 0 + } + } + ], + "nested_history_owned_session_item_refs": [], + "no_active_agent_run": true, + "original_input": "historical input", + "output_guardrail_results": [], + "pending_input": [], + "previous_response_id": null, + "reasoning_item_id_policy": null, + "session_history_mutations": [], + "session_items": [ + { + "agent": { + "name": "compat-agent" + }, + "raw_item": { + "arguments": "{\"recipient\": \"alice@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + }, + "tool_name": "send_email", + "type": "tool_call_item" + } + ], + "tool_input_guardrail_results": [], + "tool_output_guardrail_results": [], + "tool_use_tracker": {}, + "trace": null +} diff --git a/tests/fixtures/run_state/features/v1_16_durable_approval_argument_overrides.json b/tests/fixtures/run_state/features/v1_16_durable_approval_argument_overrides.json new file mode 100644 index 0000000000..1bcb2d40b3 --- /dev/null +++ b/tests/fixtures/run_state/features/v1_16_durable_approval_argument_overrides.json @@ -0,0 +1,299 @@ +{ + "$schemaVersion": "1.16", + "approval_argument_override_modes": [ + { + "call_id": "approval-override-1", + "mode": "durable" + } + ], + "auto_previous_response_id": false, + "context": { + "approvals": { + "send_email": { + "approved": [ + "approval-override-1" + ], + "rejected": [] + } + }, + "context": {}, + "context_meta": { + "omitted": false, + "original_type": "mapping", + "requires_deserializer": false, + "serialized_via": "mapping" + }, + "tool_invocations": { + "approval-override-1": { + "approval_scope": "03c0b69e3d8062f4c5ad4a51305b5ac6f32321871b2433b01ea6025ed2565956", + "completed": false, + "executed": false, + "fingerprint": "05d7f7f0356c79296d60ce0ccc75dbb2d5f990eb09664cd7a7a459182cae969a", + "type": "function_call" + } + }, + "usage": { + "input_tokens": 0, + "input_tokens_details": [ + { + "cache_write_tokens": 0, + "cached_tokens": 0 + } + ], + "output_tokens": 0, + "output_tokens_details": [ + { + "reasoning_tokens": 0 + } + ], + "request_usage_entries": [], + "requests": 0, + "total_tokens": 0 + } + }, + "conversation_id": null, + "current_agent": { + "name": "compat-agent" + }, + "current_step": { + "data": { + "interruptions": [ + { + "agent": { + "name": "compat-agent" + }, + "raw_item": { + "arguments": "{\"recipient\": \"bob@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + }, + "tool_lookup_key": { + "kind": "bare", + "name": "send_email" + }, + "tool_name": "send_email", + "type": "tool_approval_item" + } + ], + "llm_end_hooks_started": true, + "response_accepted": false + }, + "type": "next_step_interruption" + }, + "current_turn": 0, + "current_turn_persisted_item_count": 0, + "generated_items": [ + { + "agent": { + "name": "compat-agent" + }, + "raw_item": { + "arguments": "{\"recipient\": \"bob@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + }, + "tool_name": "send_email", + "type": "tool_call_item" + } + ], + "generated_prompt_cache_key": null, + "generated_session_item_indexes": [ + null + ], + "input_guardrail_results": [], + "last_model_response": { + "output": [ + { + "arguments": "{\"recipient\": \"alice@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + } + ], + "request_id": null, + "response_id": "response-override", + "usage": { + "input_tokens": 0, + "input_tokens_details": [ + { + "cache_write_tokens": 0, + "cached_tokens": 0 + } + ], + "output_tokens": 0, + "output_tokens_details": [ + { + "reasoning_tokens": 0 + } + ], + "request_usage_entries": [], + "requests": 0, + "total_tokens": 0 + } + }, + "last_processed_response": { + "apply_patch_actions": [], + "computer_actions": [], + "custom_tool_actions": [], + "functions": [ + { + "tool": { + "description": "", + "lookupKey": { + "kind": "bare", + "name": "send_email" + }, + "name": "send_email", + "paramsJsonSchema": { + "additionalProperties": false, + "properties": { + "recipient": { + "title": "Recipient", + "type": "string" + } + }, + "required": [ + "recipient" + ], + "title": "send_email_args", + "type": "object" + } + }, + "tool_call": { + "arguments": "{\"recipient\": \"bob@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + } + } + ], + "handoffs": [], + "interruptions": [ + { + "agent": { + "name": "compat-agent" + }, + "raw_item": { + "arguments": "{\"recipient\": \"bob@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + }, + "tool_lookup_key": { + "kind": "bare", + "name": "send_email" + }, + "tool_name": "send_email", + "type": "tool_approval_item" + } + ], + "local_shell_actions": [], + "mcp_approval_requests": [], + "new_items": [ + { + "agent": { + "name": "compat-agent" + }, + "raw_item": { + "arguments": "{\"recipient\": \"bob@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + }, + "tool_name": "send_email", + "type": "tool_call_item" + } + ], + "shell_actions": [], + "tools_used": [] + }, + "max_turns": 10, + "model_responses": [ + { + "output": [ + { + "arguments": "{\"recipient\": \"alice@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + } + ], + "request_id": null, + "response_id": "response-override", + "usage": { + "input_tokens": 0, + "input_tokens_details": [ + { + "cache_write_tokens": 0, + "cached_tokens": 0 + } + ], + "output_tokens": 0, + "output_tokens_details": [ + { + "reasoning_tokens": 0 + } + ], + "request_usage_entries": [], + "requests": 0, + "total_tokens": 0 + } + } + ], + "nested_history_owned_session_item_refs": [], + "no_active_agent_run": true, + "original_input": "historical input", + "output_guardrail_results": [], + "pending_input": [], + "previous_response_id": null, + "reasoning_item_id_policy": null, + "session_history_mutations": [ + { + "call_id": "approval-override-1", + "expected": { + "arguments": "{\"recipient\": \"alice@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + }, + "replacement": { + "arguments": "{\"recipient\": \"bob@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + }, + "type": "replace_function_call" + } + ], + "session_items": [ + { + "agent": { + "name": "compat-agent" + }, + "raw_item": { + "arguments": "{\"recipient\": \"bob@example.com\"}", + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call" + }, + "tool_name": "send_email", + "type": "tool_call_item" + } + ], + "tool_input_guardrail_results": [], + "tool_output_guardrail_results": [], + "tool_use_tracker": {}, + "trace": null +} diff --git a/tests/fixtures/run_state/generate_corpus.py b/tests/fixtures/run_state/generate_corpus.py index 994b7dd4f3..d4503a2a11 100644 --- a/tests/fixtures/run_state/generate_corpus.py +++ b/tests/fixtures/run_state/generate_corpus.py @@ -33,11 +33,66 @@ @dataclass(frozen=True) class Scenario: version: str - commit: str + commit: str | None name: str code: str provenance: str = "historical_writer" emitted_version: str | None = None + lineage_commit: str | None = None + + +APPROVAL_OVERRIDE_SETUP = """ +from agents.items import ModelResponse, ToolApprovalItem, ToolCallItem +from agents.run_context import RunContextWrapper +from agents.run_internal.run_loop import NextStepInterruption, ProcessedResponse, ToolRunFunction +from agents.tool import function_tool +from agents.usage import Usage +from openai.types.responses import ResponseFunctionToolCall + +@function_tool(name_override="send_email") +def send_email(recipient: str) -> str: + return f"sent:{recipient}" + +agent = Agent(name="compat-agent", tools=[send_email]) +state = RunState( + context=RunContextWrapper(context={}), + original_input="historical input", + starting_agent=agent, + max_turns=10, +) +raw_item = ResponseFunctionToolCall( + type="function_call", + id="fc-override", + call_id="approval-override-1", + name="send_email", + arguments=json.dumps({"recipient": "alice@example.com"}), +) +tool_call_item = ToolCallItem(agent=agent, raw_item=raw_item) +approval_item = ToolApprovalItem( + agent=agent, + raw_item=raw_item, + tool_name="send_email", +) +state._generated_items = [tool_call_item] +state._session_items = [ToolCallItem(agent=agent, raw_item=raw_item)] +state._current_step = NextStepInterruption(interruptions=[approval_item]) +state._model_responses = [ + ModelResponse(output=[raw_item], usage=Usage(), response_id="response-override") +] +state._last_processed_response = ProcessedResponse( + new_items=[ToolCallItem(agent=agent, raw_item=raw_item)], + handoffs=[], + functions=[ToolRunFunction(tool_call=raw_item, function_tool=send_email)], + computer_actions=[], + local_shell_calls=[], + shell_calls=[], + apply_patch_calls=[], + tools_used=[], + mcp_approval_requests=[], + interruptions=[approval_item], +) +state._mark_generated_items_merged_with_last_processed() +""" SCENARIOS = ( @@ -387,6 +442,35 @@ class Scenario: state.approve(approval) """, ), + Scenario( + "1.16", + None, + "approval_argument_overrides", + APPROVAL_OVERRIDE_SETUP + + """ +state.approve( + approval_item, + override_arguments={"recipient": "bob@example.com"}, + save_override_arguments=False, +) +""", + provenance="current_writer", + lineage_commit="672358a3c0eef02c15f62c0ba449a63f9581f6ad", + ), + Scenario( + "1.16", + None, + "durable_approval_argument_overrides", + APPROVAL_OVERRIDE_SETUP + + """ +state.approve( + approval_item, + override_arguments={"recipient": "bob@example.com"}, +) +""", + provenance="current_writer", + lineage_commit="672358a3c0eef02c15f62c0ba449a63f9581f6ad", + ), ) @@ -479,9 +563,7 @@ def _extract(commit: str, destination: Path) -> None: def _generate(scenario: Scenario) -> dict[str, object]: - with tempfile.TemporaryDirectory(prefix=f"run-state-{scenario.version}-") as temp: - tree = Path(temp) - _extract(scenario.commit, tree) + def generate_from_tree(tree: Path) -> dict[str, object]: env = dict(os.environ) env["UV_DEFAULT_INDEX"] = "https://pypi.org/simple" for variable in ( @@ -526,6 +608,14 @@ def _generate(scenario: Scenario) -> dict[str, object]: payload["$schemaVersion"] = scenario.version return cast(dict[str, object], payload) + if scenario.commit is None: + return generate_from_tree(ROOT) + + with tempfile.TemporaryDirectory(prefix=f"run-state-{scenario.version}-") as temp: + tree = Path(temp) + _extract(scenario.commit, tree) + return generate_from_tree(tree) + def main() -> None: OUTPUT.mkdir(parents=True, exist_ok=True) @@ -540,16 +630,20 @@ def main() -> None: source = { "version": scenario.version, "feature": scenario.name, - "commit": scenario.commit, "fixture": f"features/{filename}", "provenance": scenario.provenance, } + if scenario.commit is not None: + source["commit"] = scenario.commit + if scenario.lineage_commit is not None: + source["lineage_commit"] = scenario.lineage_commit if scenario.emitted_version is not None: source["emitted_version"] = scenario.emitted_version source["note"] = ( "The release-boundary schema renumbering introduced this reader version " - "without a writer that emitted it. The recorded writer emitted 1.9; only " - "the schema label is changed to exercise the canonical compatibility branch." + "without a writer that emitted it. The recorded writer emitted " + f"{scenario.emitted_version}; only the schema label is changed to exercise " + "the canonical compatibility branch." ) feature_sources.append(source) diff --git a/tests/fixtures/run_state/minimal/v1_16.json b/tests/fixtures/run_state/minimal/v1_16.json new file mode 100644 index 0000000000..88163676bb --- /dev/null +++ b/tests/fixtures/run_state/minimal/v1_16.json @@ -0,0 +1,62 @@ +{ + "$schemaVersion": "1.16", + "auto_previous_response_id": false, + "context": { + "approvals": {}, + "context": {}, + "context_meta": { + "omitted": false, + "original_type": "mapping", + "requires_deserializer": false, + "serialized_via": "mapping" + }, + "tool_invocations": {}, + "usage": { + "input_tokens": 0, + "input_tokens_details": [ + { + "cache_write_tokens": 0, + "cached_tokens": 0 + } + ], + "output_tokens": 0, + "output_tokens_details": [ + { + "reasoning_tokens": 0 + } + ], + "request_usage_entries": [], + "requests": 0, + "total_tokens": 0 + } + }, + "conversation_id": null, + "current_agent": { + "name": "compat-agent" + }, + "current_step": null, + "current_turn": 0, + "current_turn_persisted_item_count": 0, + "approval_argument_override_modes": [], + "generated_items": [], + "generated_prompt_cache_key": null, + "generated_session_item_indexes": [], + "input_guardrail_results": [], + "last_model_response": null, + "last_processed_response": null, + "max_turns": 10, + "model_responses": [], + "nested_history_owned_session_item_refs": [], + "no_active_agent_run": true, + "original_input": "historical input", + "output_guardrail_results": [], + "pending_input": [], + "previous_response_id": null, + "reasoning_item_id_policy": null, + "session_history_mutations": [], + "session_items": [], + "tool_input_guardrail_results": [], + "tool_output_guardrail_results": [], + "tool_use_tracker": {}, + "trace": null +} diff --git a/tests/fixtures/run_state/sources.json b/tests/fixtures/run_state/sources.json index cdae672998..917e36273b 100644 --- a/tests/fixtures/run_state/sources.json +++ b/tests/fixtures/run_state/sources.json @@ -109,6 +109,20 @@ "fixture": "features/v1_15_canonical_invocation_identity.json", "provenance": "historical_writer", "version": "1.15" + }, + { + "feature": "approval_argument_overrides", + "fixture": "features/v1_16_approval_argument_overrides.json", + "lineage_commit": "672358a3c0eef02c15f62c0ba449a63f9581f6ad", + "provenance": "current_writer", + "version": "1.16" + }, + { + "feature": "durable_approval_argument_overrides", + "fixture": "features/v1_16_durable_approval_argument_overrides.json", + "lineage_commit": "672358a3c0eef02c15f62c0ba449a63f9581f6ad", + "provenance": "current_writer", + "version": "1.16" } ], "resume": { @@ -163,6 +177,11 @@ "commit": "4720150fde047baa4e88b16082b282bee3a5e87d", "fixture": "minimal/v1_15.json" }, + "1.16": { + "fixture": "minimal/v1_16.json", + "lineage_commit": "672358a3c0eef02c15f62c0ba449a63f9581f6ad", + "provenance": "current_writer" + }, "1.2": { "commit": "74e8c1e22d7441bd42c58bcd4270937ccc2dca8c", "fixture": "minimal/v1_2.json" diff --git a/tests/memory/test_openai_responses_compaction_session.py b/tests/memory/test_openai_responses_compaction_session.py index ae0c7e2028..8a4d5f00f7 100644 --- a/tests/memory/test_openai_responses_compaction_session.py +++ b/tests/memory/test_openai_responses_compaction_session.py @@ -15,6 +15,7 @@ from agents.memory import ( OpenAIResponsesCompactionSession, Session, + SessionHistoryRewriteArgs, SessionSettings, SQLiteSession, is_openai_responses_compaction_aware_session, @@ -31,7 +32,24 @@ ) from tests.fake_model import FakeModel from tests.test_responses import get_function_tool, get_function_tool_call, get_text_message -from tests.utils.simple_session import SimpleListSession +from tests.utils.simple_session import ( + RewriteAwareSimpleSession, + ServerManagedSimpleSession, + SimpleListSession, +) + + +def _history_function_call(arguments: str) -> TResponseInputItem: + return cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": arguments, + }, + ) class TestIsOpenAIModelName: @@ -131,6 +149,141 @@ def test_init_accepts_valid_model(self) -> None: ) assert session.model == "gpt-4.1" + def test_init_rejects_structurally_server_managed_session(self) -> None: + underlying_session = ServerManagedSimpleSession() + + with pytest.raises(ValueError, match="server-managed conversation session"): + OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=underlying_session, + ) + + @pytest.mark.asyncio + async def test_history_rewrite_rejects_replaced_server_managed_session(self) -> None: + original_item = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "name": "test_tool", + "arguments": '{"value":"old"}', + }, + ) + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=SimpleListSession(history=[original_item]), + ) + replacement = ServerManagedSimpleSession(history=[original_item]) + session.underlying_session = replacement + + with pytest.raises(ValueError, match="cannot rewrite a server-managed conversation"): + await session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": original_item, + "replacement": cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "name": "test_tool", + "arguments": '{"value":"new"}', + }, + ), + } + ] + } + ) + + assert await replacement.get_items() == [original_item] + assert session._history_rewrite_generation == 0 + + @pytest.mark.asyncio + @pytest.mark.parametrize("falsey_snapshot", [False, True], ids=["truthy", "falsey"]) + async def test_history_rewrite_rejects_underlying_session_changed_during_transaction( + self, + falsey_snapshot: bool, + ) -> None: + original_item = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "name": "test_tool", + "arguments": '{"value":"old"}', + }, + ) + replacement_item = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "name": "test_tool", + "arguments": '{"value":"new"}', + }, + ) + storage_started = asyncio.Event() + release_storage = asyncio.Event() + + class BlockingRewriteSession(RewriteAwareSimpleSession): + async def apply_history_mutations(self, args: SessionHistoryRewriteArgs) -> bool: + result = await super().apply_history_mutations(args) + storage_started.set() + await release_storage.wait() + return result + + class FalseyBlockingSession(SimpleListSession): + def __bool__(self) -> bool: + return False + + async def get_items(self, limit: int | None = None) -> list[TResponseInputItem]: + storage_started.set() + await release_storage.wait() + return await super().get_items(limit) + + underlying: SimpleListSession + if falsey_snapshot: + underlying = FalseyBlockingSession(history=[original_item]) + else: + underlying = BlockingRewriteSession(history=[original_item]) + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=underlying, + ) + server_managed_replacement = ServerManagedSimpleSession(history=[original_item]) + + rewrite_task = asyncio.create_task( + session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": original_item, + "replacement": replacement_item, + } + ] + } + ) + ) + try: + await asyncio.wait_for(storage_started.wait(), timeout=2) + session.underlying_session = server_managed_replacement + finally: + release_storage.set() + + with pytest.raises(ValueError, match="underlying session changed during history rewrite"): + await rewrite_task + + assert await server_managed_replacement.get_items() == [original_item] + assert await underlying.get_items() == [replacement_item] + assert session._session_items is None + assert session._compaction_candidate_items is None + assert session._history_rewrite_generation == 1 + @pytest.mark.asyncio async def test_add_items_delegates(self) -> None: mock_session = self.create_mock_session() @@ -160,6 +313,136 @@ async def test_get_items_delegates(self) -> None: assert len(result) == 1 mock_session.get_items.assert_called_once() + @pytest.mark.asyncio + async def test_apply_history_mutations_rewrites_underlying_history(self) -> None: + underlying = SimpleListSession( + history=[ + cast(TResponseInputItem, {"type": "message", "role": "user", "content": "hello"}), + cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"foo"}', + }, + ), + cast( + TResponseInputItem, + { + "type": "function_call_output", + "call_id": "call-1", + "output": "ok", + }, + ), + ] + ) + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=underlying, + ) + + await session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": _history_function_call('{"value":"foo"}'), + "replacement": cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"bar"}', + }, + ), + } + ] + } + ) + + saved_items = await underlying.get_items() + assert cast(dict[str, Any], saved_items[1])["arguments"] == '{"value":"bar"}' + + @pytest.mark.asyncio + async def test_apply_history_mutations_rejects_newer_divergent_reused_call_id(self) -> None: + expected_call = _history_function_call('{"value":"expected"}') + divergent_call = _history_function_call('{"value":"divergent"}') + replacement_call = _history_function_call('{"value":"corrected"}') + underlying = SimpleListSession(history=[expected_call, divergent_call]) + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=underlying, + ) + + with pytest.raises(ValueError, match="did not match the expected function call"): + await session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": expected_call, + "replacement": replacement_call, + } + ] + } + ) + + assert await underlying.get_items() == [expected_call, divergent_call] + + @pytest.mark.asyncio + async def test_apply_history_mutations_delegates_to_rewrite_aware_underlying_session( + self, + ) -> None: + underlying = RewriteAwareSimpleSession( + history=[ + cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"foo"}', + }, + ) + ] + ) + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=underlying, + ) + + await session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": _history_function_call('{"value":"foo"}'), + "replacement": cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"bar"}', + }, + ), + } + ] + } + ) + + saved_items = await underlying.get_items() + assert cast(dict[str, Any], saved_items[0])["arguments"] == '{"value":"bar"}' + @pytest.mark.asyncio async def test_run_compaction_requires_response_id(self) -> None: mock_session = self.create_mock_session() @@ -406,6 +689,542 @@ async def test_run_compaction_auto_uses_input_when_store_false(self) -> None: assert "previous_response_id" not in call_kwargs assert call_kwargs.get("input") == items + @pytest.mark.asyncio + async def test_run_compaction_forces_input_mode_after_local_history_rewrite(self) -> None: + underlying = RewriteAwareSimpleSession( + history=[ + cast(TResponseInputItem, {"type": "message", "role": "user", "content": "hello"}), + cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"foo"}', + }, + ), + cast( + TResponseInputItem, + { + "type": "function_call_output", + "call_id": "call-1", + "output": "ok", + }, + ), + ] + ) + mock_compact_response = MagicMock() + mock_compact_response.output = [] + mock_client = MagicMock() + mock_client.responses.compact = AsyncMock(return_value=mock_compact_response) + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=underlying, + client=mock_client, + compaction_mode="auto", + ) + + await session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": _history_function_call('{"value":"foo"}'), + "replacement": cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"bar"}', + }, + ), + } + ] + } + ) + await session.run_compaction({"response_id": "resp-1", "force": True}) + + first_call_kwargs = mock_client.responses.compact.call_args.kwargs + assert "previous_response_id" not in first_call_kwargs + assert isinstance(first_call_kwargs.get("input"), list) + + mock_client.responses.compact.reset_mock() + await session.run_compaction({"response_id": "resp-2", "force": True}) + + second_call_kwargs = mock_client.responses.compact.call_args.kwargs + assert second_call_kwargs.get("previous_response_id") == "resp-2" + + @pytest.mark.asyncio + async def test_deferred_candidate_load_cannot_overwrite_newer_history_rewrite(self) -> None: + original_call = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"foo"}', + }, + ) + corrected_call = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"bar"}', + }, + ) + + class BlockingSnapshotSession(RewriteAwareSimpleSession): + def __init__(self) -> None: + super().__init__(history=[original_call]) + self.first_read_started = asyncio.Event() + self.release_first_read = asyncio.Event() + self.read_count = 0 + + async def get_items(self, limit: int | None = None) -> list[TResponseInputItem]: + self.read_count += 1 + snapshot = await super().get_items(limit) + if self.read_count == 1: + self.first_read_started.set() + await self.release_first_read.wait() + return snapshot + + underlying = BlockingSnapshotSession() + mock_compact_response = MagicMock() + mock_compact_response.output = [] + mock_client = MagicMock() + mock_client.responses.compact = AsyncMock(return_value=mock_compact_response) + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=underlying, + client=mock_client, + compaction_mode="auto", + should_trigger_compaction=lambda _context: True, + ) + + deferred = asyncio.create_task(session._defer_compaction("resp-old")) + await underlying.first_read_started.wait() + rewrite = asyncio.create_task( + session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": original_call, + "replacement": corrected_call, + } + ] + } + ) + ) + await asyncio.sleep(0) + assert not rewrite.done() + + underlying.release_first_read.set() + await deferred + await rewrite + await session.run_compaction({"response_id": "resp-new", "force": True}) + + call_kwargs = mock_client.responses.compact.call_args.kwargs + assert "previous_response_id" not in call_kwargs + assert cast(list[dict[str, Any]], call_kwargs["input"])[0]["arguments"] == ( + '{"value":"bar"}' + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("compaction_mode", ["previous_response_id", "input"]) + async def test_run_compaction_discards_result_started_before_local_history_rewrite( + self, compaction_mode: str + ) -> None: + underlying = RewriteAwareSimpleSession( + history=[ + cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"foo"}', + }, + ) + ] + ) + compact_started = asyncio.Event() + release_compact = asyncio.Event() + mock_compact_response = MagicMock() + mock_compact_response.output = [{"type": "compaction", "summary": "stale"}] + + async def compact(**_kwargs: Any) -> Any: + compact_started.set() + await release_compact.wait() + return mock_compact_response + + mock_client = MagicMock() + mock_client.responses.compact = AsyncMock(side_effect=compact) + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=underlying, + client=mock_client, + compaction_mode=cast(Any, compaction_mode), + ) + + compaction_task = asyncio.create_task( + session.run_compaction({"response_id": "resp-1", "force": True}) + ) + await compact_started.wait() + await session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": _history_function_call('{"value":"foo"}'), + "replacement": cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"bar"}', + }, + ), + } + ] + } + ) + release_compact.set() + await compaction_task + + saved_items = await underlying.get_items() + assert cast(dict[str, Any], saved_items[0])["arguments"] == '{"value":"bar"}' + assert session._has_pending_local_history_rewrite is True + + call_kwargs = mock_client.responses.compact.call_args.kwargs + if compaction_mode == "previous_response_id": + assert call_kwargs["previous_response_id"] == "resp-1" + else: + assert cast(list[dict[str, Any]], call_kwargs["input"])[0]["arguments"] == ( + '{"value":"foo"}' + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("mutation", ["add", "pop"]) + async def test_run_compaction_discards_result_started_before_newer_history_mutation( + self, mutation: str + ) -> None: + corrected_call = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"bar"}', + }, + ) + original_output = cast( + TResponseInputItem, + {"type": "function_call_output", "call_id": "call-1", "output": "old"}, + ) + underlying = RewriteAwareSimpleSession( + history=[ + cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"foo"}', + }, + ), + original_output, + ] + ) + compact_started = asyncio.Event() + release_compact = asyncio.Event() + mock_compact_response = MagicMock() + mock_compact_response.output = [{"type": "compaction", "summary": "stale"}] + + async def compact(**_kwargs: Any) -> Any: + compact_started.set() + await release_compact.wait() + return mock_compact_response + + mock_client = MagicMock() + mock_client.responses.compact = AsyncMock(side_effect=compact) + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=underlying, + client=mock_client, + compaction_mode="input", + ) + await session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": _history_function_call('{"value":"foo"}'), + "replacement": corrected_call, + } + ] + } + ) + + compaction_task = asyncio.create_task( + session.run_compaction({"response_id": "resp-1", "force": True}) + ) + await compact_started.wait() + + if mutation == "add": + newer_output = cast( + TResponseInputItem, + {"type": "function_call_output", "call_id": "call-2", "output": "new"}, + ) + await session.add_items([newer_output]) + expected_items = [corrected_call, original_output, newer_output] + else: + assert await session.pop_item() == original_output + expected_items = [corrected_call] + + release_compact.set() + await compaction_task + + assert await underlying.get_items() == expected_items + assert session._history_rewrite_generation == 2 + + @pytest.mark.asyncio + @pytest.mark.parametrize("compaction_mode", ["previous_response_id", "input"]) + async def test_run_compaction_discards_stale_result_when_rewrite_is_cancelled_after_commit( + self, compaction_mode: str + ) -> None: + class CommitThenSuspendSession(RewriteAwareSimpleSession): + def __init__(self) -> None: + super().__init__( + history=[ + cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"foo"}', + }, + ) + ] + ) + self.get_count = 0 + self.post_rewrite_read_started = asyncio.Event() + + async def get_items(self, limit: int | None = None) -> list[TResponseInputItem]: + self.get_count += 1 + if self.get_count == 2: + self.post_rewrite_read_started.set() + await asyncio.Event().wait() + return await super().get_items(limit) + + underlying = CommitThenSuspendSession() + compact_started = asyncio.Event() + release_compact = asyncio.Event() + mock_compact_response = MagicMock() + mock_compact_response.output = [{"type": "compaction", "summary": "stale"}] + + async def compact(**_kwargs: Any) -> Any: + compact_started.set() + await release_compact.wait() + return mock_compact_response + + mock_client = MagicMock() + mock_client.responses.compact = AsyncMock(side_effect=compact) + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=underlying, + client=mock_client, + compaction_mode=cast(Any, compaction_mode), + ) + + compaction_task = asyncio.create_task( + session.run_compaction({"response_id": "resp-1", "force": True}) + ) + await compact_started.wait() + rewrite_task = asyncio.create_task( + session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": _history_function_call('{"value":"foo"}'), + "replacement": cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"bar"}', + }, + ), + } + ] + } + ) + ) + await underlying.post_rewrite_read_started.wait() + rewrite_task.cancel() + with pytest.raises(asyncio.CancelledError): + await rewrite_task + + release_compact.set() + await compaction_task + + assert cast(dict[str, Any], underlying._items[0])["arguments"] == '{"value":"bar"}' + assert session._history_rewrite_generation == 1 + assert session._has_pending_local_history_rewrite is True + + @pytest.mark.asyncio + @pytest.mark.parametrize("compaction_mode", ["previous_response_id", "input"]) + async def test_run_compaction_discards_stale_result_when_clear_is_cancelled_after_commit( + self, compaction_mode: str + ) -> None: + class CommitThenSuspendClearSession(SimpleListSession): + def __init__(self) -> None: + super().__init__( + history=[ + cast( + TResponseInputItem, + {"type": "message", "role": "user", "content": "old"}, + ) + ] + ) + self.clear_committed = asyncio.Event() + self.clear_count = 0 + + async def clear_session(self) -> None: + self.clear_count += 1 + await super().clear_session() + if self.clear_count == 1: + self.clear_committed.set() + await asyncio.Event().wait() + + underlying = CommitThenSuspendClearSession() + compact_started = asyncio.Event() + release_compact = asyncio.Event() + mock_compact_response = MagicMock() + mock_compact_response.output = [{"type": "compaction", "summary": "stale"}] + + async def compact(**_kwargs: Any) -> Any: + compact_started.set() + await release_compact.wait() + return mock_compact_response + + mock_client = MagicMock() + mock_client.responses.compact = AsyncMock(side_effect=compact) + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=underlying, + client=mock_client, + compaction_mode=cast(Any, compaction_mode), + ) + + compaction_task = asyncio.create_task( + session.run_compaction({"response_id": "resp-1", "force": True}) + ) + await compact_started.wait() + clear_task = asyncio.create_task(session.clear_session()) + await underlying.clear_committed.wait() + clear_task.cancel() + with pytest.raises(asyncio.CancelledError): + await clear_task + + release_compact.set() + await compaction_task + + assert underlying._items == [] + assert session._history_rewrite_generation == 1 + + @pytest.mark.asyncio + async def test_run_compaction_keeps_local_rewrite_pending_until_input_compaction_succeeds( + self, + ) -> None: + underlying = RewriteAwareSimpleSession( + history=[ + cast(TResponseInputItem, {"type": "message", "role": "user", "content": "hello"}), + cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"foo"}', + }, + ), + cast( + TResponseInputItem, + { + "type": "function_call_output", + "call_id": "call-1", + "output": "ok", + }, + ), + ] + ) + mock_compact_response = MagicMock() + mock_compact_response.output = [] + mock_client = MagicMock() + mock_client.responses.compact = AsyncMock(return_value=mock_compact_response) + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=underlying, + client=mock_client, + compaction_mode="auto", + ) + + await session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": _history_function_call('{"value":"foo"}'), + "replacement": cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"bar"}', + }, + ), + } + ] + } + ) + + await session.run_compaction({"response_id": "resp-1"}) + mock_client.responses.compact.assert_not_called() + + await session.run_compaction({"response_id": "resp-2", "force": True}) + + call_kwargs = mock_client.responses.compact.call_args.kwargs + assert "previous_response_id" not in call_kwargs + assert isinstance(call_kwargs.get("input"), list) + assert cast(dict[str, Any], call_kwargs["input"][1])["arguments"] == '{"value":"bar"}' + @pytest.mark.asyncio async def test_run_compaction_auto_uses_default_store_when_unset(self) -> None: mock_session = self.create_mock_session() diff --git a/tests/memory/test_session.py b/tests/memory/test_session.py index 8c8bf5ea0d..50c1b8624c 100644 --- a/tests/memory/test_session.py +++ b/tests/memory/test_session.py @@ -1,6 +1,7 @@ """Tests for session memory functionality.""" import asyncio +import json import sqlite3 import tempfile import threading @@ -10,11 +11,160 @@ import pytest from agents import Agent, RunConfig, Runner, SessionSettings, SQLiteSession, TResponseInputItem +from agents.memory import SessionHistoryMutation, apply_session_history_mutations from agents.memory.sqlite_session import _await_mutation from tests.fake_model import FakeModel from tests.test_responses import get_text_message +class _RecordingLock: + def __init__(self, lock): + self._lock = lock + self.enter_count = 0 + + def __enter__(self): + self.enter_count += 1 + return self._lock.__enter__() + + def __exit__(self, exc_type, exc, tb): + return self._lock.__exit__(exc_type, exc, tb) + + +def test_history_mutation_rewrites_only_latest_expected_reused_call_id() -> None: + earlier_call = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "reused-call", + "name": "tool", + "arguments": '{"value":"earlier"}', + }, + ) + expected_call = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "reused-call", + "name": "tool", + "arguments": '{"value":"current"}', + }, + ) + replacement_call = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "reused-call", + "name": "tool", + "arguments": '{"value":"corrected"}', + }, + ) + items: list[TResponseInputItem] = [ + earlier_call, + {"role": "user", "content": "later"}, + expected_call, + ] + mutation: SessionHistoryMutation = { + "type": "replace_function_call", + "call_id": "reused-call", + "expected": expected_call, + "replacement": replacement_call, + } + + rewritten = apply_session_history_mutations(items, [mutation]) + + assert rewritten == [earlier_call, {"role": "user", "content": "later"}, replacement_call] + assert apply_session_history_mutations(rewritten, [mutation]) == rewritten + + +def test_history_mutation_rejects_divergent_reused_call_id() -> None: + expected_call = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "reused-call", + "name": "tool", + "arguments": '{"value":"expected"}', + }, + ) + persisted_call = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "reused-call", + "name": "tool", + "arguments": '{"value":"different"}', + }, + ) + replacement_call = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "reused-call", + "name": "tool", + "arguments": '{"value":"corrected"}', + }, + ) + + with pytest.raises(ValueError, match="did not match the expected function call"): + apply_session_history_mutations( + [persisted_call], + [ + { + "type": "replace_function_call", + "call_id": "reused-call", + "expected": expected_call, + "replacement": replacement_call, + } + ], + ) + + +@pytest.mark.parametrize("older_arguments", ['{"value":"expected"}', '{"value":"corrected"}']) +def test_history_mutation_rejects_newer_divergent_reused_call_id( + older_arguments: str, +) -> None: + """A newer same-ID call is the only eligible canonical mutation target.""" + expected_call = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "reused-call", + "name": "tool", + "arguments": '{"value":"expected"}', + }, + ) + replacement_call = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "reused-call", + "name": "tool", + "arguments": '{"value":"corrected"}', + }, + ) + older_call = cast( + TResponseInputItem, + {**expected_call, "arguments": older_arguments}, + ) + newer_call = cast( + TResponseInputItem, + {**expected_call, "arguments": '{"value":"divergent"}'}, + ) + + with pytest.raises(ValueError, match="did not match the expected function call"): + apply_session_history_mutations( + [older_call, newer_call], + [ + { + "type": "replace_function_call", + "call_id": "reused-call", + "expected": expected_call, + "replacement": replacement_call, + } + ], + ) + + @pytest.mark.asyncio async def test_await_mutation_cancellation_hides_later_failure_without_loop_error() -> None: """A failed mutation must not leak a false loop error after caller cancellation.""" @@ -812,6 +962,190 @@ async def test_sqlite_session_failed_add_items_releases_write_lock(): session.close() +@pytest.mark.asyncio +async def test_sqlite_session_apply_history_mutations_uses_file_lock(): + """File-backed history rewrites should reuse the session lock.""" + with tempfile.TemporaryDirectory() as temp_dir: + db_path = Path(temp_dir) / "test_rewrite_lock.db" + session = SQLiteSession("rewrite_lock_test", db_path) + function_call: TResponseInputItem = { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"before"}', + } + replacement: TResponseInputItem = { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"after"}', + } + + await session.add_items([function_call]) + recording_lock = _RecordingLock(session._lock) + session.__dict__["_lock"] = recording_lock + + await session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": function_call, + "replacement": replacement, + } + ] + } + ) + assert recording_lock.enter_count == 1 + + retrieved = await session.get_items() + assert retrieved == [replacement] + session.close() + + +@pytest.mark.asyncio +async def test_sqlite_session_history_rewrite_requires_matching_target(): + """A missing mutation target should leave SQLite history unchanged.""" + with tempfile.TemporaryDirectory() as temp_dir: + session = SQLiteSession("rewrite_missing_test", Path(temp_dir) / "rewrite_missing.db") + function_call: TResponseInputItem = { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"before"}', + } + await session.add_items([function_call]) + + with pytest.raises(ValueError, match="did not match the expected function call"): + await session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "missing-call", + "expected": { + "type": "function_call", + "call_id": "missing-call", + "id": "fc_missing", + "name": "test_tool", + "arguments": '{"value":"before"}', + }, + "replacement": { + "type": "function_call", + "call_id": "missing-call", + "id": "fc_missing", + "name": "test_tool", + "arguments": '{"value":"after"}', + }, + } + ] + } + ) + + assert await session.get_items() == [function_call] + session.close() + + +@pytest.mark.asyncio +async def test_sqlite_session_history_rewrite_rejects_newer_divergent_reused_call_id(): + """SQLite should not rewrite an older call shadowed by a newer same-ID call.""" + with tempfile.TemporaryDirectory() as temp_dir: + session = SQLiteSession("rewrite_reused_test", Path(temp_dir) / "rewrite_reused.db") + expected_call: TResponseInputItem = { + "type": "function_call", + "call_id": "call-1", + "name": "test_tool", + "arguments": '{"value":"expected"}', + } + divergent_call: TResponseInputItem = { + **expected_call, + "arguments": '{"value":"divergent"}', + } + replacement_call: TResponseInputItem = { + **expected_call, + "arguments": '{"value":"corrected"}', + } + original_items = [expected_call, divergent_call] + await session.add_items(original_items) + + with pytest.raises(ValueError, match="did not match the expected function call"): + await session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": expected_call, + "replacement": replacement_call, + } + ] + } + ) + + assert await session.get_items() == original_items + session.close() + + +@pytest.mark.asyncio +async def test_sqlite_session_history_rewrite_preserves_unrelated_corrupt_rows(): + """A targeted rewrite should not delete an unrelated undecodable record.""" + with tempfile.TemporaryDirectory() as temp_dir: + session = SQLiteSession("rewrite_corrupt_test", Path(temp_dir) / "rewrite_corrupt.db") + function_call: TResponseInputItem = { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"before"}', + } + replacement = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "id": "fc_1", + "name": "test_tool", + "arguments": '{"value":"after"}', + }, + ) + await session.add_items([function_call]) + with session._write_connection() as conn: + conn.execute( + f"INSERT INTO {session.messages_table} (session_id, message_data) VALUES (?, ?)", + (session.session_id, "not-json"), + ) + conn.commit() + + await session.apply_history_mutations( + { + "mutations": [ + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": function_call, + "replacement": replacement, + } + ] + } + ) + + with session._write_connection() as conn: + rows = [ + row[0] + for row in conn.execute( + f"SELECT message_data FROM {session.messages_table} " + "WHERE session_id = ? ORDER BY id", + (session.session_id,), + ).fetchall() + ] + assert rows == [json.dumps(replacement), "not-json"] + session.close() + + @pytest.mark.asyncio async def test_session_add_items_exception_propagates_in_streamed(): """Test that exceptions from session.add_items are properly propagated diff --git a/tests/test_agent_runner.py b/tests/test_agent_runner.py index 1cce41035c..9c1c0e89d2 100644 --- a/tests/test_agent_runner.py +++ b/tests/test_agent_runner.py @@ -8,7 +8,7 @@ from collections.abc import Callable from pathlib import Path from typing import Any, cast -from unittest.mock import patch +from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest @@ -67,7 +67,7 @@ TResponseInputItem, ) from agents.lifecycle import RunHooks -from agents.memory import SessionSettings +from agents.memory import OpenAIResponsesCompactionSession, SessionSettings from agents.models.fake_id import FAKE_RESPONSES_ID from agents.run import AgentRunner, get_default_agent_runner, set_default_agent_runner from agents.run_config import _default_trace_include_sensitive_data @@ -110,8 +110,14 @@ get_text_message, ) from .utils.factories import make_run_state -from .utils.hitl import make_context_wrapper, make_model_and_agent, make_shell_call -from .utils.simple_session import CountingSession, IdStrippingSession, SimpleListSession +from .utils.hitl import consume_stream, make_context_wrapper, make_model_and_agent, make_shell_call +from .utils.simple_session import ( + CountingSession, + IdStrippingSession, + RewriteAwareSimpleSession, + ServerManagedSimpleSession, + SimpleListSession, +) class _DummyRunItem: @@ -1652,191 +1658,1322 @@ def capture_model_input(data): assert "function_call" not in handoff_input_types assert "function_call_output" not in handoff_input_types - assert any(isinstance(item, ToolCallOutputItem) for item in result.new_items) - assert any(isinstance(item, HandoffOutputItem) for item in result.new_items) + assert any(isinstance(item, ToolCallOutputItem) for item in result.new_items) + assert any(isinstance(item, HandoffOutputItem) for item in result.new_items) + + session_items = await session.get_items() + has_function_call_output = any( + isinstance(item, dict) and item.get("type") == "function_call_output" + for item in session_items + ) + assert has_function_call_output + + +@pytest.mark.asyncio +async def test_nested_handoff_filters_reasoning_items_from_model_input(): + model = FakeModel() + delegate = Agent( + name="delegate", + model=model, + ) + triage = Agent( + name="triage", + model=model, + handoffs=[delegate], + ) + + model.add_multiple_turn_outputs( + [ + [ + ResponseReasoningItem( + id="reasoning_1", + type="reasoning", + summary=[Summary(text="Thinking about a handoff.", type="summary_text")], + ), + get_handoff_tool_call(delegate), + ], + [get_text_message("done")], + ] + ) + + captured_inputs: list[list[dict[str, Any]]] = [] + + def capture_model_input(data): + if isinstance(data.model_data.input, list): + captured_inputs.append( + [item for item in data.model_data.input if isinstance(item, dict)] + ) + return data.model_data + + result = await Runner.run( + triage, + input="user_message", + run_config=RunConfig( + nest_handoff_history=True, + call_model_input_filter=capture_model_input, + ), + ) + + assert result.final_output == "done" + assert len(captured_inputs) >= 2 + handoff_input = captured_inputs[1] + handoff_input_types = [ + item["type"] for item in handoff_input if isinstance(item.get("type"), str) + ] + assert "reasoning" not in handoff_input_types + + +@pytest.mark.asyncio +async def test_resume_preserves_filtered_model_input_after_handoff(): + model = FakeModel() + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool() -> str: + return "ok" + + delegate = Agent( + name="delegate", + model=model, + tools=[approval_tool], + ) + triage = Agent( + name="triage", + model=model, + handoffs=[delegate], + tools=[get_function_tool("some_function", "result")], + ) + + model.add_multiple_turn_outputs( + [ + [ + get_function_tool_call( + "some_function", json.dumps({"a": "b"}), call_id="triage-call" + ) + ], + [get_text_message("a_message"), get_handoff_tool_call(delegate)], + [get_function_tool_call("approval_tool", json.dumps({}), call_id="delegate-call")], + [get_text_message("done")], + ] + ) + + model_input_call_ids: list[set[str]] = [] + model_input_output_call_ids: list[set[str]] = [] + + def capture_model_input(data): + call_ids: set[str] = set() + output_call_ids: set[str] = set() + for item in data.model_data.input: + if not isinstance(item, dict): + continue + item_type = item.get("type") + call_id = item.get("call_id") + if not isinstance(call_id, str): + continue + if item_type == "function_call": + call_ids.add(call_id) + elif item_type == "function_call_output": + output_call_ids.add(call_id) + model_input_call_ids.append(call_ids) + model_input_output_call_ids.append(output_call_ids) + return data.model_data + + run_config = RunConfig( + nest_handoff_history=True, + call_model_input_filter=capture_model_input, + ) + + first = await Runner.run(triage, input="user_message", run_config=run_config) + assert first.interruptions + + state = first.to_state() + state.approve(first.interruptions[0]) + + resumed = await Runner.run(triage, state, run_config=run_config) + + last_call_ids = model_input_call_ids[-1] + last_output_call_ids = model_input_output_call_ids[-1] + assert "triage-call" not in last_call_ids + assert "triage-call" not in last_output_call_ids + assert "delegate-call" in last_call_ids + assert "delegate-call" in last_output_call_ids + assert resumed.final_output == "done" + + +@pytest.mark.asyncio +async def test_resumed_state_updates_agent_after_handoff() -> None: + model = FakeModel() + + @function_tool(name_override="triage_tool", needs_approval=True) + def triage_tool() -> str: + return "ok" + + @function_tool(name_override="delegate_tool", needs_approval=True) + def delegate_tool() -> str: + return "ok" + + delegate = Agent( + name="delegate", + model=model, + tools=[delegate_tool], + ) + triage = Agent( + name="triage", + model=model, + handoffs=[delegate], + tools=[triage_tool], + ) + + model.add_multiple_turn_outputs( + [ + [get_function_tool_call("triage_tool", "{}", call_id="triage-1")], + [get_text_message("handoff"), get_handoff_tool_call(delegate)], + [get_function_tool_call("delegate_tool", "{}", call_id="delegate-1")], + ] + ) + + first = await Runner.run(triage, input="user_message") + assert first.interruptions + + state = first.to_state() + state.approve(first.interruptions[0]) + + second = await Runner.run(triage, state) + assert second.interruptions + assert any(item.tool_name == delegate_tool.name for item in second.interruptions), ( + "handoff should switch approvals to the delegate agent" + ) + assert state._current_agent is delegate + + +@pytest.mark.asyncio +async def test_resume_with_durable_override_rewrites_local_session_history() -> None: + model = FakeModel() + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + return f"result:{test}" + + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = RewriteAwareSimpleSession() + + model.add_multiple_turn_outputs( + [ + [ + get_function_tool_call( + "approval_tool", json.dumps({"test": "foo"}), call_id="call-1" + ) + ], + [get_text_message("done")], + ] + ) + + first = await Runner.run(agent, input="user_message", session=session) + assert first.interruptions + + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + + resumed = await Runner.run(agent, state, session=session) + + assert resumed.final_output == "done" + assert cast(Any, resumed.raw_responses[0].output[0]).arguments == json.dumps({"test": "foo"}) + saved_items = await session.get_items() + assert saved_items[1]["type"] == "function_call" + assert cast(dict[str, Any], saved_items[1])["arguments"] == json.dumps({"test": "bar"}) + assert saved_items[2]["type"] == "function_call_output" + assert saved_items[2]["call_id"] == "call-1" + assert state.to_json()["approval_argument_override_modes"] == [] + + +async def _run_approval_recovery_case( + agent: Agent[Any], + input: str | RunState[Any, Agent[Any]], + *, + streamed: bool, + run_config: RunConfig | None = None, + session: RewriteAwareSimpleSession | None = None, + previous_response_id: str | None = "resp-root", +) -> Any: + if streamed: + result = Runner.run_streamed( + agent, + input, + session=session, + previous_response_id=previous_response_id, + run_config=run_config, + ) + await consume_stream(result) + return result + return await Runner.run( + agent, + input, + session=session, + previous_response_id=previous_response_id, + run_config=run_config, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streamed", [False, True]) +async def test_execution_only_override_failure_state_roundtrips(streamed: bool) -> None: + @function_tool(needs_approval=True, failure_error_function=None) + async def failing_tool(value: str) -> str: + raise RuntimeError(f"failed:{value}") + + model = FakeModel() + model.set_next_output( + [ + get_function_tool_call( + "failing_tool", + json.dumps({"value": "original"}), + call_id="call-failure", + ) + ] + ) + agent = Agent(name="approval_agent", model=model, tools=[failing_tool]) + first = await _run_approval_recovery_case(agent, "go", streamed=streamed) + state = first.to_state() + state.approve( + first.interruptions[0], + override_arguments={"value": "corrected"}, + save_override_arguments=False, + ) + + with pytest.raises(UserError, match="Error running tool failing_tool: failed:corrected"): + await _run_approval_recovery_case(agent, state, streamed=streamed) + + restored = await RunState.from_json(agent, state.to_json()) + assert restored.to_json()["approval_argument_override_modes"] == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streamed", [False, True]) +async def test_execution_only_override_cancelled_state_roundtrips(streamed: bool) -> None: + started = asyncio.Event() + + @function_tool(needs_approval=True, failure_error_function=None) + async def blocking_tool(value: str) -> str: + started.set() + await asyncio.Event().wait() + return value + + model = FakeModel() + model.set_next_output( + [ + get_function_tool_call( + "blocking_tool", + json.dumps({"value": "original"}), + call_id="call-cancelled", + ) + ] + ) + agent = Agent(name="approval_agent", model=model, tools=[blocking_tool]) + first = await _run_approval_recovery_case(agent, "go", streamed=streamed) + state = first.to_state() + state.approve( + first.interruptions[0], + override_arguments={"value": "corrected"}, + save_override_arguments=False, + ) + + if streamed: + resumed = Runner.run_streamed(agent, state, previous_response_id="resp-root") + consumer = asyncio.create_task(consume_stream(resumed)) + await started.wait() + resumed.cancel(mode="immediate") + await consumer + else: + task = asyncio.create_task(Runner.run(agent, state, previous_response_id="resp-root")) + await started.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + restored = await RunState.from_json(agent, state.to_json()) + assert restored.to_json()["approval_argument_override_modes"] == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streamed", [False, True]) +async def test_execution_only_multi_call_partial_failure_preserves_unstarted_mode( + streamed: bool, +) -> None: + calls: list[tuple[str, str]] = [] + + @function_tool(needs_approval=True, failure_error_function=None) + async def failing_tool(value: str) -> str: + calls.append(("failing", value)) + raise RuntimeError(f"failed:{value}") + + @function_tool(needs_approval=True, failure_error_function=None) + async def pending_tool(value: str) -> str: + calls.append(("pending", value)) + return value + + model = FakeModel() + model.set_next_output( + [ + get_function_tool_call( + "failing_tool", + json.dumps({"value": "original-failure"}), + call_id="call-failure", + ), + get_function_tool_call( + "pending_tool", + json.dumps({"value": "original-pending"}), + call_id="call-pending", + ), + ] + ) + agent = Agent( + name="approval_agent", + model=model, + tools=[failing_tool, pending_tool], + ) + run_config = RunConfig(tool_execution=ToolExecutionConfig(max_function_tool_concurrency=1)) + first = await _run_approval_recovery_case( + agent, + "go", + streamed=streamed, + run_config=run_config, + ) + state = first.to_state() + for interruption in first.interruptions: + state.approve( + interruption, + override_arguments={ + "value": ( + "corrected-failure" + if interruption.tool_name == "failing_tool" + else "corrected-pending" + ) + }, + save_override_arguments=False, + ) + + with pytest.raises(UserError, match="Error running tool failing_tool: failed:corrected"): + await _run_approval_recovery_case( + agent, + state, + streamed=streamed, + run_config=run_config, + ) + + assert calls == [("failing", "corrected-failure")] + assert state.to_json()["approval_argument_override_modes"] == [ + {"call_id": "call-pending", "mode": "execution_only"} + ] + restored = await RunState.from_json(agent, state.to_json()) + assert restored.to_json()["approval_argument_override_modes"] == [ + {"call_id": "call-pending", "mode": "execution_only"} + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streamed", [False, True]) +async def test_durable_multi_call_partial_failure_preserves_unstarted_override( + streamed: bool, +) -> None: + calls: list[tuple[str, str]] = [] + + @function_tool(needs_approval=True, failure_error_function=None) + async def failing_tool(value: str) -> str: + calls.append(("failing", value)) + raise RuntimeError(f"failed:{value}") + + @function_tool(needs_approval=True, failure_error_function=None) + async def pending_tool(value: str) -> str: + calls.append(("pending", value)) + return value + + model = FakeModel() + model.add_multiple_turn_outputs( + [ + [ + get_function_tool_call( + "failing_tool", + json.dumps({"value": "original-failure"}), + call_id="call-failure", + ), + get_function_tool_call( + "pending_tool", + json.dumps({"value": "original-pending"}), + call_id="call-pending", + ), + ], + [get_text_message("done")], + ] + ) + agent = Agent( + name="approval_agent", + model=model, + tools=[failing_tool, pending_tool], + ) + session = RewriteAwareSimpleSession() + run_config = RunConfig(tool_execution=ToolExecutionConfig(max_function_tool_concurrency=1)) + first = await _run_approval_recovery_case( + agent, + "go", + streamed=streamed, + run_config=run_config, + session=session, + previous_response_id=None, + ) + state = first.to_state() + for interruption in first.interruptions: + state.approve( + interruption, + override_arguments={ + "value": ( + "corrected-failure" + if interruption.tool_name == "failing_tool" + else "corrected-pending" + ) + }, + ) + + with pytest.raises(UserError, match="Error running tool failing_tool: failed:corrected"): + await _run_approval_recovery_case( + agent, + state, + streamed=streamed, + run_config=run_config, + session=session, + previous_response_id=None, + ) + + assert calls == [("failing", "corrected-failure")] + serialized = state.to_json() + assert serialized["approval_argument_override_modes"] == [ + {"call_id": "call-pending", "mode": "durable"} + ] + assert [mutation["call_id"] for mutation in serialized["session_history_mutations"]] == [ + "call-pending" + ] + + restored = await RunState.from_json(agent, serialized) + assert restored.to_json()["approval_argument_override_modes"] == [ + {"call_id": "call-pending", "mode": "durable"} + ] + assert [ + mutation["call_id"] for mutation in restored.to_json()["session_history_mutations"] + ] == ["call-pending"] + + with pytest.raises( + ModelBehaviorError, + match="A tool call already executed, but its output was not committed", + ): + await _run_approval_recovery_case( + agent, + restored, + streamed=streamed, + run_config=run_config, + session=session, + previous_response_id=None, + ) + + assert calls == [("failing", "corrected-failure")] + assert [ + mutation["call_id"] for mutation in restored.to_json()["session_history_mutations"] + ] == ["call-pending"] + + +class _TrackingRewriteSession(RewriteAwareSimpleSession): + def __init__(self) -> None: + super().__init__() + self.rewrite_calls = 0 + + async def apply_history_mutations(self, args: Any) -> bool: + self.rewrite_calls += 1 + return await super().apply_history_mutations(args) + + +async def _build_nested_approval_override_state( + case_id: str, +) -> tuple[Agent[Any], RunState[Any, Agent[Any]], _TrackingRewriteSession, list[str]]: + calls: list[str] = [] + + @function_tool(needs_approval=True) + async def inner_tool(value: str) -> str: + calls.append(value) + return value + + nested_model = FakeModel() + nested_agent = Agent(name="nested", model=nested_model, tools=[inner_tool]) + nested_model.add_multiple_turn_outputs( + [ + [ + get_function_tool_call( + "inner_tool", + json.dumps({"value": "original"}), + call_id=f"inner-call-{case_id}", + ) + ], + [get_text_message("nested done")], + ] + ) + + session = _TrackingRewriteSession() + run_config = RunConfig(tracing_disabled=True) + outer_model = FakeModel() + outer_agent = Agent( + name="outer", + model=outer_model, + tools=[ + nested_agent.as_tool( + tool_name="nested_tool", + tool_description="Run the nested agent.", + needs_approval=True, + session=session, + run_config=run_config, + failure_error_function=None, + ) + ], + ) + outer_model.add_multiple_turn_outputs( + [ + [ + get_function_tool_call( + "nested_tool", + json.dumps({"input": "go"}), + call_id=f"outer-call-{case_id}", + ) + ], + [get_text_message("done")], + ] + ) + + first = await Runner.run(outer_agent, "start", run_config=run_config) + first_state = first.to_state() + first_state.approve(first.interruptions[0], always_approve=True) + second = await Runner.run(outer_agent, first_state, run_config=run_config) + assert second.interruptions + return outer_agent, second.to_state(), session, calls + + +@pytest.mark.asyncio +@pytest.mark.parametrize("round_trip", [False, True], ids=["live", "serialized"]) +async def test_nested_durable_override_rewrites_local_session_history( + round_trip: bool, +) -> None: + outer_agent, state, session, calls = await _build_nested_approval_override_state( + f"durable-{round_trip}" + ) + state.approve(state.get_interruptions()[0], override_arguments={"value": "corrected"}) + if round_trip: + state = await RunState.from_json(outer_agent, state.to_json()) + + resumed = await Runner.run( + outer_agent, + state, + run_config=RunConfig(tracing_disabled=True), + ) + + assert resumed.final_output == "done" + assert resumed.interruptions == [] + assert calls == ["corrected"] + assert session.rewrite_calls == 1 + saved_calls = [ + item + for item in await session.get_items() + if isinstance(item, dict) and item.get("type") == "function_call" + ] + assert len(saved_calls) == 1 + assert saved_calls[0]["arguments"] == json.dumps({"value": "corrected"}) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("round_trip", [False, True], ids=["live", "serialized"]) +async def test_nested_execution_only_override_rejects_local_session_history( + round_trip: bool, +) -> None: + outer_agent, state, session, calls = await _build_nested_approval_override_state( + f"execution-only-{round_trip}" + ) + state.approve( + state.get_interruptions()[0], + override_arguments={"value": "corrected"}, + save_override_arguments=False, + ) + if round_trip: + state = await RunState.from_json(outer_agent, state.to_json()) + + with pytest.raises(UserError, match="save_override_arguments=False is only supported"): + await Runner.run( + outer_agent, + state, + run_config=RunConfig(tracing_disabled=True), + ) + + assert calls == [] + assert session.rewrite_calls == 0 + + +@pytest.mark.asyncio +async def test_invalid_nested_override_preserves_cached_run_result() -> None: + from agents.agent_tool_state import peek_agent_tool_run_result + + outer_agent, state, _session, calls = await _build_nested_approval_override_state( + "invalid-override" + ) + assert state._last_processed_response is not None + outer_call = state._last_processed_response.functions[0].tool_call + pending_result = peek_agent_tool_run_result( + outer_call, + scope_id=state._agent_tool_state_scope_id, + ) + assert pending_result is not None + + with pytest.raises(UserError, match="cannot be used together with always_approve"): + state.approve( + state.get_interruptions()[0], + always_approve=True, + override_arguments={"value": "corrected"}, + ) + + assert ( + peek_agent_tool_run_result( + outer_call, + scope_id=state._agent_tool_state_scope_id, + ) + is pending_result + ) + assert calls == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("round_trip", [False, True], ids=["live", "serialized"]) +async def test_nested_execution_only_sync_cancellation_state_roundtrips( + round_trip: bool, +) -> None: + from agents.agent_tool_state import peek_agent_tool_run_result + + config = RunConfig(tracing_disabled=True) + case_id = "serialized" if round_trip else "live" + tool_started = asyncio.Event() + tool_stopped = asyncio.Event() + calls: list[str] = [] + + @function_tool(needs_approval=True, failure_error_function=None) + async def inner_tool(value: str) -> str: + calls.append(value) + tool_started.set() + try: + await asyncio.Event().wait() + finally: + tool_stopped.set() + return value + + nested_model = FakeModel() + nested_agent = Agent(name="nested", model=nested_model, tools=[inner_tool]) + nested_model.add_multiple_turn_outputs( + [ + [ + get_function_tool_call( + "inner_tool", + json.dumps({"value": "original"}), + call_id=f"inner-call-{case_id}", + ) + ], + [get_text_message("nested done")], + ] + ) + + outer_model = FakeModel() + outer_agent = Agent( + name="outer", + model=outer_model, + tools=[ + nested_agent.as_tool( + tool_name="nested_tool", + tool_description="Run the nested agent.", + needs_approval=True, + session=ServerManagedSimpleSession(), + run_config=config, + failure_error_function=None, + ) + ], + ) + outer_model.add_multiple_turn_outputs( + [ + [ + get_function_tool_call( + "nested_tool", + json.dumps({"input": "go"}), + call_id=f"outer-call-{case_id}", + ) + ], + [get_text_message("outer done")], + ] + ) + + first = await Runner.run(outer_agent, "go", run_config=config) + state = first.to_state() + state.approve(first.interruptions[0], always_approve=True) + second = await Runner.run(outer_agent, state, run_config=config) + state = second.to_state() + state.approve( + state.get_interruptions()[0], + override_arguments={"value": "corrected"}, + save_override_arguments=False, + ) + if round_trip: + state = await RunState.from_json(outer_agent, state.to_json()) + + resume_task = asyncio.create_task(Runner.run(outer_agent, state, run_config=config)) + try: + await asyncio.wait_for(tool_started.wait(), timeout=2) + resume_task.cancel() + with pytest.raises(asyncio.CancelledError): + await resume_task + + immediate_payload = state.to_json() + nested_payload = immediate_payload["last_processed_response"]["functions"][0][ + "agent_run_state" + ] + assert nested_payload["approval_argument_override_modes"] == [] + inner_invocation = nested_payload["context"]["tool_invocations"][f"inner-call-{case_id}"] + assert inner_invocation["executed"] is True + assert inner_invocation["completed"] is False + + restored = await RunState.from_json(outer_agent, immediate_payload) + assert restored._last_processed_response is not None + restored_outer_call = restored._last_processed_response.functions[0].tool_call + assert ( + peek_agent_tool_run_result( + restored_outer_call, + scope_id=restored._agent_tool_state_scope_id, + ) + is not None + ) + finally: + if not resume_task.done(): + resume_task.cancel() + await asyncio.gather(resume_task, return_exceptions=True) + await asyncio.wait_for(tool_stopped.wait(), timeout=2) + + assert calls == ["corrected"] + + +@pytest.mark.asyncio +async def test_durable_override_rewrite_fails_before_tool_execution() -> None: + calls: list[str] = [] + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + calls.append(test) + return f"result:{test}" + + class FailingRewriteSession(RewriteAwareSimpleSession): + async def apply_history_mutations(self, _args: Any) -> bool: + raise RuntimeError("rewrite failed") + + model = FakeModel() + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = FailingRewriteSession() + model.set_next_output( + [get_function_tool_call("approval_tool", '{"test":"foo"}', call_id="call-1")] + ) + + first = await Runner.run(agent, input="user_message", session=session) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + + with pytest.raises(RuntimeError, match="rewrite failed"): + await Runner.run(agent, state, session=session) + + assert calls == [] + assert state._get_session_history_mutations() + + +@pytest.mark.asyncio +async def test_durable_override_missing_history_fails_before_tool_execution() -> None: + calls: list[str] = [] + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + calls.append(test) + return f"result:{test}" + + model = FakeModel() + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = RewriteAwareSimpleSession() + model.set_next_output( + [get_function_tool_call("approval_tool", '{"test":"foo"}', call_id="call-1")] + ) + + first = await Runner.run(agent, input="user_message", session=session) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + await session.clear_session() + + with pytest.raises(ValueError, match="did not match the expected function call"): + await Runner.run(agent, state, session=session) + + assert calls == [] + assert state._get_session_history_mutations() + + +@pytest.mark.asyncio +async def test_durable_override_requires_rewrite_confirmation_before_tool_execution() -> None: + calls: list[str] = [] + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + calls.append(test) + return f"result:{test}" + + class UnconfirmedRewriteSession(RewriteAwareSimpleSession): + async def apply_history_mutations(self, _args: Any) -> bool: + return False + + model = FakeModel() + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = UnconfirmedRewriteSession() + model.set_next_output( + [get_function_tool_call("approval_tool", '{"test":"foo"}', call_id="call-1")] + ) + + first = await Runner.run(agent, input="user_message", session=session) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + + with pytest.raises(UserError, match="did not confirm that every target call was rewritten"): + await Runner.run(agent, state, session=session) + + assert calls == [] + assert state._get_session_history_mutations() + + +@pytest.mark.asyncio +async def test_malformed_durable_override_state_fails_before_tool_execution() -> None: + calls: list[str] = [] + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + calls.append(test) + return f"result:{test}" + + model = FakeModel() + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = RewriteAwareSimpleSession() + model.set_next_output( + [get_function_tool_call("approval_tool", '{"test":"foo"}', call_id="call-1")] + ) + first = await Runner.run(agent, input="user_message", session=session) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + serialized = state.to_json() + serialized["session_history_mutations"][0]["replacement"] = None + + with pytest.raises(UserError, match="invalid function-call replacement"): + await RunState.from_json(agent, serialized) + + assert calls == [] + + +@pytest.mark.asyncio +async def test_missing_durable_override_mutation_fails_before_tool_execution() -> None: + calls: list[str] = [] + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + calls.append(test) + return f"result:{test}" + + model = FakeModel() + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = RewriteAwareSimpleSession() + model.set_next_output( + [get_function_tool_call("approval_tool", '{"test":"foo"}', call_id="call-1")] + ) + first = await Runner.run(agent, input="user_message", session=session) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + serialized = state.to_json() + serialized["session_history_mutations"] = [] + + with pytest.raises(UserError, match="is inconsistent with pending function calls"): + await RunState.from_json(agent, serialized) + + assert calls == [] + + +@pytest.mark.asyncio +async def test_mismatched_durable_override_state_fails_before_tool_execution() -> None: + calls: list[str] = [] + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + calls.append(test) + return f"result:{test}" + + model = FakeModel() + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = RewriteAwareSimpleSession() + model.set_next_output( + [get_function_tool_call("approval_tool", '{"test":"foo"}', call_id="call-1")] + ) + first = await Runner.run(agent, input="user_message", session=session) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + serialized = state.to_json() + serialized["session_history_mutations"][0]["replacement"]["arguments"] = '{"test":"different"}' + + with pytest.raises(UserError, match="does not match the pending function call"): + await RunState.from_json(agent, serialized) + + assert calls == [] + + +@pytest.mark.asyncio +async def test_processed_durable_override_mismatch_fails_before_session_rewrite() -> None: + calls: list[str] = [] + + class TrackingRewriteSession(RewriteAwareSimpleSession): + def __init__(self) -> None: + super().__init__() + self.rewrite_calls = 0 + + async def apply_history_mutations(self, args: Any) -> bool: + self.rewrite_calls += 1 + return await super().apply_history_mutations(args) + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + calls.append(test) + return f"result:{test}" + + model = FakeModel() + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = TrackingRewriteSession() + model.set_next_output( + [get_function_tool_call("approval_tool", '{"test":"foo"}', call_id="call-1")] + ) + first = await Runner.run(agent, input="user_message", session=session) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + serialized = state.to_json() + serialized["last_processed_response"]["functions"][0]["tool_call"]["arguments"] = ( + '{"test":"different"}' + ) + + with pytest.raises(UserError, match="does not match the pending function call"): + await RunState.from_json(agent, serialized) + + assert session.rewrite_calls == 0 + assert calls == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("remove_canonical_call", [False, True]) +async def test_unmarked_execution_only_override_state_fails_before_tool_execution( + remove_canonical_call: bool, +) -> None: + calls: list[str] = [] + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + calls.append(test) + return f"result:{test}" + + model = FakeModel() + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = SimpleListSession() + model.set_next_output( + [get_function_tool_call("approval_tool", '{"test":"foo"}', call_id="call-1")] + ) + first = await Runner.run(agent, input="user_message", session=session) + state = first.to_state() + state.approve( + first.interruptions[0], + override_arguments={"test": "bar"}, + save_override_arguments=False, + ) + serialized = state.to_json() + serialized["approval_argument_override_modes"] = [] + if remove_canonical_call: + serialized["generated_items"] = [] + serialized["generated_session_item_indexes"] = [] + + with pytest.raises(UserError, match="is inconsistent with pending function calls"): + await RunState.from_json(agent, serialized) + + assert calls == [] + + +@pytest.mark.asyncio +async def test_resume_rejects_execution_only_override_without_server_managed_history() -> None: + model = FakeModel() + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + return f"result:{test}" + + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + model.add_multiple_turn_outputs( + [ + [ + get_function_tool_call( + "approval_tool", json.dumps({"test": "foo"}), call_id="call-1" + ) + ], + [get_text_message("done")], + ] + ) + + first = await Runner.run(agent, input="user_message") + assert first.interruptions + + state = first.to_state() + state.approve( + first.interruptions[0], + override_arguments={"test": "bar"}, + save_override_arguments=False, + ) + + with pytest.raises(UserError, match="save_override_arguments=False is only supported"): + await Runner.run(agent, state) + + +@pytest.mark.asyncio +async def test_resume_supports_execution_only_override_with_server_managed_session() -> None: + model = FakeModel() + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + return f"result:{test}" + + agent = Agent( + name="approval_agent", + model=model, + tools=[approval_tool], + tool_use_behavior="stop_on_first_tool", + ) + session = ServerManagedSimpleSession() + model.add_multiple_turn_outputs( + [ + [ + get_function_tool_call( + "approval_tool", json.dumps({"test": "foo"}), call_id="call-1" + ) + ], + ] + ) + + first = await Runner.run(agent, input="user_message", session=session) + assert first.interruptions - session_items = await session.get_items() - has_function_call_output = any( - isinstance(item, dict) and item.get("type") == "function_call_output" - for item in session_items + state = first.to_state() + state.approve( + first.interruptions[0], + override_arguments={"test": "bar"}, + save_override_arguments=False, ) - assert has_function_call_output + + resumed = await Runner.run(agent, state, session=session) + + assert resumed.final_output == "result:bar" + saved_items = await session.get_items() + assert cast(dict[str, Any], saved_items[1])["arguments"] == json.dumps({"test": "foo"}) + assert state.to_json()["approval_argument_override_modes"] == [] @pytest.mark.asyncio -async def test_nested_handoff_filters_reasoning_items_from_model_input(): +async def test_resume_rejects_durable_override_with_server_managed_session() -> None: model = FakeModel() - delegate = Agent( - name="delegate", - model=model, - ) - triage = Agent( - name="triage", - model=model, - handoffs=[delegate], - ) + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + return f"result:{test}" + + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = ServerManagedSimpleSession() model.add_multiple_turn_outputs( [ [ - ResponseReasoningItem( - id="reasoning_1", - type="reasoning", - summary=[Summary(text="Thinking about a handoff.", type="summary_text")], - ), - get_handoff_tool_call(delegate), + get_function_tool_call( + "approval_tool", json.dumps({"test": "foo"}), call_id="call-1" + ) ], - [get_text_message("done")], ] ) - captured_inputs: list[list[dict[str, Any]]] = [] - - def capture_model_input(data): - if isinstance(data.model_data.input, list): - captured_inputs.append( - [item for item in data.model_data.input if isinstance(item, dict)] - ) - return data.model_data + first = await Runner.run(agent, input="user_message", session=session) + assert first.interruptions - result = await Runner.run( - triage, - input="user_message", - run_config=RunConfig( - nest_handoff_history=True, - call_model_input_filter=capture_model_input, - ), - ) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) - assert result.final_output == "done" - assert len(captured_inputs) >= 2 - handoff_input = captured_inputs[1] - handoff_input_types = [ - item["type"] for item in handoff_input if isinstance(item.get("type"), str) - ] - assert "reasoning" not in handoff_input_types + with pytest.raises(UserError, match="Server-managed conversations cannot persist"): + await Runner.run(agent, state, session=session) @pytest.mark.asyncio -async def test_resume_preserves_filtered_model_input_after_handoff(): +async def test_resume_supports_execution_only_override_with_previous_response_id() -> None: model = FakeModel() @function_tool(name_override="approval_tool", needs_approval=True) - def approval_tool() -> str: - return "ok" + def approval_tool(test: str) -> str: + return f"result:{test}" - delegate = Agent( - name="delegate", + agent = Agent( + name="approval_agent", model=model, tools=[approval_tool], + tool_use_behavior="stop_on_first_tool", ) - triage = Agent( - name="triage", - model=model, - handoffs=[delegate], - tools=[get_function_tool("some_function", "result")], - ) - model.add_multiple_turn_outputs( [ [ get_function_tool_call( - "some_function", json.dumps({"a": "b"}), call_id="triage-call" + "approval_tool", json.dumps({"test": "foo"}), call_id="call-1" ) ], - [get_text_message("a_message"), get_handoff_tool_call(delegate)], - [get_function_tool_call("approval_tool", json.dumps({}), call_id="delegate-call")], - [get_text_message("done")], ] ) - model_input_call_ids: list[set[str]] = [] - model_input_output_call_ids: list[set[str]] = [] - - def capture_model_input(data): - call_ids: set[str] = set() - output_call_ids: set[str] = set() - for item in data.model_data.input: - if not isinstance(item, dict): - continue - item_type = item.get("type") - call_id = item.get("call_id") - if not isinstance(call_id, str): - continue - if item_type == "function_call": - call_ids.add(call_id) - elif item_type == "function_call_output": - output_call_ids.add(call_id) - model_input_call_ids.append(call_ids) - model_input_output_call_ids.append(output_call_ids) - return data.model_data - - run_config = RunConfig( - nest_handoff_history=True, - call_model_input_filter=capture_model_input, - ) - - first = await Runner.run(triage, input="user_message", run_config=run_config) + first = await Runner.run(agent, input="user_message", previous_response_id="resp-root") assert first.interruptions state = first.to_state() - state.approve(first.interruptions[0]) + state.approve( + first.interruptions[0], + override_arguments={"test": "bar"}, + save_override_arguments=False, + ) - resumed = await Runner.run(triage, state, run_config=run_config) + resumed = await Runner.run(agent, state) - last_call_ids = model_input_call_ids[-1] - last_output_call_ids = model_input_output_call_ids[-1] - assert "triage-call" not in last_call_ids - assert "triage-call" not in last_output_call_ids - assert "delegate-call" in last_call_ids - assert "delegate-call" in last_output_call_ids - assert resumed.final_output == "done" + assert resumed.final_output == "result:bar" @pytest.mark.asyncio -async def test_resumed_state_updates_agent_after_handoff() -> None: +async def test_resume_rejects_durable_override_for_non_rewrite_aware_session() -> None: model = FakeModel() - @function_tool(name_override="triage_tool", needs_approval=True) - def triage_tool() -> str: - return "ok" - - @function_tool(name_override="delegate_tool", needs_approval=True) - def delegate_tool() -> str: - return "ok" - - delegate = Agent( - name="delegate", - model=model, - tools=[delegate_tool], - ) - triage = Agent( - name="triage", - model=model, - handoffs=[delegate], - tools=[triage_tool], - ) + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + return f"result:{test}" + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = SimpleListSession() model.add_multiple_turn_outputs( [ - [get_function_tool_call("triage_tool", "{}", call_id="triage-1")], - [get_text_message("handoff"), get_handoff_tool_call(delegate)], - [get_function_tool_call("delegate_tool", "{}", call_id="delegate-1")], + [ + get_function_tool_call( + "approval_tool", json.dumps({"test": "foo"}), call_id="call-1" + ) + ], + [get_text_message("done")], ] ) - first = await Runner.run(triage, input="user_message") + first = await Runner.run(agent, input="user_message", session=session) assert first.interruptions state = first.to_state() - state.approve(first.interruptions[0]) + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) - second = await Runner.run(triage, state) - assert second.interruptions - assert any(item.tool_name == delegate_tool.name for item in second.interruptions), ( - "handoff should switch approvals to the delegate agent" + with pytest.raises(UserError, match="supports expected history rewrites"): + await Runner.run(agent, state, session=session) + + +@pytest.mark.asyncio +async def test_resume_rejects_method_only_history_rewrite_session() -> None: + """A method-name collision must not opt a session into expected rewrites.""" + + class MethodOnlyRewriteSession(SimpleListSession): + async def apply_history_mutations(self, _args: Any) -> bool: + return True + + model = FakeModel() + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + return f"result:{test}" + + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = MethodOnlyRewriteSession() + model.set_next_output( + [get_function_tool_call("approval_tool", '{"test":"foo"}', call_id="call-1")] ) - assert state._current_agent is delegate + + first = await Runner.run(agent, input="user_message", session=session) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + + with pytest.raises(UserError, match="supports expected history rewrites"): + await Runner.run(agent, state, session=session) + + +@pytest.mark.asyncio +async def test_resume_rejects_dynamically_forwarded_rewrite_capability() -> None: + """A transforming proxy must explicitly implement compatible rewrite semantics.""" + + class ForwardingSession: + def __init__(self) -> None: + self.session_id = "forwarding" + self.underlying_session = RewriteAwareSimpleSession(self.session_id) + + def __getattr__(self, name: str) -> Any: + return getattr(self.underlying_session, name) + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + return f"result:{test}" + + model = FakeModel() + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = ForwardingSession() + model.set_next_output( + [get_function_tool_call("approval_tool", '{"test":"foo"}', call_id="call-1")] + ) + + first = await Runner.run(agent, input="user_message", session=cast(Any, session)) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + + with pytest.raises(UserError, match="supports expected history rewrites"): + await Runner.run(agent, state, session=cast(Any, session)) + + +@pytest.mark.asyncio +async def test_resume_rejects_dynamically_forwarded_server_managed_capability() -> None: + """A transforming proxy must explicitly declare server-managed history semantics.""" + calls: list[str] = [] + + class ForwardingSession: + def __init__(self) -> None: + self.session_id = "forwarding-server" + self.underlying_session = ServerManagedSimpleSession(self.session_id) + + def __getattr__(self, name: str) -> Any: + return getattr(self.underlying_session, name) + + @function_tool(name_override="approval_tool", needs_approval=True) + def approval_tool(test: str) -> str: + calls.append(test) + return f"result:{test}" + + model = FakeModel() + agent = Agent(name="approval_agent", model=model, tools=[approval_tool]) + session = ForwardingSession() + model.set_next_output( + [get_function_tool_call("approval_tool", '{"test":"foo"}', call_id="call-1")] + ) + first = await Runner.run(agent, input="user_message", session=cast(Any, session)) + state = first.to_state() + state.approve( + first.interruptions[0], + override_arguments={"test": "bar"}, + save_override_arguments=False, + ) + + with pytest.raises(UserError, match="save_override_arguments=False is only supported"): + await Runner.run(agent, state, session=cast(Any, session)) + + assert calls == [] class Foo(TypedDict): @@ -3843,6 +4980,100 @@ async def test_save_result_to_session_does_not_increment_counter_when_nothing_sa assert session.saved_items == [] +@pytest.mark.asyncio +async def test_save_result_to_session_does_not_compact_empty_no_override_save() -> None: + client = MagicMock() + client.responses.compact = AsyncMock() + session = OpenAIResponsesCompactionSession( + session_id="test", + underlying_session=SimpleListSession(), + client=client, + should_trigger_compaction=lambda _: True, + ) + agent = Agent(name="agent", model=FakeModel()) + approval_item = ToolApprovalItem( + agent=agent, + raw_item={"type": "function_call", "call_id": "call-1", "name": "tool"}, + ) + run_state: RunState[Any] = RunState( + context=RunContextWrapper(context={}), + original_input="input", + starting_agent=agent, + max_turns=1, + ) + + saved_count = await save_result_to_session( + session, + [], + cast(list[RunItem], [approval_item]), + run_state, + response_id="resp-test", + ) + + assert saved_count == 0 + assert await session.get_items() == [] + client.responses.compact.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_save_result_to_session_compacts_empty_save_and_retains_history_mutation() -> None: + class CompactingRewriteSession(RewriteAwareSimpleSession): + def __init__(self, history: list[TResponseInputItem]) -> None: + super().__init__(history=history) + self.compaction_calls = 0 + + async def run_compaction(self, _args: Any = None) -> None: + self.compaction_calls += 1 + + original_call = cast( + TResponseInputItem, + { + "type": "function_call", + "call_id": "call-1", + "name": "tool", + "arguments": '{"value":"old"}', + }, + ) + session = CompactingRewriteSession([original_call]) + agent = Agent(name="agent", model=FakeModel()) + approval_item = ToolApprovalItem(agent=agent, raw_item=original_call) + run_state: RunState[Any] = RunState( + context=RunContextWrapper(context={}), + original_input="input", + starting_agent=agent, + max_turns=1, + ) + run_state._session_history_mutations.append( + { + "type": "replace_function_call", + "call_id": "call-1", + "expected": original_call, + "replacement": { + "type": "function_call", + "call_id": "call-1", + "name": "tool", + "arguments": '{"value":"new"}', + }, + } + ) + run_state._approval_argument_override_modes["call-1"] = "durable" + + saved_count = await save_result_to_session( + session, + [], + cast(list[RunItem], [approval_item]), + run_state, + response_id="resp-test", + ) + + assert saved_count == 0 + assert session.compaction_calls == 1 + assert [mutation["call_id"] for mutation in run_state._get_session_history_mutations()] == [ + "call-1" + ] + assert cast(dict[str, Any], (await session.get_items())[0])["arguments"] == '{"value":"new"}' + + @pytest.mark.asyncio async def test_save_result_to_session_returns_count_and_updates_state() -> None: session = SimpleListSession() diff --git a/tests/test_agent_runner_streamed.py b/tests/test_agent_runner_streamed.py index 8bbf4ec2f6..875f2753a1 100644 --- a/tests/test_agent_runner_streamed.py +++ b/tests/test_agent_runner_streamed.py @@ -37,6 +37,7 @@ OutputGuardrailTripwireTriggered, RunContextWrapper, Runner, + RunState, SQLiteSession, ToolGuardrailFunctionOutput, ToolInputGuardrailData, @@ -78,7 +79,7 @@ queue_function_call_and_text, resume_streamed_after_first_approval, ) -from .utils.simple_session import CountingSession, SimpleListSession +from .utils.simple_session import CountingSession, RewriteAwareSimpleSession, SimpleListSession def _conversation_locked_error() -> BadRequestError: @@ -2285,6 +2286,177 @@ async def test_tool() -> str: assert output_count == 1 +@pytest.mark.asyncio +async def test_streaming_resume_with_durable_override_rewrites_session_history() -> None: + async def test_tool(test: str) -> str: + return f"result:{test}" + + tool = function_tool(test_tool, name_override="test_tool", needs_approval=True) + model, agent = make_model_and_agent(name="test", tools=[tool]) + session = RewriteAwareSimpleSession() + + queue_function_call_and_text( + model, + get_function_tool_call("test_tool", json.dumps({"test": "foo"}), call_id="call-resume"), + followup=[get_text_message("done")], + ) + + first = Runner.run_streamed(agent, input="Use test_tool", session=session) + await consume_stream(first) + assert first.interruptions + + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + + resumed = Runner.run_streamed(agent, state, session=session) + await consume_stream(resumed) + + assert resumed.final_output == "done" + saved_items = await session.get_items() + assert cast(dict[str, Any], saved_items[1])["arguments"] == json.dumps({"test": "bar"}) + + +@pytest.mark.asyncio +async def test_streaming_durable_override_rewrite_fails_before_tool_execution() -> None: + calls: list[str] = [] + + async def test_tool(test: str) -> str: + calls.append(test) + return f"result:{test}" + + class FailingRewriteSession(RewriteAwareSimpleSession): + async def apply_history_mutations(self, _args: Any) -> bool: + raise RuntimeError("rewrite failed") + + tool = function_tool(test_tool, name_override="test_tool", needs_approval=True) + model = FakeModel() + agent = Agent(name="test", model=model, tools=[tool]) + session = FailingRewriteSession() + model.set_next_output( + [get_function_tool_call("test_tool", '{"test":"foo"}', call_id="call-resume")] + ) + + first = Runner.run_streamed(agent, input="Use test_tool", session=session) + await consume_stream(first) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + + resumed = Runner.run_streamed(agent, state, session=session) + with pytest.raises(RuntimeError, match="rewrite failed"): + await consume_stream(resumed) + + assert calls == [] + assert state._get_session_history_mutations() + + +@pytest.mark.asyncio +async def test_streaming_processed_override_mismatch_fails_before_session_rewrite() -> None: + calls: list[str] = [] + + class TrackingRewriteSession(RewriteAwareSimpleSession): + def __init__(self) -> None: + super().__init__() + self.rewrite_calls = 0 + + async def apply_history_mutations(self, args: Any) -> bool: + self.rewrite_calls += 1 + return await super().apply_history_mutations(args) + + async def test_tool(test: str) -> str: + calls.append(test) + return f"result:{test}" + + tool = function_tool(test_tool, name_override="test_tool", needs_approval=True) + model = FakeModel() + agent = Agent(name="test", model=model, tools=[tool]) + session = TrackingRewriteSession() + model.set_next_output( + [get_function_tool_call("test_tool", '{"test":"foo"}', call_id="call-resume")] + ) + + first = Runner.run_streamed(agent, input="Use test_tool", session=session) + await consume_stream(first) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + serialized = state.to_json() + serialized["last_processed_response"]["functions"][0]["tool_call"]["arguments"] = ( + '{"test":"different"}' + ) + + with pytest.raises(UserError, match="does not match the pending function call"): + await RunState.from_json(agent, serialized) + + assert session.rewrite_calls == 0 + assert calls == [] + + +@pytest.mark.asyncio +async def test_streaming_durable_override_missing_history_fails_before_tool_execution() -> None: + calls: list[str] = [] + + async def test_tool(test: str) -> str: + calls.append(test) + return f"result:{test}" + + tool = function_tool(test_tool, name_override="test_tool", needs_approval=True) + model = FakeModel() + agent = Agent(name="test", model=model, tools=[tool]) + session = RewriteAwareSimpleSession() + model.set_next_output( + [get_function_tool_call("test_tool", '{"test":"foo"}', call_id="call-resume")] + ) + + first = Runner.run_streamed(agent, input="Use test_tool", session=session) + await consume_stream(first) + state = first.to_state() + state.approve(first.interruptions[0], override_arguments={"test": "bar"}) + await session.clear_session() + + resumed = Runner.run_streamed(agent, state, session=session) + with pytest.raises(ValueError, match="did not match the expected function call"): + await consume_stream(resumed) + + assert calls == [] + assert state._get_session_history_mutations() + + +@pytest.mark.asyncio +async def test_streaming_resume_supports_execution_only_override_with_previous_response_id() -> ( + None +): + async def test_tool(test: str) -> str: + return f"result:{test}" + + tool = function_tool(test_tool, name_override="test_tool", needs_approval=True) + model = FakeModel() + agent = Agent( + name="test", + model=model, + tools=[tool], + tool_use_behavior="stop_on_first_tool", + ) + + model.add_multiple_turn_outputs( + [[get_function_tool_call("test_tool", json.dumps({"test": "foo"}), call_id="call-resume")]] + ) + + first = Runner.run_streamed(agent, input="Use test_tool", previous_response_id="resp-root") + await consume_stream(first) + assert first.interruptions + + state = first.to_state() + state.approve( + first.interruptions[0], + override_arguments={"test": "bar"}, + save_override_arguments=False, + ) + + resumed = Runner.run_streamed(agent, state) + await consume_stream(resumed) + + assert resumed.final_output == "result:bar" + + @pytest.mark.parametrize("mode", ["non_streamed", "streamed"]) @pytest.mark.parametrize("tripwire", [False, True], ids=["passes", "trips"]) @pytest.mark.asyncio diff --git a/tests/test_run_state.py b/tests/test_run_state.py index a69edbb242..6846b63830 100644 --- a/tests/test_run_state.py +++ b/tests/test_run_state.py @@ -304,6 +304,56 @@ def set_last_processed_response( state._last_processed_response = make_processed_response(new_items=new_items) +def build_overrideable_approval_state( + *, + conversation_id: str | None = None, + previous_response_id: str | None = None, + auto_previous_response_id: bool = False, +) -> tuple[RunState[Any, Agent[Any]], ToolApprovalItem, ResponseFunctionToolCall]: + """Build a RunState whose interruption can override function-call arguments.""" + + @function_tool(name_override="send_email") + def send_email(recipient: str) -> str: + return f"sent:{recipient}" + + context: RunContextWrapper[dict[str, Any]] = RunContextWrapper(context={}) + agent = Agent(name="OverrideAgent", tools=[send_email]) + raw_item = ResponseFunctionToolCall( + type="function_call", + id="fc_override", + call_id="call-override", + name="send_email", + arguments=json.dumps({"recipient": "alice@example.com"}), + ) + tool_call_item = ToolCallItem(agent=agent, raw_item=raw_item) + approval_item = ToolApprovalItem( + agent=agent, + raw_item=raw_item, + tool_name="send_email", + ) + state = make_state(agent, context=context, original_input="input", max_turns=2) + state._conversation_id = conversation_id + state._previous_response_id = previous_response_id + state._auto_previous_response_id = auto_previous_response_id + state._generated_items = [tool_call_item] + state._session_items = [ToolCallItem(agent=agent, raw_item=raw_item)] + state._current_step = NextStepInterruption(interruptions=[approval_item]) + state._model_responses = [ + ModelResponse( + output=[raw_item], + usage=Usage(), + response_id="resp-override", + ) + ] + state._last_processed_response = make_processed_response( + new_items=[ToolCallItem(agent=agent, raw_item=raw_item)], + functions=[ToolRunFunction(tool_call=raw_item, function_tool=send_email)], + interruptions=[approval_item], + ) + state._mark_generated_items_merged_with_last_processed() + return state, approval_item, raw_item + + class TestRunState: """Test RunState initialization, serialization, and core functionality.""" @@ -904,6 +954,573 @@ def test_approve_updates_context_approvals_correctly(self): assert state._context is not None assert state._context.is_tool_approved(tool_name="toolX", call_id="cid123") is True + @pytest.mark.parametrize( + "method_name", + [ + "get_session_history_mutations", + "has_pending_execution_only_approval_overrides", + "clear_execution_only_approval_overrides", + "clear_session_history_mutations", + ], + ) + def test_approval_override_bookkeeping_methods_are_not_public(self, method_name: str) -> None: + """Override bookkeeping should remain internal to validated resume boundaries.""" + assert not hasattr(RunState, method_name) + + def test_approve_with_override_arguments_updates_durable_replay_state(self): + """approve() should update replay history and record a session history mutation.""" + state, approval_item, _ = build_overrideable_approval_state() + + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + + assert approval_item.arguments == json.dumps({"recipient": "bob@example.com"}) + assert cast(Any, state._generated_items[0].raw_item).arguments == json.dumps( + {"recipient": "bob@example.com"} + ) + assert cast(Any, state._session_items[0].raw_item).arguments == json.dumps( + {"recipient": "bob@example.com"} + ) + assert ( + state._last_processed_response is not None + and state._last_processed_response.functions[0].tool_call.arguments + == json.dumps({"recipient": "bob@example.com"}) + ) + assert cast(Any, state._model_responses[0].output[0]).arguments == json.dumps( + {"recipient": "alice@example.com"} + ) + assert state._get_session_history_mutations() == [ + { + "type": "replace_function_call", + "call_id": "call-override", + "expected": { + "type": "function_call", + "id": "fc_override", + "call_id": "call-override", + "name": "send_email", + "arguments": json.dumps({"recipient": "alice@example.com"}), + }, + "replacement": { + "type": "function_call", + "id": "fc_override", + "call_id": "call-override", + "name": "send_email", + "arguments": json.dumps({"recipient": "bob@example.com"}), + }, + } + ] + assert state.to_json()["approval_argument_override_modes"] == [ + {"call_id": "call-override", "mode": "durable"} + ] + assert state._has_pending_execution_only_approval_overrides() is False + + def test_approve_with_execution_only_override_keeps_replay_history_unchanged(self): + """Execution-only overrides should only affect the pending execution surface.""" + state, approval_item, raw_item = build_overrideable_approval_state() + + state.approve( + approval_item, + override_arguments={"recipient": "bob@example.com"}, + save_override_arguments=False, + ) + + assert approval_item.arguments == json.dumps({"recipient": "bob@example.com"}) + assert cast(Any, state._generated_items[0].raw_item).arguments == json.dumps( + {"recipient": "alice@example.com"} + ) + assert cast(Any, state._session_items[0].raw_item).arguments == json.dumps( + {"recipient": "alice@example.com"} + ) + assert state._last_processed_response is not None + assert state._last_processed_response.functions[0].tool_call.arguments == json.dumps( + {"recipient": "bob@example.com"} + ) + assert state._model_responses[0].output[0] == raw_item + assert state._get_session_history_mutations() == [] + assert state.to_json()["approval_argument_override_modes"] == [ + {"call_id": "call-override", "mode": "execution_only"} + ] + assert state._has_pending_execution_only_approval_overrides() is True + + @pytest.mark.parametrize("save_override_arguments", [None, False]) + def test_repeated_argument_override_is_rejected_before_rebinding( + self, + save_override_arguments: bool | None, + ) -> None: + """A second argument override should leave the first approved state unchanged.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve( + approval_item, + override_arguments={"recipient": "bob@example.com"}, + save_override_arguments=save_override_arguments, + ) + serialized = state.to_json() + + with pytest.raises(UserError, match="Cannot replace an existing argument override"): + state.approve( + approval_item, + override_arguments={"recipient": "carol@example.com"}, + save_override_arguments=save_override_arguments, + ) + + assert approval_item.arguments == json.dumps({"recipient": "bob@example.com"}) + assert state.to_json() == serialized + + def test_approve_with_override_arguments_rejects_server_managed_conversation_defaults(self): + """Durable overrides should fail fast for server-managed conversations.""" + state, approval_item, _ = build_overrideable_approval_state(previous_response_id="resp-1") + + with pytest.raises( + UserError, match="save_override_arguments requires local canonical history" + ): + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + + def test_approve_with_override_arguments_validates_options(self): + """approve() should reject invalid override option combinations.""" + state, approval_item, _ = build_overrideable_approval_state() + + with pytest.raises(UserError, match="save_override_arguments can only be used"): + state.approve(approval_item, save_override_arguments=False) + with pytest.raises(UserError, match="cannot be used together with always_approve"): + state.approve( + approval_item, + always_approve=True, + override_arguments={"recipient": "bob@example.com"}, + ) + with pytest.raises(UserError, match="plain JSON object"): + state.approve(approval_item, override_arguments=cast(Any, ["not", "a", "dict"])) + + @pytest.mark.parametrize("value", [float("nan"), float("inf"), float("-inf")]) + def test_approve_with_override_arguments_rejects_non_finite_numbers(self, value: float) -> None: + """Approval overrides should contain only standards-compliant JSON values.""" + state, approval_item, _ = build_overrideable_approval_state() + + with pytest.raises(UserError, match="must contain only JSON-serializable values"): + state.approve(approval_item, override_arguments={"value": value}) + + def test_approve_with_override_arguments_redacts_serialization_errors(self) -> None: + """Serialization failures should not retain payload-bearing exception context.""" + + class LeakyList(list[Any]): + def __iter__(self): + raise ValueError("replacement-secret-sentinel") + + state, approval_item, _ = build_overrideable_approval_state() + + with pytest.raises(UserError) as exc_info: + state.approve( + approval_item, + override_arguments={"value": LeakyList(["secret"])}, + ) + + assert str(exc_info.value) == ( + "override_arguments must contain only JSON-serializable values." + ) + assert "replacement-secret-sentinel" not in str(exc_info.value) + assert exc_info.value.__cause__ is None + assert exc_info.value.__context__ is None + + @pytest.mark.parametrize( + "override_arguments", + [ + {1: "integer"}, + {1: "integer", "1": "string"}, + {"nested": [{False: "boolean"}]}, + ], + ) + def test_approve_with_override_arguments_rejects_non_string_object_keys( + self, override_arguments: Any + ) -> None: + """JSON object keys should remain identical through serialization.""" + state, approval_item, _ = build_overrideable_approval_state() + + with pytest.raises(UserError, match="must contain only JSON-serializable values"): + state.approve(approval_item, override_arguments=override_arguments) + + @pytest.mark.parametrize( + "override_arguments", + [ + {"value": (1, 2)}, + {"value": [{"nested": (1,)}]}, + ], + ) + def test_approve_with_override_arguments_rejects_non_json_containers( + self, override_arguments: Any + ) -> None: + """Approval arguments should not coerce Python-only containers.""" + state, approval_item, _ = build_overrideable_approval_state() + + with pytest.raises(UserError, match="must contain only JSON-serializable values"): + state.approve(approval_item, override_arguments=override_arguments) + + async def test_override_state_roundtrips_through_serialization(self): + """Override-specific bookkeeping should survive RunState serialization.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve( + approval_item, + override_arguments={"recipient": "bob@example.com"}, + save_override_arguments=False, + ) + + restored = await RunState.from_string(state._current_agent, state.to_string()) # type: ignore[arg-type] + + assert restored._has_pending_execution_only_approval_overrides() is True + assert restored._get_session_history_mutations() == [] + + @pytest.mark.parametrize( + "serialized_modes", + [ + None, + [123], + [{"call_id": 123, "mode": "execution_only"}], + [{"call_id": "", "mode": "execution_only"}], + [{"call_id": "call-override", "mode": "unknown"}], + [ + {"call_id": "call-override", "mode": "execution_only"}, + {"call_id": "call-override", "mode": "durable"}, + ], + ], + ) + async def test_current_deserialization_rejects_malformed_approval_override_modes( + self, serialized_modes: Any + ) -> None: + """Current-schema override modes should be exact and unambiguous.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve( + approval_item, + override_arguments={"recipient": "bob@example.com"}, + save_override_arguments=False, + ) + serialized = state.to_json() + serialized["approval_argument_override_modes"] = serialized_modes + + with pytest.raises(UserError, match="must be a list of unique valid per-call modes"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + @pytest.mark.parametrize("remove_field", [False, True]) + async def test_current_deserialization_requires_execution_only_mode_for_divergence( + self, remove_field: bool + ) -> None: + """Divergent pending execution should retain its execution-only mode.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve( + approval_item, + override_arguments={"recipient": "bob@example.com"}, + save_override_arguments=False, + ) + serialized = state.to_json() + if remove_field: + serialized.pop("approval_argument_override_modes") + else: + serialized["approval_argument_override_modes"] = [] + + with pytest.raises(UserError, match="is inconsistent with pending function calls"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + async def test_current_deserialization_requires_canonical_call_for_pending_approval( + self, + ) -> None: + """A pending function approval should retain exactly one canonical replay call.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve( + approval_item, + override_arguments={"recipient": "bob@example.com"}, + save_override_arguments=False, + ) + serialized = state.to_json() + serialized["approval_argument_override_modes"] = [] + serialized["generated_items"] = [] + serialized["generated_session_item_indexes"] = [] + + with pytest.raises(UserError, match="is inconsistent with pending function calls"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + @pytest.mark.parametrize( + "surface", + ["generated_items", "session_items", "processed_new_items", "model_responses"], + ) + async def test_current_deserialization_rejects_execution_only_history_mismatch( + self, surface: str + ) -> None: + """Execution-only historical copies should retain one original model call.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve( + approval_item, + override_arguments={"recipient": "bob@example.com"}, + save_override_arguments=False, + ) + serialized = state.to_json() + if surface == "generated_items": + raw_item = serialized["generated_items"][0]["raw_item"] + elif surface == "session_items": + raw_item = serialized["session_items"][0]["raw_item"] + elif surface == "processed_new_items": + raw_item = serialized["last_processed_response"]["new_items"][0]["raw_item"] + else: + raw_item = serialized["model_responses"][0]["output"][0] + raw_item["arguments"] = json.dumps({"recipient": "mallory@example.com"}) + + with pytest.raises(UserError, match="is inconsistent with pending function calls"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + async def test_current_deserialization_rejects_execution_only_processed_call_mismatch( + self, + ) -> None: + """Execution-only processed calls should match the approved pending execution.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve( + approval_item, + override_arguments={"recipient": "bob@example.com"}, + save_override_arguments=False, + ) + serialized = state.to_json() + serialized["last_processed_response"]["functions"][0]["tool_call"]["arguments"] = ( + json.dumps({"recipient": "mallory@example.com"}) + ) + + with pytest.raises(UserError, match="is inconsistent with pending function calls"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + async def test_durable_override_state_roundtrips_through_serialization(self): + """Durable override mutations should survive RunState serialization.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + + restored = await RunState.from_string(state._current_agent, state.to_string()) # type: ignore[arg-type] + + assert restored._get_session_history_mutations() == [ + { + "type": "replace_function_call", + "call_id": "call-override", + "expected": { + "type": "function_call", + "id": "fc_override", + "call_id": "call-override", + "name": "send_email", + "arguments": json.dumps({"recipient": "alice@example.com"}), + }, + "replacement": { + "type": "function_call", + "id": "fc_override", + "call_id": "call-override", + "name": "send_email", + "arguments": json.dumps({"recipient": "bob@example.com"}), + }, + } + ] + + @pytest.mark.parametrize( + "missing_field", ["approval_argument_override_modes", "session_history_mutations"] + ) + async def test_current_deserialization_requires_complete_durable_override_provenance( + self, missing_field: str + ) -> None: + """Durable overrides should retain both mode provenance and the rewrite payload.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + serialized = state.to_json() + serialized[missing_field] = [] + + with pytest.raises(UserError, match="is inconsistent with pending function calls"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + async def test_current_deserialization_rejects_durable_override_marked_execution_only( + self, + ) -> None: + """A durable mutation should not be accepted under execution-only provenance.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + serialized = state.to_json() + serialized["approval_argument_override_modes"][0]["mode"] = "execution_only" + + with pytest.raises(UserError, match="is inconsistent with pending function calls"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + async def test_current_deserialization_rejects_mutation_execution_mismatch(self) -> None: + """A durable mutation should exactly match the pending function call.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + serialized = state.to_json() + serialized["session_history_mutations"][0]["replacement"]["arguments"] = json.dumps( + {"recipient": "mallory@example.com"} + ) + + with pytest.raises(UserError, match="does not match the pending function call"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + async def test_current_deserialization_rejects_mutation_expected_audit_mismatch(self) -> None: + """A durable mutation should remain bound to the original model call.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + serialized = state.to_json() + serialized["session_history_mutations"][0]["expected"]["arguments"] = json.dumps( + {"recipient": "mallory@example.com"} + ) + + with pytest.raises(UserError, match="does not match the pending function call"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + async def test_current_deserialization_rejects_mutation_item_identity_mismatch(self) -> None: + """A durable replacement should preserve the complete canonical call payload.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + serialized = state.to_json() + serialized["session_history_mutations"][0]["replacement"]["id"] = "fc-tampered" + + with pytest.raises(UserError, match="does not match the pending function call"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + async def test_current_deserialization_rejects_mutation_processed_call_mismatch(self) -> None: + """A durable replacement must match the function call that resume will execute.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + serialized = state.to_json() + serialized["last_processed_response"]["functions"][0]["tool_call"]["arguments"] = ( + json.dumps({"recipient": "mallory@example.com"}) + ) + + with pytest.raises(UserError, match="does not match the pending function call"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + @pytest.mark.parametrize( + "surface", + ["session_items", "processed_new_items"], + ) + async def test_current_deserialization_rejects_durable_replay_surface_mismatch( + self, surface: str + ) -> None: + """A durable replacement should match every stored replay copy.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + serialized = state.to_json() + if surface == "session_items": + raw_item = serialized["session_items"][0]["raw_item"] + else: + raw_item = serialized["last_processed_response"]["new_items"][0]["raw_item"] + raw_item["arguments"] = json.dumps({"recipient": "mallory@example.com"}) + + with pytest.raises(UserError, match="does not match the pending function call"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + async def test_current_deserialization_preserves_raw_model_response_provenance(self) -> None: + """Raw model responses should retain the model-requested arguments for auditability.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + + restored = await RunState.from_json(state._current_agent, state.to_json()) # type: ignore[arg-type] + + assert cast(Any, restored._model_responses[0].output[0]).arguments == json.dumps( + {"recipient": "alice@example.com"} + ) + assert restored._get_session_history_mutations()[0]["replacement"][ + "arguments" + ] == json.dumps({"recipient": "bob@example.com"}) + + async def test_current_deserialization_requires_durable_raw_model_audit_call(self) -> None: + """A durable override should retain its matching raw model audit call.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + serialized = json.loads(state.to_string()) + serialized["model_responses"][0]["output"] = [] + serialized["last_model_response"] = serialized["model_responses"][-1] + + with pytest.raises(UserError, match="session_history_mutations does not match"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + async def test_current_deserialization_rejects_durable_raw_model_identity_mismatch( + self, + ) -> None: + """A durable raw audit call should preserve the overridden call's identity.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + serialized = json.loads(state.to_string()) + serialized["model_responses"][0]["output"][0]["id"] = "fc-tampered" + serialized["last_model_response"] = serialized["model_responses"][-1] + + with pytest.raises(UserError, match="session_history_mutations does not match"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + async def test_current_deserialization_rejects_last_model_response_mismatch(self) -> None: + """The duplicate last model response should match the serialized response list.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + serialized = json.loads(state.to_string()) + serialized["last_model_response"]["output"][0]["arguments"] = json.dumps( + {"recipient": "mallory@example.com"} + ) + + with pytest.raises(UserError, match="is inconsistent with pending function calls"): + await RunState.from_json(state._current_agent, serialized) # type: ignore[arg-type] + + @pytest.mark.parametrize( + "replacement", + [ + None, + {}, + { + "type": "function_call", + "call_id": "different-call", + "name": "send_email", + "arguments": "{}", + }, + { + "type": "function_call", + "call_id": "call-override", + "arguments": "{}", + }, + { + "type": "function_call", + "call_id": "call-override", + "name": "send_email", + "arguments": {"recipient": "bob@example.com"}, + }, + ], + ) + async def test_deserialization_rejects_invalid_current_session_history_mutations( + self, replacement: Any + ) -> None: + """Malformed current-schema mutations should fail closed.""" + state, _, _ = build_overrideable_approval_state() + serialized = state.to_json() + serialized["session_history_mutations"] = [ + { + "type": "replace_function_call", + "call_id": "call-override", + "expected": serialized["model_responses"][-1]["output"][0], + "replacement": replacement, + } + ] + + with pytest.raises(UserError, match="contains an invalid function-call replacement"): + await RunState.from_string( # type: ignore[arg-type] + state._current_agent, + json.dumps(serialized), + ) + + @pytest.mark.parametrize( + "schema_version", + [ + version + for version in sorted(SUPPORTED_SCHEMA_VERSIONS) + if tuple(int(part) for part in version.split(".")) < (1, 16) + ], + ) + async def test_legacy_deserialization_ignores_approval_override_fields( + self, schema_version: str + ) -> None: + """Schemas predating override bookkeeping should ignore its fields.""" + state, approval_item, _ = build_overrideable_approval_state() + state.approve(approval_item, override_arguments={"recipient": "bob@example.com"}) + serialized = state.to_json() + serialized["$schemaVersion"] = schema_version + + restored = await RunState.from_string( # type: ignore[arg-type] + state._current_agent, + json.dumps(serialized), + ) + + assert restored._get_session_history_mutations() == [] + assert restored.to_json()["approval_argument_override_modes"] == [] + def test_returns_undefined_when_approval_status_is_unknown(self): """Test that isToolApproved returns None for unknown tools.""" context: RunContextWrapper[dict[str, str]] = RunContextWrapper(context={}) @@ -7198,6 +7815,7 @@ def test_supported_schema_versions_match_released_boundary(self): "1.12", "1.13", "1.14", + "1.15", CURRENT_SCHEMA_VERSION, } ) diff --git a/tests/test_run_state_compatibility_corpus.py b/tests/test_run_state_compatibility_corpus.py index 19232245cb..e667ba9f92 100644 --- a/tests/test_run_state_compatibility_corpus.py +++ b/tests/test_run_state_compatibility_corpus.py @@ -165,14 +165,66 @@ 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 all(entry["commit"] for entry in SOURCES["versions"].values()) + assert all( + entry.get("commit") or entry.get("lineage_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"} } assert {entry["provenance"] for entry in SOURCES["features"]} == { "historical_writer", "canonical_compatibility", + "current_writer", + } + + +def test_current_approval_override_fixtures_cover_both_persistence_modes() -> None: + entries = { + entry["feature"]: entry for entry in SOURCES["features"] if entry["version"] == "1.16" } + assert set(entries) == { + "approval_argument_overrides", + "durable_approval_argument_overrides", + } + + execution_only = json.loads( + (FIXTURE_ROOT / entries["approval_argument_overrides"]["fixture"]).read_text( + encoding="utf-8" + ) + ) + durable = json.loads( + (FIXTURE_ROOT / entries["durable_approval_argument_overrides"]["fixture"]).read_text( + encoding="utf-8" + ) + ) + + assert execution_only["approval_argument_override_modes"] == [ + {"call_id": "approval-override-1", "mode": "execution_only"} + ] + assert execution_only["session_history_mutations"] == [] + assert durable["approval_argument_override_modes"] == [ + {"call_id": "approval-override-1", "mode": "durable"} + ] + assert durable["session_history_mutations"] == [ + { + "call_id": "approval-override-1", + "expected": { + "arguments": json.dumps({"recipient": "alice@example.com"}), + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call", + }, + "replacement": { + "arguments": json.dumps({"recipient": "bob@example.com"}), + "call_id": "approval-override-1", + "id": "fc-override", + "name": "send_email", + "type": "function_call", + }, + "type": "replace_function_call", + } + ] @pytest.mark.parametrize( diff --git a/tests/utils/simple_session.py b/tests/utils/simple_session.py index 94bcc97e9e..137fdaa873 100644 --- a/tests/utils/simple_session.py +++ b/tests/utils/simple_session.py @@ -1,9 +1,13 @@ from __future__ import annotations -from typing import cast +from typing import Literal, cast from agents.items import TResponseInputItem -from agents.memory.session import Session +from agents.memory.session import ( + Session, + SessionHistoryRewriteArgs, + apply_session_history_mutations, +) from agents.memory.session_settings import SessionSettings @@ -80,3 +84,20 @@ async def add_items(self, items: list[TResponseInputItem]) -> None: else: sanitized.append(item) await super().add_items(sanitized) + + +class RewriteAwareSimpleSession(SimpleListSession): + """In-memory test session that supports persisted-history rewrites.""" + + supports_expected_history_mutations: Literal[True] = True + + async def apply_history_mutations(self, args: SessionHistoryRewriteArgs) -> bool: + self._items = apply_session_history_mutations(self._items, args.get("mutations", [])) + self.saved_items = self._items + return True + + +class ServerManagedSimpleSession(SimpleListSession): + """In-memory test session that advertises server-managed history semantics.""" + + _server_managed_conversation_session = True