From af69e3739eeba07cae38ff1539ef2f229ace4249 Mon Sep 17 00:00:00 2001 From: Justin Beckwith Date: Wed, 7 Oct 2026 07:26:18 -0700 Subject: [PATCH 1/8] fix: preserve completed tool outputs when a sibling fails --- src/agents/run.py | 27 ++ src/agents/run_internal/run_loop.py | 50 ++++ src/agents/run_internal/turn_resolution.py | 77 ++++-- tests/test_tool_batch_failure_history.py | 274 +++++++++++++++++++++ 4 files changed, 411 insertions(+), 17 deletions(-) create mode 100644 tests/test_tool_batch_failure_history.py diff --git a/src/agents/run.py b/src/agents/run.py index d0785268f3..453e4e5020 100644 --- a/src/agents/run.py +++ b/src/agents/run.py @@ -135,6 +135,7 @@ NextStepInterruption, NextStepRunAgain, ProcessedResponse, + SingleStepResult, ) from .run_internal.session_persistence import ( _session_get_items, @@ -1709,6 +1710,7 @@ async def _save_max_turns_handler_output( ) if current_turn_span is not None: current_turn_span.start(mark_as_current=True) + partial_tool_results: list[SingleStepResult] = [] try: if current_turn <= 1: try: @@ -1758,6 +1760,7 @@ async def _save_max_turns_handler_output( on_response_accepted=_commit_pending_server_response, on_response_hooks_started=_mark_response_hooks_started, run_state=run_state, + on_tool_execution_error=partial_tool_results.append, ) ) @@ -1832,7 +1835,31 @@ async def _save_max_turns_handler_output( on_response_accepted=_commit_pending_server_response, on_response_hooks_started=_mark_response_hooks_started, run_state=run_state, + on_tool_execution_error=partial_tool_results.append, ) + except Exception: + input_accepted = len(_attempt_input_guardrail_results()) >= len( + all_input_guardrails + ) and not input_guardrails_triggered(_attempt_input_guardrail_results()) + if partial_tool_results and input_accepted: + partial_result = partial_tool_results[0] + generated_items.extend(partial_result.new_step_items) + session_items.extend(partial_result.new_step_items) + model_responses.append(partial_result.model_response) + try: + await save_turn_items_if_needed( + session=session, + run_state=run_state, + session_persistence_enabled=session_persistence_enabled, + input_guardrail_results=input_guardrail_results, + items=partial_result.new_step_items, + response_id=partial_result.model_response.response_id, + store=store_setting, + wrapper=context_wrapper, + ) + except Exception: + logger.warning("Failed to save completed tools after a tool error") + raise finally: if current_turn_span is not None: attach_usage_to_span( diff --git a/src/agents/run_internal/run_loop.py b/src/agents/run_internal/run_loop.py index 688a3a9004..732c65d91b 100644 --- a/src/agents/run_internal/run_loop.py +++ b/src/agents/run_internal/run_loop.py @@ -1789,6 +1789,7 @@ def _record_max_turns_handler_output( ) if current_turn_span is not None: current_turn_span.start(mark_as_current=True) + partial_tool_results: list[SingleStepResult] = [] try: if ( session is not None @@ -1818,7 +1819,52 @@ def _record_max_turns_handler_output( on_response_accepted=_commit_pending_server_response, on_response_hooks_started=_mark_response_hooks_started, run_state=run_state, + on_tool_execution_error=partial_tool_results.append, ) + except Exception as tool_error: + input_task = streamed_result._input_guardrails_task + input_accepted = input_task is None or ( + input_task.done() + and not input_task.cancelled() + and input_task.exception() is None + ) + if ( + partial_tool_results + and input_accepted + and not any( + result.output.tripwire_triggered + for result in streamed_result.input_guardrail_results + ) + ): + partial_result = partial_tool_results[0] + streamed_result._model_input_items.extend(partial_result.new_step_items) + streamed_result.new_items.extend(partial_result.new_step_items) + streamed_result.raw_responses.append(partial_result.model_response) + if run_state is not None: + run_state._current_step = NextStepRunAgain() + run_state._generated_items = list(streamed_result._model_input_items) + run_state._session_items = list(streamed_result.new_items) + run_state._model_responses = list(streamed_result.raw_responses) + stream_step_items_to_queue( + [ + item + for item in partial_result.new_step_items + if item.type == "tool_call_output_item" + ], + streamed_result._event_queue, + ) + _mark_error_to_drain_stream_events(tool_error) + try: + await _save_stream_items_with_count( + partial_result.new_step_items, + partial_result.model_response.response_id, + current_agent.model_settings.resolve( + run_config.model_settings + ).store, + ) + except Exception: + logger.warning("Failed to save completed tools after a tool error") + raise finally: if current_turn_span is not None: attach_usage_to_span( @@ -2114,6 +2160,7 @@ async def run_single_turn_streamed( on_response_accepted: Callable[[ModelResponse, ProcessedResponse | None], bool] | None = None, on_response_hooks_started: Callable[[], None] | None = None, run_state: RunState[Any] | None = None, + on_tool_execution_error: Callable[[SingleStepResult], None] | None = None, ) -> SingleStepResult: """Run a single streamed turn and emit events as results arrive.""" public_agent = bindings.public_agent @@ -2438,6 +2485,7 @@ async def check_input_guardrails_before_side_effects() -> None: after_invocation_validation=after_invocation_validation, before_side_effects=check_input_guardrails_before_side_effects, run_state=run_state, + on_tool_execution_error=on_tool_execution_error, ) items_to_filter = session_items_for_turn(single_step_result) @@ -2473,6 +2521,7 @@ async def run_single_turn( on_response_accepted: Callable[[ModelResponse, ProcessedResponse | None], bool] | None = None, on_response_hooks_started: Callable[[], None] | None = None, run_state: RunState[Any] | None = None, + on_tool_execution_error: Callable[[SingleStepResult], None] | None = None, ) -> SingleStepResult: """Run a single non-streaming turn of the agent loop.""" public_agent = bindings.public_agent @@ -2585,6 +2634,7 @@ async def after_invocation_validation( server_manages_conversation=server_conversation_tracker is not None, after_invocation_validation=after_invocation_validation, run_state=run_state, + on_tool_execution_error=on_tool_execution_error, ) diff --git a/src/agents/run_internal/turn_resolution.py b/src/agents/run_internal/turn_resolution.py index e143de5c21..d0b8332f64 100644 --- a/src/agents/run_internal/turn_resolution.py +++ b/src/agents/run_internal/turn_resolution.py @@ -793,6 +793,25 @@ async def check_for_final_output_from_tools( raise UserError(f"Invalid tool_use_behavior: {agent.tool_use_behavior}") +def _completed_tool_step_items(model_items: list[RunItem], outputs: list[RunItem]) -> list[RunItem]: + """Keep accepted call/output pairs in model order and their preceding reasoning.""" + outputs_by_call_id = {extract_tool_call_id(item.raw_item): item for item in outputs} + retained: list[RunItem] = [] + ordered_outputs: list[RunItem] = [] + reasoning: list[RunItem] = [] + for item in model_items: + if isinstance(item, ReasoningItem): + reasoning.append(item) + elif isinstance(item, ToolCallItem): + output = outputs_by_call_id.get(extract_tool_call_id(item.raw_item)) + if output is not None: + retained.extend(reasoning) + reasoning.clear() + retained.append(item) + ordered_outputs.append(output) + return [*retained, *ordered_outputs] + + async def execute_tools_and_side_effects( *, bindings: AgentBindings[TContext], @@ -808,6 +827,7 @@ async def execute_tools_and_side_effects( server_manages_conversation: bool = False, precomputed_skipped_raw_item_ids: set[int] | None = None, run_state: RunState[Any] | None = None, + on_tool_execution_error: Callable[[SingleStepResult], None] | None = None, ) -> SingleStepResult: """Run one turn of the loop, coordinating tools, approvals, guardrails, and handoffs.""" public_agent = bindings.public_agent @@ -849,7 +869,10 @@ async def execute_tools_and_side_effects( skipped_raw_item_ids=skipped_raw_item_ids, ) + completed_outputs: list[RunItem] = [] + def _commit_accepted_response_tool_output(item: RunItem) -> None: + completed_outputs.append(item) if run_state is None or not isinstance(run_state._current_step, NextStepInterruption): return if not run_state._current_step.response_accepted: @@ -858,23 +881,41 @@ def _commit_accepted_response_tool_output(item: RunItem) -> None: if item not in target: target.append(item) - ( - function_results, - tool_input_guardrail_results, - tool_output_guardrail_results, - computer_results, - custom_tool_results, - shell_results, - apply_patch_results, - local_shell_results, - ) = await _execute_tool_plan( - plan=plan, - bindings=bindings, - hooks=hooks, - context_wrapper=context_wrapper, - run_config=run_config, - tool_output_committer=_commit_accepted_response_tool_output, - ) + try: + ( + function_results, + tool_input_guardrail_results, + tool_output_guardrail_results, + computer_results, + custom_tool_results, + shell_results, + apply_patch_results, + local_shell_results, + ) = await _execute_tool_plan( + plan=plan, + bindings=bindings, + hooks=hooks, + context_wrapper=context_wrapper, + run_config=run_config, + tool_output_committer=_commit_accepted_response_tool_output, + ) + except Exception: + # Accepted server responses already have their own resumable checkpoint. + if on_tool_execution_error is not None and not server_manages_conversation: + retained_items = _completed_tool_step_items(new_step_items, completed_outputs) + if retained_items: + on_tool_execution_error( + SingleStepResult( + original_input=original_input, + model_response=new_response, + pre_step_items=pre_step_items, + new_step_items=retained_items, + next_step=NextStepRunAgain(), + tool_input_guardrail_results=[], + tool_output_guardrail_results=[], + ) + ) + raise new_step_items.extend( _build_tool_result_items( function_results=function_results, @@ -3682,6 +3723,7 @@ async def get_single_step_result_from_response( | None = None, before_side_effects: Callable[[], Awaitable[None]] | None = None, run_state: RunState[Any] | None = None, + on_tool_execution_error: Callable[[SingleStepResult], None] | None = None, ) -> SingleStepResult: item_agent = bindings.public_agent try: @@ -3744,4 +3786,5 @@ async def get_single_step_result_from_response( server_manages_conversation=server_manages_conversation, precomputed_skipped_raw_item_ids=skipped_raw_item_ids, run_state=run_state, + on_tool_execution_error=on_tool_execution_error, ) diff --git a/tests/test_tool_batch_failure_history.py b/tests/test_tool_batch_failure_history.py new file mode 100644 index 0000000000..331e0ccd20 --- /dev/null +++ b/tests/test_tool_batch_failure_history.py @@ -0,0 +1,274 @@ +"""Completed side effects survive a sibling failure within the same model turn.""" + +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest +from openai.types.responses import ResponseReasoningItem + +from agents import ( + Agent, + AgentsException, + GuardrailFunctionOutput, + InputGuardrailTripwireTriggered, + RunHooks, + Runner, + RunState, + SQLiteSession, + ToolGuardrailFunctionOutput, + ToolInputGuardrailTripwireTriggered, + ToolOutputGuardrailTripwireTriggered, + UserError, + input_guardrail, + tool_input_guardrail, + tool_output_guardrail, +) +from agents.decorators import tool +from agents.testing import ScriptedModel, assistant_message, function_call + + +def _shape(items: list[Any]) -> list[str]: + return [ + f"{item['type']}:{item['call_id']}" + if "call_id" in item + else item.get("type", item.get("role", "?")) + for item in items + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("failure", ["input", "output", "handler", "hook"]) +async def test_completed_sibling_survives_tool_batch_failure(failure: str, streaming: bool): + finished = asyncio.Event() + effects: list[str] = [] + + @tool + async def create_ticket() -> str: + effects.append("ticket") + return "ticket T-1" + + @tool_input_guardrail + async def reject_input(data): + await finished.wait() + return ToolGuardrailFunctionOutput.raise_exception(output_info="blocked") + + @tool_output_guardrail + async def reject_output(data): + return ToolGuardrailFunctionOutput.raise_exception(output_info="blocked") + + @tool( + failure_error_function=None, + tool_input_guardrails=[reject_input] if failure == "input" else [], + tool_output_guardrails=[reject_output] if failure == "output" else [], + ) + async def send_email() -> str: + await finished.wait() + if failure == "handler": + raise ValueError("synthetic tool failure") + return "email sent" + + class Hooks(RunHooks): + async def on_tool_end(self, context, agent, tool, result): + if tool.name == "create_ticket": + finished.set() + elif failure == "hook": + raise ValueError("synthetic hook failure") + + model = ScriptedModel( + [ + [ + ResponseReasoningItem(id="rs_before", type="reasoning", summary=[]), + function_call("send_email", {}, call_id="email"), + function_call("create_ticket", {}, call_id="ticket"), + ResponseReasoningItem(id="rs_after", type="reasoning", summary=[]), + ], + [assistant_message("done")], + ] + ) + agent = Agent(name="support", model=model, tools=[send_email, create_ticket]) + session = SQLiteSession("test") + expected_error = { + "input": ToolInputGuardrailTripwireTriggered, + "output": ToolOutputGuardrailTripwireTriggered, + "handler": UserError, + "hook": UserError, + }[failure] + result = None + output_events = [] + try: + with pytest.raises(expected_error) as caught: + if streaming: + result = Runner.run_streamed(agent, "go", session=session, hooks=Hooks()) + async for event in result.stream_events(): + if ( + event.type == "run_item_stream_event" + and event.item.type == "tool_call_output_item" + ): + output_events.append(event.item.to_input_item()) + else: + await Runner.run(agent, "go", session=session, hooks=Hooks()) + assert effects == ["ticket"] + expected = ["reasoning"] + if failure == "hook": + # An end hook runs after the output has passed its tool guardrails. + expected += ["function_call:email"] + expected += ["function_call:ticket"] + if failure == "hook": + expected += ["function_call_output:email"] + expected += ["function_call_output:ticket"] + assert isinstance(caught.value, AgentsException) + assert caught.value.run_data is not None + assert _shape([i.to_input_item() for i in caught.value.run_data.new_items]) == expected + history = await session.get_items() + assert _shape(history) == ["user", *expected] + assert history[1]["id"] == "rs_before" + assert history[-1]["output"] == "ticket T-1" + if result is not None: + assert _shape(output_events) == [ + item for item in expected if item.startswith("function_call_output:") + ] + assert _shape(result.to_input_list()) == ["user", *expected] + state = await RunState.from_json(agent, result.to_state().to_json()) + replay_model = ScriptedModel([[assistant_message("resumed")]]) + agent.model = replay_model + await Runner.run(agent, state) + assert replay_model.last_call is not None + assert _shape(replay_model.last_call.input) == ["user", *expected] + assert effects == ["ticket"] + agent.model = model + await Runner.run(agent, "finish", session=session) + assert model.last_call is not None + assert _shape(model.last_call.input) == ["user", *expected, "user"] + assert effects == ["ticket"] + finally: + session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +async def test_input_guardrail_rejection_does_not_publish_partial_tools(streaming: bool): + completed = asyncio.Event() + blocked = asyncio.Event() + result = None + + @tool + async def completed_tool() -> str: + return "speculative output" + + @tool(failure_error_function=None) + async def failed_tool() -> str: + await blocked.wait() + if streaming: + assert result is not None + assert result._input_guardrails_task is not None + await result._input_guardrails_task + raise ValueError("synthetic failure") + + class Hooks(RunHooks): + async def on_tool_end(self, context, agent, tool, result): + completed.set() + + @input_guardrail + async def reject(context, agent, input): + await completed.wait() + blocked.set() + return GuardrailFunctionOutput(output_info="blocked", tripwire_triggered=True) + + agent = Agent( + name="support", + model=ScriptedModel( + [ + [ + function_call("completed_tool", {}, call_id="done"), + function_call("failed_tool", {}, call_id="failed"), + ] + ] + ), + tools=[completed_tool, failed_tool], + input_guardrails=[reject], + ) + session = SQLiteSession("test") + try: + with pytest.raises((InputGuardrailTripwireTriggered, UserError)): + if streaming: + result = Runner.run_streamed(agent, "go", session=session, hooks=Hooks()) + # Let the failure settle before the consumer reacts to the tripwire. + assert result.run_loop_task is not None + with pytest.raises(UserError): + await result.run_loop_task + async for _ in result.stream_events(): + pass + else: + await Runner.run(agent, "go", session=session, hooks=Hooks()) + assert _shape(await session.get_items()) == ["user"] + if result is not None: + assert result.new_items == [] + assert result._model_input_items == [] + assert _shape(result.to_input_list()) == ["user"] + state = result.to_state() + assert state._generated_items == [] + assert state._session_items == [] + finally: + session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +async def test_session_save_failure_preserves_primary_tool_error(streaming: bool): + class FailingSession(SQLiteSession): + async def add_items(self, items): + if any(item.get("type") == "function_call_output" for item in items): + raise RuntimeError("synthetic session failure") + await super().add_items(items) + + finished = asyncio.Event() + + @tool + async def completed_tool() -> str: + return "completed" + + @tool_input_guardrail + async def reject(data): + await finished.wait() + return ToolGuardrailFunctionOutput.raise_exception(output_info="blocked") + + @tool(tool_input_guardrails=[reject]) + async def blocked_tool() -> str: + raise AssertionError("blocked tool must not execute") + + class Hooks(RunHooks): + async def on_tool_end(self, context, agent, tool, result): + finished.set() + + agent = Agent( + name="support", + model=ScriptedModel( + [ + [ + function_call("completed_tool", {}, call_id="done"), + function_call("blocked_tool", {}, call_id="blocked"), + ] + ] + ), + tools=[completed_tool, blocked_tool], + ) + session = FailingSession("test") + try: + with pytest.raises(ToolInputGuardrailTripwireTriggered) as caught: + if streaming: + result = Runner.run_streamed(agent, "go", session=session, hooks=Hooks()) + async for _ in result.stream_events(): + pass + else: + await Runner.run(agent, "go", session=session, hooks=Hooks()) + assert caught.value.run_data is not None + assert _shape([item.to_input_item() for item in caught.value.run_data.new_items]) == [ + "function_call:done", + "function_call_output:done", + ] + assert _shape(await session.get_items()) == ["user"] + finally: + session.close() From 4e334fbc96c070daabe871c93bc6de81b700a15e Mon Sep 17 00:00:00 2001 From: Justin Beckwith Date: Wed, 7 Oct 2026 08:28:14 -0700 Subject: [PATCH 2/8] fix: preserve resumed state and diagnostics after tool batch failures --- src/agents/run.py | 41 ++++-- src/agents/run_internal/run_loop.py | 35 ++++- src/agents/run_internal/turn_resolution.py | 4 +- tests/test_tool_batch_failure_history.py | 163 ++++++++++++++++++++- 4 files changed, 225 insertions(+), 18 deletions(-) diff --git a/src/agents/run.py b/src/agents/run.py index de980f8207..eb2e2e9588 100644 --- a/src/agents/run.py +++ b/src/agents/run.py @@ -1896,17 +1896,38 @@ async def _save_max_turns_handler_output( generated_items.extend(partial_result.new_step_items) session_items.extend(partial_result.new_step_items) model_responses.append(partial_result.model_response) - try: - await save_turn_items_if_needed( - session=session, - run_state=run_state, - session_persistence_enabled=session_persistence_enabled, - input_guardrail_results=input_guardrail_results, - items=partial_result.new_step_items, - response_id=partial_result.model_response.response_id, - store=store_setting, - wrapper=context_wrapper, + tool_input_guardrail_results.extend( + partial_result.tool_input_guardrail_results + ) + tool_output_guardrail_results.extend( + partial_result.tool_output_guardrail_results + ) + if run_state is not None: + _synchronize_accepted_run_state( + run_state, + generated_items=generated_items, + session_items=session_items, + model_responses=model_responses, + tool_input_guardrail_results=tool_input_guardrail_results, + tool_output_guardrail_results=tool_output_guardrail_results, + current_turn=current_turn, + ) + run_state._current_step = NextStepRunAgain() + run_state.set_tool_use_tracker_snapshot( + _tool_use_tracker_snapshot() ) + try: + if session_persistence_enabled: + await save_result_to_session( + session, + [], + partial_result.new_step_items, + run_state, + response_id=partial_result.model_response.response_id, + store=store_setting, + wrapper=context_wrapper, + resumed_write_state=run_state, + ) except Exception: logger.warning("Failed to save completed tools after a tool error") raise diff --git a/src/agents/run_internal/run_loop.py b/src/agents/run_internal/run_loop.py index 7c53d7a991..a73bba9043 100644 --- a/src/agents/run_internal/run_loop.py +++ b/src/agents/run_internal/run_loop.py @@ -2013,11 +2013,38 @@ def _record_max_turns_handler_output( streamed_result._model_input_items.extend(partial_result.new_step_items) streamed_result.new_items.extend(partial_result.new_step_items) streamed_result.raw_responses.append(partial_result.model_response) + _accumulate_tool_guardrail_results( + streamed_result, + partial_result, + accepted_input_results=accepted_tool_input_guardrail_results, + accepted_output_results=accepted_tool_output_guardrail_results, + ) + streamed_result._tool_use_tracker_snapshot = serialize_tool_use_tracker( + tool_use_tracker, + starting_agent=( + run_state._starting_agent + if run_state is not None and run_state._starting_agent is not None + else starting_agent + ), + ) if run_state is not None: + _synchronize_accepted_run_state( + run_state, + generated_items=streamed_result._model_input_items, + session_items=streamed_result.new_items, + model_responses=streamed_result.raw_responses, + tool_input_guardrail_results=( + streamed_result.tool_input_guardrail_results + ), + tool_output_guardrail_results=( + streamed_result.tool_output_guardrail_results + ), + current_turn=current_turn, + ) run_state._current_step = NextStepRunAgain() - run_state._generated_items = list(streamed_result._model_input_items) - run_state._session_items = list(streamed_result.new_items) - run_state._model_responses = list(streamed_result.raw_responses) + run_state.set_tool_use_tracker_snapshot( + streamed_result._tool_use_tracker_snapshot + ) stream_step_items_to_queue( [ item @@ -2028,7 +2055,7 @@ def _record_max_turns_handler_output( ) _mark_error_to_drain_stream_events(tool_error) try: - await _save_stream_items_with_count( + await _save_stream_items_without_count( partial_result.new_step_items, partial_result.model_response.response_id, current_agent.model_settings.resolve( diff --git a/src/agents/run_internal/turn_resolution.py b/src/agents/run_internal/turn_resolution.py index 26ad9236bf..6672ff2576 100644 --- a/src/agents/run_internal/turn_resolution.py +++ b/src/agents/run_internal/turn_resolution.py @@ -937,8 +937,8 @@ def _commit_accepted_response_tool_output(item: RunItem) -> None: pre_step_items=pre_step_items, new_step_items=retained_items, next_step=NextStepRunAgain(), - tool_input_guardrail_results=[], - tool_output_guardrail_results=[], + tool_input_guardrail_results=tool_input_guardrail_results, + tool_output_guardrail_results=tool_output_guardrail_results, ) ) raise diff --git a/tests/test_tool_batch_failure_history.py b/tests/test_tool_batch_failure_history.py index 331e0ccd20..bc473220c9 100644 --- a/tests/test_tool_batch_failure_history.py +++ b/tests/test_tool_batch_failure_history.py @@ -13,6 +13,7 @@ AgentsException, GuardrailFunctionOutput, InputGuardrailTripwireTriggered, + ModelSettings, RunHooks, Runner, RunState, @@ -45,7 +46,15 @@ async def test_completed_sibling_survives_tool_batch_failure(failure: str, strea finished = asyncio.Event() effects: list[str] = [] - @tool + @tool_input_guardrail + async def allow_input(data): + return ToolGuardrailFunctionOutput.allow(output_info="input accepted") + + @tool_output_guardrail + async def allow_output(data): + return ToolGuardrailFunctionOutput.allow(output_info="output accepted") + + @tool(tool_input_guardrails=[allow_input], tool_output_guardrails=[allow_output]) async def create_ticket() -> str: effects.append("ticket") return "ticket T-1" @@ -88,7 +97,12 @@ async def on_tool_end(self, context, agent, tool, result): [assistant_message("done")], ] ) - agent = Agent(name="support", model=model, tools=[send_email, create_ticket]) + agent = Agent( + name="support", + model=model, + tools=[send_email, create_ticket], + model_settings=ModelSettings(tool_choice="required"), + ) session = SQLiteSession("test") expected_error = { "input": ToolInputGuardrailTripwireTriggered, @@ -121,6 +135,14 @@ async def on_tool_end(self, context, agent, tool, result): expected += ["function_call_output:ticket"] assert isinstance(caught.value, AgentsException) assert caught.value.run_data is not None + assert any( + decision.output.output_info == "input accepted" + for decision in caught.value.run_data.tool_input_guardrail_results + ) + assert any( + decision.output.output_info == "output accepted" + for decision in caught.value.run_data.tool_output_guardrail_results + ) assert _shape([i.to_input_item() for i in caught.value.run_data.new_items]) == expected history = await session.get_items() assert _shape(history) == ["user", *expected] @@ -132,10 +154,15 @@ async def on_tool_end(self, context, agent, tool, result): ] assert _shape(result.to_input_list()) == ["user", *expected] state = await RunState.from_json(agent, result.to_state().to_json()) + assert any( + decision.output.output_info == "output accepted" + for decision in state._tool_output_guardrail_results + ) replay_model = ScriptedModel([[assistant_message("resumed")]]) agent.model = replay_model await Runner.run(agent, state) assert replay_model.last_call is not None + assert replay_model.last_call.model_settings.tool_choice is None assert _shape(replay_model.last_call.input) == ["user", *expected] assert effects == ["ticket"] agent.model = model @@ -272,3 +299,135 @@ async def on_tool_end(self, context, agent, tool, result): assert _shape(await session.get_items()) == ["user"] finally: session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("write_failure", [None, "before_append", "after_append"]) +async def test_resumed_failure_keeps_accepted_history_and_reconciles_session( + streaming: bool, write_failure: str | None +): + class FailingSession(SQLiteSession): + fail_once = write_failure + + async def add_items(self, items): + failure = self.fail_once + if failure and any(item.get("call_id") == "done" for item in items): + self.fail_once = None + if failure == "after_append": + await super().add_items(items) + raise RuntimeError("synthetic append failure") + await super().add_items(items) + + finished = asyncio.Event() + effects: list[str] = [] + admissions: list[str] = [] + + @input_guardrail + async def admit(context, agent, input): + admissions.append("admitted") + return GuardrailFunctionOutput(output_info="safe", tripwire_triggered=False) + + @tool(needs_approval=True) + async def approved_tool() -> str: + effects.append("approved") + return "approved" + + @tool_output_guardrail + async def accept_output(data): + return ToolGuardrailFunctionOutput.allow(output_info="accepted side effect") + + @tool(tool_output_guardrails=[accept_output]) + async def completed_tool() -> str: + effects.append("done") + return "completed" + + @tool(failure_error_function=None) + async def failed_tool() -> str: + await finished.wait() + raise ValueError("synthetic sibling failure") + + class Hooks(RunHooks): + async def on_tool_end(self, context, agent, tool, result): + if tool.name == "completed_tool": + finished.set() + + model = ScriptedModel( + [ + [function_call("approved_tool", {}, call_id="approved")], + [ + function_call("completed_tool", {}, call_id="done"), + function_call("failed_tool", {}, call_id="failed"), + ], + [assistant_message("resumed")], + ] + ) + agent = Agent( + name="support", + model=model, + tools=[approved_tool, completed_tool, failed_tool], + input_guardrails=[admit], + model_settings=ModelSettings(tool_choice="required"), + ) + session = FailingSession("resume") + try: + interrupted = await Runner.run(agent, "go", session=session) + state = interrupted.to_state() + state.approve(interrupted.interruptions[0]) + stream_result = None + with pytest.raises(UserError, match="synthetic sibling failure") as caught: + if streaming: + stream_result = Runner.run_streamed(agent, state, session=session, hooks=Hooks()) + async for _ in stream_result.stream_events(): + pass + else: + await Runner.run(agent, state, session=session, hooks=Hooks()) + expected = [ + "function_call:approved", + "function_call_output:approved", + "function_call:done", + "function_call_output:done", + ] + assert admissions == ["admitted"] + assert caught.value.run_data is not None + assert ( + _shape( + [ + item.to_input_item() + for item in caught.value.run_data.new_items + if item.type != "tool_approval_item" + ] + ) + == expected + ) + assert ( + _shape( + [ + item.to_input_item() + for item in state._generated_items + if item.type != "tool_approval_item" + ] + ) + == expected + ) + assert len(state._model_responses) == 2 + assert state._current_turn == 2 + assert [r.output.output_info for r in state._tool_output_guardrail_results] == [ + "accepted side effect" + ] + assert (state._pending_session_write is not None) == (write_failure is not None) + if stream_result is not None: + state = stream_result.to_state() + assert (state._pending_session_write is not None) == (write_failure is not None) + state = await RunState.from_json(agent, state.to_json()) + resumed = await Runner.run(agent, state, session=session) + assert resumed.final_output == "resumed" + assert model.last_call is not None + assert _shape(model.last_call.input) == ["user", *expected] + assert model.last_call.model_settings.tool_choice is None + assert _shape(await session.get_items()) == ["user", *expected, "message"] + assert effects == ["approved", "done"] + assert admissions == ["admitted"] + assert state._pending_session_write is None + finally: + session.close() From 8493f588e198d22aafd3dfb05eaa352047f98ee1 Mon Sep 17 00:00:00 2001 From: Justin Beckwith Date: Wed, 7 Oct 2026 09:20:01 -0700 Subject: [PATCH 3/8] fix: preserve completed tool history across failure lifecycles --- src/agents/exceptions.py | 11 + src/agents/items.py | 3 + src/agents/result.py | 10 +- src/agents/run.py | 4 +- src/agents/run_internal/guardrails.py | 1 + src/agents/run_internal/run_loop.py | 19 +- src/agents/run_internal/tool_actions.py | 201 ++++++---- src/agents/run_internal/tool_execution.py | 46 ++- src/agents/run_internal/tool_planning.py | 20 +- src/agents/run_internal/turn_resolution.py | 98 +++-- src/agents/util/_asyncio_tasks.py | 18 +- tests/test_asyncio_tasks.py | 8 +- tests/test_native_tool_failure_history.py | 416 +++++++++++++++++++++ tests/test_tool_batch_failure_history.py | 99 +++++ tests/test_tool_failure_boundaries.py | 146 ++++++++ 15 files changed, 970 insertions(+), 130 deletions(-) create mode 100644 tests/test_native_tool_failure_history.py create mode 100644 tests/test_tool_failure_boundaries.py diff --git a/src/agents/exceptions.py b/src/agents/exceptions.py index c490edc77d..799dcdaa9d 100644 --- a/src/agents/exceptions.py +++ b/src/agents/exceptions.py @@ -29,6 +29,7 @@ from .util._pretty_print import pretty_print_run_error_details +_TOOL_LOCAL_CANCELLATION_ATTR = "_agents_tool_local_cancellation" _DRAIN_STREAM_EVENTS_ATTR = "_agents_drain_queued_stream_events" _DATA_REDACTED_ATTR = "_agents_data_redacted" _DATA_REDACTED_ERROR_MESSAGE = "Error details are redacted." @@ -41,6 +42,16 @@ class _RedactedExceptionCancellationError(asyncio.CancelledError, Exception): """Payload-free cancellation that remains catchable as an Exception.""" +def _mark_tool_local_cancellation(error: asyncio.CancelledError) -> None: + setattr(error, _TOOL_LOCAL_CANCELLATION_ATTR, True) + + +def _is_tool_local_cancellation(error: BaseException) -> bool: + return isinstance(error, asyncio.CancelledError) and bool( + getattr(error, _TOOL_LOCAL_CANCELLATION_ATTR, False) + ) + + def _mark_error_to_drain_stream_events(error: BaseException) -> None: setattr(error, _DRAIN_STREAM_EVENTS_ATTR, True) diff --git a/src/agents/items.py b/src/agents/items.py index 49f734b1d0..7a350e62ad 100644 --- a/src/agents/items.py +++ b/src/agents/items.py @@ -451,6 +451,9 @@ class ToolCallOutputItem(RunItemBase[Any]): replayed as input. """ + _custom_data_pending: bool = field(default=False, init=False, repr=False, compare=False) + """Live finalization state; excluded from fresh partial history, not serialized to RunState.""" + @property def call_id(self) -> str | None: """Return the call identifier from the raw item, if available.""" diff --git a/src/agents/result.py b/src/agents/result.py index 92fcd96259..b2c0f93573 100644 --- a/src/agents/result.py +++ b/src/agents/result.py @@ -703,6 +703,8 @@ class RunResultStreaming(RunResultBase): _triggered_input_guardrail_result: InputGuardrailResult | None = field(default=None, repr=False) _output_guardrails_task: asyncio.Task[Any] | None = field(default=None, repr=False) _stored_exception: BaseException | None = field(default=None, repr=False) + _tool_error_selected: bool = field(default=False, init=False, repr=False) + """A selected tool error owns failure reporting while its input verdict settles.""" _cancel_mode: Literal["none", "immediate", "after_turn"] = field(default="none", repr=False) _last_processed_response: ProcessedResponse | None = field(default=None, repr=False) """The last processed model response. This is needed for resuming from interruptions.""" @@ -1174,7 +1176,7 @@ def _check_errors(self): # Fetch all the completed guardrail results from the queue and raise if needed while not self._input_guardrail_queue.empty(): guardrail_result = self._input_guardrail_queue.get_nowait() - if guardrail_result.output.tripwire_triggered: + if guardrail_result.output.tripwire_triggered and not self._tool_error_selected: tripwire_exc = InputGuardrailTripwireTriggered(guardrail_result) tripwire_exc.run_data = self._create_error_details() self._stored_exception = tripwire_exc @@ -1192,7 +1194,11 @@ def _check_errors(self): run_impl_exc.run_data = self._create_error_details() self._stored_exception = run_impl_exc - if self._input_guardrails_task and self._input_guardrails_task.done(): + if ( + not self._tool_error_selected + and self._input_guardrails_task + and self._input_guardrails_task.done() + ): if not self._input_guardrails_task.cancelled(): in_guard_exc = self._input_guardrails_task.exception() if isinstance(in_guard_exc, Exception): diff --git a/src/agents/run.py b/src/agents/run.py index eb2e2e9588..d819380afa 100644 --- a/src/agents/run.py +++ b/src/agents/run.py @@ -1887,7 +1887,9 @@ async def _save_max_turns_handler_output( run_state=run_state, on_tool_execution_error=partial_tool_results.append, ) - except Exception: + except (Exception, asyncio.CancelledError): + if not partial_tool_results: + raise input_accepted = len(_attempt_input_guardrail_results()) >= len( all_input_guardrails ) and not input_guardrails_triggered(_attempt_input_guardrail_results()) diff --git a/src/agents/run_internal/guardrails.py b/src/agents/run_internal/guardrails.py index d1a1c24d2c..a87d486983 100644 --- a/src/agents/run_internal/guardrails.py +++ b/src/agents/run_internal/guardrails.py @@ -105,6 +105,7 @@ async def run_input_guardrails_with_queue( isinstance(error, Exception) and asyncio.current_task() is streamed_result._input_guardrails_task and not streamed_result.is_complete + and not streamed_result._tool_error_selected ): if streamed_result.run_loop_task and not streamed_result.run_loop_task.done(): streamed_result.run_loop_task.cancel() diff --git a/src/agents/run_internal/run_loop.py b/src/agents/run_internal/run_loop.py index a73bba9043..97e536fcf5 100644 --- a/src/agents/run_internal/run_loop.py +++ b/src/agents/run_internal/run_loop.py @@ -1994,8 +1994,20 @@ def _record_max_turns_handler_output( run_state=run_state, on_tool_execution_error=partial_tool_results.append, ) - except Exception as tool_error: + except (Exception, asyncio.CancelledError) as tool_error: + if not partial_tool_results: + raise input_task = streamed_result._input_guardrails_task + if partial_tool_results and input_task is not None and not input_task.done(): + # The finalizer already waits for this verdict. Settle it before + # publishing, without letting a late verdict replace the tool error. + streamed_result._tool_error_selected = True + try: + await input_guardrail_tripwire_triggered_for_stream( + streamed_result, ignore_cancelled=True + ) + except Exception: + pass input_accepted = input_task is None or ( input_task.done() and not input_task.cancelled() @@ -2704,6 +2716,10 @@ async def after_invocation_validation( async def check_input_guardrails_before_side_effects() -> None: await raise_if_input_guardrail_tripwire_known() + def on_tool_error_selected() -> None: + # The category has selected its failure, before draining native finalization. + streamed_result._tool_error_selected = True + single_step_result = await get_single_step_result_from_response( bindings=bindings, original_input=streamed_result.input, @@ -2722,6 +2738,7 @@ async def check_input_guardrails_before_side_effects() -> None: before_side_effects=check_input_guardrails_before_side_effects, run_state=run_state, on_tool_execution_error=on_tool_execution_error, + on_tool_error_selected=on_tool_error_selected, ) items_to_filter = session_items_for_turn(single_step_result) diff --git a/src/agents/run_internal/tool_actions.py b/src/agents/run_internal/tool_actions.py index 771e7ac974..3b938bd27b 100644 --- a/src/agents/run_internal/tool_actions.py +++ b/src/agents/run_internal/tool_actions.py @@ -5,6 +5,7 @@ from __future__ import annotations +import asyncio import copy import dataclasses import inspect @@ -57,6 +58,7 @@ render_shell_outputs, resolve_approval_rejection_message, resolve_approval_status, + run_native_tool_post_invoke, serialize_shell_output, truncate_shell_outputs, with_tool_function_span, @@ -114,6 +116,7 @@ async def execute( config: RunConfig, acknowledged_safety_checks: list[ComputerCallOutputAcknowledgedSafetyCheck] | None = None, tool_output_committer: Callable[[RunItem], None] | None = None, + sibling_category_failure: asyncio.Event | None = None, ) -> RunItem: """Run a computer action, capturing a screenshot and notifying hooks.""" trace_tool_name = get_tool_trace_name_for_tool(action.computer_tool) or cls.TRACE_TOOL_NAME @@ -177,28 +180,38 @@ async def _run_action(span: Any | None) -> RunItem: output=image_url, raw_item=raw_item, ) + + # Preserve the pre-metadata recovery checkpoint for approved/server-owned calls. + output_item._custom_data_pending = True if tool_output_committer is not None: tool_output_committer(output_item) - custom_data = await maybe_extract_custom_data( - action.computer_tool.custom_data_extractor, - ComputerToolCustomDataContext( - run_context=context_wrapper, - tool=action.computer_tool, - tool_call=action.tool_call, - output=image_url, - raw_item=copy.deepcopy(raw_item), - ), - ) - output_item.custom_data = custom_data - await gather_with_cancel( - hooks.on_tool_end(context_wrapper, agent, action.computer_tool, output), - ( - agent_hooks.on_tool_end(context_wrapper, agent, action.computer_tool, output) - if agent_hooks is not None - else _coro.noop_coroutine() - ), - ) + async def finalize_output() -> None: + custom_data = await maybe_extract_custom_data( + action.computer_tool.custom_data_extractor, + ComputerToolCustomDataContext( + run_context=context_wrapper, + tool=action.computer_tool, + tool_call=action.tool_call, + output=image_url, + raw_item=copy.deepcopy(raw_item), + ), + ) + output_item.custom_data = custom_data + output_item._custom_data_pending = False + + await gather_with_cancel( + hooks.on_tool_end(context_wrapper, agent, action.computer_tool, output), + ( + agent_hooks.on_tool_end( + context_wrapper, agent, action.computer_tool, output + ) + if agent_hooks is not None + else _coro.noop_coroutine() + ), + ) + + await run_native_tool_post_invoke(finalize_output(), sibling_category_failure) if span is not None and config.trace_include_sensitive_data: span.span_data.output = image_url @@ -400,6 +413,7 @@ async def execute( context_wrapper: RunContextWrapper[Any], config: RunConfig, tool_output_committer: Callable[[RunItem], None] | None = None, + sibling_category_failure: asyncio.Event | None = None, ) -> RunItem: """Run a local shell tool call and wrap the result as a ToolCallOutputItem.""" agent_hooks = agent.hooks @@ -433,17 +447,22 @@ async def execute( output=result, raw_item=raw_payload, ) - if tool_output_committer is not None: - tool_output_committer(output_item) - await gather_with_cancel( - hooks.on_tool_end(context_wrapper, agent, call.local_shell_tool, result), - ( - agent_hooks.on_tool_end(context_wrapper, agent, call.local_shell_tool, result) - if agent_hooks is not None - else _coro.noop_coroutine() - ), - ) + async def finalize_output() -> None: + if tool_output_committer is not None: + tool_output_committer(output_item) + + await gather_with_cancel( + hooks.on_tool_end(context_wrapper, agent, call.local_shell_tool, result), + ( + agent_hooks.on_tool_end(context_wrapper, agent, call.local_shell_tool, result) + if agent_hooks is not None + else _coro.noop_coroutine() + ), + ) + + await run_native_tool_post_invoke(finalize_output(), sibling_category_failure) + return output_item @@ -460,6 +479,7 @@ async def execute( context_wrapper: RunContextWrapper[Any], config: RunConfig, tool_output_committer: Callable[[RunItem], None] | None = None, + sibling_category_failure: asyncio.Event | None = None, ) -> RunItem: """Run a shell tool call and return a normalized ToolCallOutputItem.""" shell_call = coerce_shell_call(call.tool_call) @@ -639,17 +659,23 @@ async def _run_call(span: Any | None) -> RunItem: output=output_text, raw_item=raw_item, ) - if tool_output_committer is not None: - tool_output_committer(output_item) - await gather_with_cancel( - hooks.on_tool_end(context_wrapper, agent, call.shell_tool, output_text), - ( - agent_hooks.on_tool_end(context_wrapper, agent, call.shell_tool, output_text) - if agent_hooks is not None - else _coro.noop_coroutine() - ), - ) + async def finalize_output() -> None: + if tool_output_committer is not None: + tool_output_committer(output_item) + + await gather_with_cancel( + hooks.on_tool_end(context_wrapper, agent, call.shell_tool, output_text), + ( + agent_hooks.on_tool_end( + context_wrapper, agent, call.shell_tool, output_text + ) + if agent_hooks is not None + else _coro.noop_coroutine() + ), + ) + + await run_native_tool_post_invoke(finalize_output(), sibling_category_failure) if span is not None and config.trace_include_sensitive_data: span.span_data.output = output_text @@ -676,6 +702,7 @@ async def execute( context_wrapper: RunContextWrapper[Any], config: RunConfig, tool_output_committer: Callable[[RunItem], None] | None = None, + sibling_category_failure: asyncio.Event | None = None, ) -> RunItem: custom_tool: CustomTool = call.custom_tool agent_hooks = agent.hooks @@ -801,28 +828,36 @@ async def _run_call(span: Any | None) -> RunItem: output_text, raw_item=raw_item, ) + + # Preserve the pre-metadata recovery checkpoint for approved/server-owned calls. + output_item._custom_data_pending = True if tool_output_committer is not None: tool_output_committer(output_item) - custom_data = await maybe_extract_custom_data( - custom_tool.custom_data_extractor, - CustomToolCustomDataContext( - tool_context=tool_context, - tool=custom_tool, - input=tool_input, - output=output_text, - raw_item=copy.deepcopy(raw_item), - ), - ) - output_item.custom_data = custom_data - await gather_with_cancel( - hooks.on_tool_end(tool_context, agent, custom_tool, output_text), - ( - agent_hooks.on_tool_end(tool_context, agent, custom_tool, output_text) - if agent_hooks is not None - else _coro.noop_coroutine() - ), - ) + async def finalize_output() -> None: + custom_data = await maybe_extract_custom_data( + custom_tool.custom_data_extractor, + CustomToolCustomDataContext( + tool_context=tool_context, + tool=custom_tool, + input=tool_input, + output=output_text, + raw_item=copy.deepcopy(raw_item), + ), + ) + output_item.custom_data = custom_data + output_item._custom_data_pending = False + + await gather_with_cancel( + hooks.on_tool_end(tool_context, agent, custom_tool, output_text), + ( + agent_hooks.on_tool_end(tool_context, agent, custom_tool, output_text) + if agent_hooks is not None + else _coro.noop_coroutine() + ), + ) + + await run_native_tool_post_invoke(finalize_output(), sibling_category_failure) if span is not None and config.trace_include_sensitive_data: span.span_data.output = output_text @@ -885,6 +920,7 @@ async def execute( context_wrapper: RunContextWrapper[Any], config: RunConfig, tool_output_committer: Callable[[RunItem], None] | None = None, + sibling_category_failure: asyncio.Event | None = None, ) -> RunItem: """Run an apply_patch call and serialize the editor result for the model.""" apply_patch_tool: ApplyPatchTool = call.apply_patch_tool @@ -1037,30 +1073,39 @@ async def _run_call(span: Any | None) -> RunItem: output=output_text, raw_item=raw_item, ) + + # Preserve the pre-metadata recovery checkpoint for approved/server-owned calls. + output_item._custom_data_pending = True if tool_output_committer is not None: tool_output_committer(output_item) - custom_data = await maybe_extract_custom_data( - apply_patch_tool.custom_data_extractor, - ApplyPatchToolCustomDataContext( - run_context=context_wrapper, - tool=apply_patch_tool, - operations=operations, - output=output_text, - status=status, - raw_item=copy.deepcopy(raw_item), - ), - ) - output_item.custom_data = custom_data + async def finalize_output() -> None: + custom_data = await maybe_extract_custom_data( + apply_patch_tool.custom_data_extractor, + ApplyPatchToolCustomDataContext( + run_context=context_wrapper, + tool=apply_patch_tool, + operations=operations, + output=output_text, + status=status, + raw_item=copy.deepcopy(raw_item), + ), + ) + output_item.custom_data = custom_data + output_item._custom_data_pending = False + + await gather_with_cancel( + hooks.on_tool_end(context_wrapper, agent, apply_patch_tool, output_text), + ( + agent_hooks.on_tool_end( + context_wrapper, agent, apply_patch_tool, output_text + ) + if agent_hooks is not None + else _coro.noop_coroutine() + ), + ) - await gather_with_cancel( - hooks.on_tool_end(context_wrapper, agent, apply_patch_tool, output_text), - ( - agent_hooks.on_tool_end(context_wrapper, agent, apply_patch_tool, output_text) - if agent_hooks is not None - else _coro.noop_coroutine() - ), - ) + await run_native_tool_post_invoke(finalize_output(), sibling_category_failure) if span is not None and config.trace_include_sensitive_data: span.span_data.output = output_text diff --git a/src/agents/run_internal/tool_execution.py b/src/agents/run_internal/tool_execution.py index 9ba6bce821..597a1fca3f 100644 --- a/src/agents/run_internal/tool_execution.py +++ b/src/agents/run_internal/tool_execution.py @@ -54,6 +54,7 @@ ToolInputGuardrailTripwireTriggered, ToolOutputGuardrailTripwireTriggered, UserError, + _mark_tool_local_cancellation, ) from ..items import ( ItemHelpers, @@ -100,7 +101,7 @@ from ..tracing import Span, SpanError, function_span, get_current_trace from ..util import _coro, _error_tracing from ..util._approvals import evaluate_function_tool_approval -from ..util._asyncio_tasks import gather_with_cancel +from ..util._asyncio_tasks import _consume_future_exception, gather_with_cancel from ..util._custom_data import maybe_extract_custom_data, merge_custom_data from ..util._tool_errors import get_trace_tool_error from ..util._types import MaybeAwaitable @@ -1557,6 +1558,7 @@ def __init__( config: RunConfig, isolate_parallel_failures: bool | None, sibling_category_failure: asyncio.Event | None, + on_tool_error_selected: Callable[[], None] | None, tool_output_committer: Callable[[RunItem], None] | None, tool_input_guardrail_results: list[ToolInputGuardrailResult] | None, tool_output_guardrail_results: list[ToolOutputGuardrailResult] | None, @@ -1571,6 +1573,7 @@ def __init__( len(tool_runs) > 1 if isolate_parallel_failures is None else isolate_parallel_failures ) self.sibling_category_failure = sibling_category_failure + self.on_tool_error_selected = on_tool_error_selected self.tool_output_committer = tool_output_committer self.tool_input_guardrail_results = ( tool_input_guardrail_results if tool_input_guardrail_results is not None else [] @@ -1628,6 +1631,7 @@ async def execute( await self._drain_pending_tasks(pending_tool_runs) except asyncio.CancelledError as exc: if self.propagating_failure is exc: + _mark_tool_local_cancellation(exc) raise if self.sibling_category_failure is not None and self.sibling_category_failure.is_set(): await self._drain_pending_tasks_for_sibling_category_failure() @@ -1687,6 +1691,8 @@ async def _raise_failure_after_draining_siblings( self, failure: _FunctionToolFailure, ) -> None: + if self.on_tool_error_selected is not None: + self.on_tool_error_selected() cancellable_tasks, post_invoke_tasks = self._partition_pending_tasks() self.teardown_cancelled_tasks.update(cancellable_tasks) _cancel_function_tool_tasks(cancellable_tasks) @@ -2342,6 +2348,7 @@ async def execute_function_tool_calls( config: RunConfig, isolate_parallel_failures: bool | None = None, sibling_category_failure: asyncio.Event | None = None, + on_tool_error_selected: Callable[[], None] | None = None, tool_output_committer: Callable[[RunItem], None] | None = None, tool_input_guardrail_results: list[ToolInputGuardrailResult] | None = None, tool_output_guardrail_results: list[ToolOutputGuardrailResult] | None = None, @@ -2357,12 +2364,39 @@ async def execute_function_tool_calls( config=config, isolate_parallel_failures=isolate_parallel_failures, sibling_category_failure=sibling_category_failure, + on_tool_error_selected=on_tool_error_selected, tool_output_committer=tool_output_committer, tool_input_guardrail_results=tool_input_guardrail_results, tool_output_guardrail_results=tool_output_guardrail_results, ).execute() +async def run_native_tool_post_invoke( + post_invoke: Awaitable[None], + sibling_category_failure: asyncio.Event | None, +) -> None: + """Settle native output finalization on sibling failure without delaying parent cancellation.""" + if sibling_category_failure is None: + await post_invoke + return + + task = asyncio.ensure_future(post_invoke) + task.add_done_callback(_consume_future_exception) + try: + await asyncio.shield(task) + except asyncio.CancelledError: + if sibling_category_failure.is_set(): + try: + await asyncio.wait((task,), timeout=_FUNCTION_TOOL_POST_INVOKE_WAIT_SECONDS) + except BaseException: + task.cancel() + raise + else: + task.cancel() + # Do not start the next native invocation after the category was cancelled. + raise + + async def execute_custom_tool_calls( *, public_agent: Agent[Any], @@ -2371,6 +2405,7 @@ async def execute_custom_tool_calls( hooks: RunHooks[Any], config: RunConfig, tool_output_committer: Callable[[RunItem], None] | None = None, + sibling_category_failure: asyncio.Event | None = None, ) -> list[RunItem]: """Run Responses custom tool calls serially and wrap outputs.""" from .tool_actions import CustomToolAction @@ -2385,6 +2420,7 @@ async def execute_custom_tool_calls( context_wrapper=context_wrapper, config=config, tool_output_committer=tool_output_committer, + sibling_category_failure=sibling_category_failure, ) ) return results @@ -2398,6 +2434,7 @@ async def execute_local_shell_calls( hooks: RunHooks[Any], config: RunConfig, tool_output_committer: Callable[[RunItem], None] | None = None, + sibling_category_failure: asyncio.Event | None = None, ) -> list[RunItem]: """Run local shell tool calls serially and wrap outputs.""" from .tool_actions import LocalShellAction @@ -2412,6 +2449,7 @@ async def execute_local_shell_calls( context_wrapper=context_wrapper, config=config, tool_output_committer=tool_output_committer, + sibling_category_failure=sibling_category_failure, ) ) return results @@ -2425,6 +2463,7 @@ async def execute_shell_calls( hooks: RunHooks[Any], config: RunConfig, tool_output_committer: Callable[[RunItem], None] | None = None, + sibling_category_failure: asyncio.Event | None = None, ) -> list[RunItem]: """Run shell tool calls serially and wrap outputs.""" from .tool_actions import ShellAction @@ -2439,6 +2478,7 @@ async def execute_shell_calls( context_wrapper=context_wrapper, config=config, tool_output_committer=tool_output_committer, + sibling_category_failure=sibling_category_failure, ) ) return results @@ -2452,6 +2492,7 @@ async def execute_apply_patch_calls( hooks: RunHooks[Any], config: RunConfig, tool_output_committer: Callable[[RunItem], None] | None = None, + sibling_category_failure: asyncio.Event | None = None, ) -> list[RunItem]: """Run apply_patch tool calls serially and normalize outputs.""" from .tool_actions import ApplyPatchAction @@ -2466,6 +2507,7 @@ async def execute_apply_patch_calls( context_wrapper=context_wrapper, config=config, tool_output_committer=tool_output_committer, + sibling_category_failure=sibling_category_failure, ) ) return results @@ -2479,6 +2521,7 @@ async def execute_computer_actions( context_wrapper: RunContextWrapper[Any], config: RunConfig, tool_output_committer: Callable[[RunItem], None] | None = None, + sibling_category_failure: asyncio.Event | None = None, ) -> list[RunItem]: """Run computer actions serially and emit screenshot outputs.""" from .tool_actions import ComputerAction @@ -2530,6 +2573,7 @@ async def execute_computer_actions( config=config, acknowledged_safety_checks=acknowledged, tool_output_committer=tool_output_committer, + sibling_category_failure=sibling_category_failure, ) ) diff --git a/src/agents/run_internal/tool_planning.py b/src/agents/run_internal/tool_planning.py index 887cb32cbc..4f247c92fc 100644 --- a/src/agents/run_internal/tool_planning.py +++ b/src/agents/run_internal/tool_planning.py @@ -23,7 +23,7 @@ tool_output_identity, ) from ..agent import Agent -from ..exceptions import ModelBehaviorError, UserError +from ..exceptions import ModelBehaviorError, UserError, _mark_tool_local_cancellation from ..items import ( HandoffCallItem, HandoffOutputItem, @@ -957,6 +957,7 @@ async def _execute_tool_plan( run_config, parallel: bool = True, tool_output_committer: Callable[[RunItem], None] | None = None, + on_tool_error_selected: Callable[[], None] | None = None, tool_input_guardrail_results: list[ToolInputGuardrailResult] | None = None, tool_output_guardrail_results: list[ToolOutputGuardrailResult] | None = None, ) -> tuple[ @@ -983,6 +984,14 @@ async def _execute_tool_plan( ) if parallel: sibling_category_failure = asyncio.Event() + + def on_category_failure(error: BaseException) -> None: + sibling_category_failure.set() + if on_tool_error_selected is not None: + on_tool_error_selected() + if isinstance(error, asyncio.CancelledError): + _mark_tool_local_cancellation(error) + ( (function_results, tool_input_guardrail_results, tool_output_guardrail_results), computer_results, @@ -998,6 +1007,7 @@ async def _execute_tool_plan( context_wrapper=context_wrapper, config=run_config, isolate_parallel_failures=isolate_function_tool_failures, + on_tool_error_selected=on_tool_error_selected, sibling_category_failure=sibling_category_failure, tool_output_committer=tool_output_committer, tool_input_guardrail_results=tool_input_guardrail_results, @@ -1010,6 +1020,7 @@ async def _execute_tool_plan( context_wrapper=context_wrapper, config=run_config, tool_output_committer=tool_output_committer, + sibling_category_failure=sibling_category_failure, ), execute_custom_tool_calls( public_agent=public_agent, @@ -1018,6 +1029,7 @@ async def _execute_tool_plan( context_wrapper=context_wrapper, config=run_config, tool_output_committer=tool_output_committer, + sibling_category_failure=sibling_category_failure, ), execute_shell_calls( public_agent=public_agent, @@ -1026,6 +1038,7 @@ async def _execute_tool_plan( context_wrapper=context_wrapper, config=run_config, tool_output_committer=tool_output_committer, + sibling_category_failure=sibling_category_failure, ), execute_apply_patch_calls( public_agent=public_agent, @@ -1034,6 +1047,7 @@ async def _execute_tool_plan( context_wrapper=context_wrapper, config=run_config, tool_output_committer=tool_output_committer, + sibling_category_failure=sibling_category_failure, ), execute_local_shell_calls( public_agent=public_agent, @@ -1042,8 +1056,9 @@ async def _execute_tool_plan( context_wrapper=context_wrapper, config=run_config, tool_output_committer=tool_output_committer, + sibling_category_failure=sibling_category_failure, ), - on_child_failure=sibling_category_failure.set, + on_child_failure=on_category_failure, ) else: ( @@ -1057,6 +1072,7 @@ async def _execute_tool_plan( context_wrapper=context_wrapper, config=run_config, isolate_parallel_failures=isolate_function_tool_failures, + on_tool_error_selected=on_tool_error_selected, tool_output_committer=tool_output_committer, tool_input_guardrail_results=tool_input_guardrail_results, tool_output_guardrail_results=tool_output_guardrail_results, diff --git a/src/agents/run_internal/turn_resolution.py b/src/agents/run_internal/turn_resolution.py index 6672ff2576..d588854d15 100644 --- a/src/agents/run_internal/turn_resolution.py +++ b/src/agents/run_internal/turn_resolution.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import inspect from collections.abc import Awaitable, Callable, Container, Mapping, Sequence from copy import deepcopy @@ -70,6 +71,7 @@ UserError, _detach_data_redacted_error_traceback, _is_error_data_redacted, + _is_tool_local_cancellation, _mark_error_data_redacted, ) from ..handoffs import Handoff, HandoffInputData, HandoffInputFilter, nest_handoff_history @@ -803,7 +805,11 @@ async def check_for_final_output_from_tools( def _completed_tool_step_items(model_items: list[RunItem], outputs: list[RunItem]) -> list[RunItem]: """Keep accepted call/output pairs in model order and their preceding reasoning.""" - outputs_by_call_id = {extract_tool_call_id(item.raw_item): item for item in outputs} + outputs_by_call_id = { + extract_tool_call_id(item.raw_item): item + for item in outputs + if not (isinstance(item, ToolCallOutputItem) and item._custom_data_pending) + } retained: list[RunItem] = [] ordered_outputs: list[RunItem] = [] reasoning: list[RunItem] = [] @@ -812,11 +818,25 @@ def _completed_tool_step_items(model_items: list[RunItem], outputs: list[RunItem reasoning.append(item) elif isinstance(item, ToolCallItem): output = outputs_by_call_id.get(extract_tool_call_id(item.raw_item)) - if output is not None: + provider_completed = ( + isinstance( + item.raw_item, + ( + ResponseFileSearchToolCall, + ResponseFunctionWebSearch, + ResponseCodeInterpreterToolCall, + ImageGenerationCall, + McpCall, + ), + ) + and item.raw_item.status == "completed" + ) + if output is not None or provider_completed: retained.extend(reasoning) reasoning.clear() retained.append(item) - ordered_outputs.append(output) + if output is not None: + ordered_outputs.append(output) return [*retained, *ordered_outputs] @@ -836,6 +856,7 @@ async def execute_tools_and_side_effects( precomputed_skipped_raw_item_ids: set[int] | None = None, run_state: RunState[Any] | None = None, on_tool_execution_error: Callable[[SingleStepResult], None] | None = None, + on_tool_error_selected: Callable[[], None] | None = None, ) -> SingleStepResult: """Run one turn of the loop, coordinating tools, approvals, guardrails, and handoffs.""" public_agent = bindings.public_agent @@ -905,6 +926,23 @@ def _commit_accepted_response_tool_output(item: RunItem) -> None: *tool_output_guardrail_results, ] + def _publish_completed_tools() -> None: + # Accepted server responses already have their own resumable checkpoint. + if on_tool_execution_error is not None and not server_manages_conversation: + retained_items = _completed_tool_step_items(new_step_items, completed_outputs) + if retained_items: + on_tool_execution_error( + SingleStepResult( + original_input=original_input, + model_response=new_response, + pre_step_items=pre_step_items, + new_step_items=retained_items, + next_step=NextStepRunAgain(), + tool_input_guardrail_results=tool_input_guardrail_results, + tool_output_guardrail_results=tool_output_guardrail_results, + ) + ) + try: ( function_results, @@ -922,25 +960,13 @@ def _commit_accepted_response_tool_output(item: RunItem) -> None: context_wrapper=context_wrapper, run_config=run_config, tool_output_committer=_commit_accepted_response_tool_output, + on_tool_error_selected=on_tool_error_selected, tool_input_guardrail_results=tool_input_guardrail_results, tool_output_guardrail_results=tool_output_guardrail_results, ) - except Exception: - # Accepted server responses already have their own resumable checkpoint. - if on_tool_execution_error is not None and not server_manages_conversation: - retained_items = _completed_tool_step_items(new_step_items, completed_outputs) - if retained_items: - on_tool_execution_error( - SingleStepResult( - original_input=original_input, - model_response=new_response, - pre_step_items=pre_step_items, - new_step_items=retained_items, - next_step=NextStepRunAgain(), - tool_input_guardrail_results=tool_input_guardrail_results, - tool_output_guardrail_results=tool_output_guardrail_results, - ) - ) + except (Exception, asyncio.CancelledError) as error: + if not isinstance(error, asyncio.CancelledError) or _is_tool_local_cancellation(error): + _publish_completed_tools() raise new_step_items.extend( @@ -999,20 +1025,24 @@ def _commit_accepted_response_tool_output(item: RunItem) -> None: _register_tool_call_items(context_wrapper, new_step_items) if run_handoffs := processed_response.handoffs: - return await execute_handoffs_call( - public_agent=public_agent, - original_input=original_input, - pre_step_items=pre_step_items, - new_step_items=new_step_items, - new_response=new_response, - run_handoffs=run_handoffs, - hooks=hooks, - context_wrapper=context_wrapper, - run_config=run_config, - server_manages_conversation=server_manages_conversation, - tool_input_guardrail_results=tool_input_guardrail_results, - tool_output_guardrail_results=tool_output_guardrail_results, - ) + try: + return await execute_handoffs_call( + public_agent=public_agent, + original_input=original_input, + pre_step_items=pre_step_items, + new_step_items=new_step_items, + new_response=new_response, + run_handoffs=run_handoffs, + hooks=hooks, + context_wrapper=context_wrapper, + run_config=run_config, + server_manages_conversation=server_manages_conversation, + tool_input_guardrail_results=tool_input_guardrail_results, + tool_output_guardrail_results=tool_output_guardrail_results, + ) + except Exception: + _publish_completed_tools() + raise tool_final_output = await _maybe_finalize_from_tool_results( public_agent=public_agent, @@ -3842,6 +3872,7 @@ async def get_single_step_result_from_response( before_side_effects: Callable[[], Awaitable[None]] | None = None, run_state: RunState[Any] | None = None, on_tool_execution_error: Callable[[SingleStepResult], None] | None = None, + on_tool_error_selected: Callable[[], None] | None = None, ) -> SingleStepResult: item_agent = bindings.public_agent try: @@ -3905,4 +3936,5 @@ async def get_single_step_result_from_response( precomputed_skipped_raw_item_ids=skipped_raw_item_ids, run_state=run_state, on_tool_execution_error=on_tool_execution_error, + on_tool_error_selected=on_tool_error_selected, ) diff --git a/src/agents/util/_asyncio_tasks.py b/src/agents/util/_asyncio_tasks.py index 338bd4e04d..f557f672c8 100644 --- a/src/agents/util/_asyncio_tasks.py +++ b/src/agents/util/_asyncio_tasks.py @@ -29,7 +29,7 @@ async def gather_with_cancel( awaitable_2: Awaitable[T2], /, *, - on_child_failure: Callable[[], None] | None = None, + on_child_failure: Callable[[BaseException], None] | None = None, ) -> tuple[T1, T2]: ... @@ -40,7 +40,7 @@ async def gather_with_cancel( awaitable_3: Awaitable[T3], /, *, - on_child_failure: Callable[[], None] | None = None, + on_child_failure: Callable[[BaseException], None] | None = None, ) -> tuple[T1, T2, T3]: ... @@ -52,7 +52,7 @@ async def gather_with_cancel( awaitable_4: Awaitable[T4], /, *, - on_child_failure: Callable[[], None] | None = None, + on_child_failure: Callable[[BaseException], None] | None = None, ) -> tuple[T1, T2, T3, T4]: ... @@ -65,7 +65,7 @@ async def gather_with_cancel( awaitable_5: Awaitable[T5], /, *, - on_child_failure: Callable[[], None] | None = None, + on_child_failure: Callable[[BaseException], None] | None = None, ) -> tuple[T1, T2, T3, T4, T5]: ... @@ -79,20 +79,20 @@ async def gather_with_cancel( awaitable_6: Awaitable[T6], /, *, - on_child_failure: Callable[[], None] | None = None, + on_child_failure: Callable[[BaseException], None] | None = None, ) -> tuple[T1, T2, T3, T4, T5, T6]: ... @overload async def gather_with_cancel( *awaitables: Awaitable[T], - on_child_failure: Callable[[], None] | None = None, + on_child_failure: Callable[[BaseException], None] | None = None, ) -> tuple[T, ...]: ... async def gather_with_cancel( *awaitables: Awaitable[Any], - on_child_failure: Callable[[], None] | None = None, + on_child_failure: Callable[[BaseException], None] | None = None, ) -> tuple[Any, ...]: """Gather awaitables, cancelling and draining siblings when one raises.""" tasks = [asyncio.ensure_future(awaitable) for awaitable in awaitables] @@ -102,9 +102,9 @@ async def gather_with_cancel( await asyncio.wait((gather_future,)) try: return tuple(gather_future.result()) - except BaseException: + except BaseException as error: if on_child_failure is not None: - on_child_failure() + on_child_failure(error) raise except GeneratorExit: # Coroutine closure cannot suspend; the owner must tear down child tasks. diff --git a/tests/test_asyncio_tasks.py b/tests/test_asyncio_tasks.py index 1c4858748f..fb001b9c9c 100644 --- a/tests/test_asyncio_tasks.py +++ b/tests/test_asyncio_tasks.py @@ -32,7 +32,7 @@ async def fail_after_sibling_starts() -> None: await gather_with_cancel( sibling(), fail_after_sibling_starts(), - on_child_failure=child_failure_reported.set, + on_child_failure=lambda _error: child_failure_reported.set(), ) assert child_failure_reported.is_set() @@ -61,7 +61,7 @@ async def child() -> None: gather_with_cancel( child(), child(), - on_child_failure=child_failure_reported.set, + on_child_failure=lambda _error: child_failure_reported.set(), ) ) await all_children_started.wait() @@ -170,7 +170,9 @@ async def test_closing_task_helper_leaves_child_cleanup_to_owner(producer_consum coro = ( run_producer_consumer(*children, on_failure=child_failure_reported.set) if producer_consumer - else gather_with_cancel(*children, on_child_failure=child_failure_reported.set) + else gather_with_cancel( + *children, on_child_failure=lambda _error: child_failure_reported.set() + ) ) try: # Drive the coroutine as its owner; do not close a live asyncio Task's coroutine. diff --git a/tests/test_native_tool_failure_history.py b/tests/test_native_tool_failure_history.py new file mode 100644 index 0000000000..bbf2ad4a3b --- /dev/null +++ b/tests/test_native_tool_failure_history.py @@ -0,0 +1,416 @@ +"""Native output finalization is owned by the already-invoked tool on sibling failure.""" + +from __future__ import annotations + +import asyncio +import json +from typing import Any + +import pytest +from openai.types.responses import ResponseApplyPatchToolCall, ResponseCustomToolCall +from openai.types.responses.response_computer_tool_call import ( + ActionScreenshot, + ResponseComputerToolCall, +) +from openai.types.responses.response_output_item import LocalShellCall, LocalShellCallAction + +from agents import ( + Agent, + ApplyPatchTool, + ComputerTool, + CustomTool, + LocalShellTool, + RunHooks, + Runner, + RunState, + ShellTool, + SQLiteSession, + UserError, +) +from agents.decorators import tool +from agents.run_internal import tool_planning +from agents.testing import ScriptedModel, function_call + +from .model_test_helpers import get_exact_output_stream_step +from .test_computer_tool_lifecycle import FakeComputer +from .test_tool_custom_data import RecordingEditor +from .utils.hitl import make_shell_call + + +def _native_tool(kind, extractor, effects): + def execute(*_args): + effects.append(kind) + return "completed" + + if kind == "custom": + native = CustomTool( + name="native", + description="Synthetic native tool", + on_invoke_tool=execute, + custom_data_extractor=extractor, + ) + call = ResponseCustomToolCall( + type="custom_tool_call", name="native", call_id="native", input="synthetic" + ) + elif kind == "computer": + + class Computer(FakeComputer): + def screenshot(self): + return execute() + + native = ComputerTool(computer=Computer(), custom_data_extractor=extractor) + call = ResponseComputerToolCall( + id="native", + type="computer_call", + action=ActionScreenshot(type="screenshot"), + call_id="native", + pending_safety_checks=[], + status="completed", + ) + elif kind == "patch": + + class Editor(RecordingEditor): + def update_file(self, operation): + execute() + return super().update_file(operation) + + native = ApplyPatchTool(editor=Editor(), custom_data_extractor=extractor) + call = ResponseApplyPatchToolCall( + type="apply_patch_call", + id="native", + call_id="native", + status="completed", + operation={"type": "update_file", "path": "synthetic.txt", "diff": "-a\n+b\n"}, + ) + elif kind == "shell": + native = ShellTool(executor=execute) + call = make_shell_call("native") + else: + native = LocalShellTool(executor=execute) + call = LocalShellCall( + id="native", + type="local_shell_call", + call_id="native", + status="completed", + action=LocalShellCallAction(type="exec", command=["synthetic"], env={}), + ) + return native, call + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize( + "kind,phase,hook_raises", + [ + ("computer", "extractor", False), + ("custom", "extractor", False), + ("patch", "extractor", False), + ("custom", "hook", False), + ("shell", "hook", False), + ("local_shell", "hook", False), + ("custom", "hook", True), + ], +) +async def test_native_finalization_survives_sibling_failure( + monkeypatch, streaming, kind, phase, hook_raises +): + entered = asyncio.Event() + release = asyncio.Event() + category_failed = asyncio.Event() + finalized = asyncio.Event() + effects: list[str] = [] + gather = tool_planning.gather_with_cancel + + async def observe_failure(*args, on_child_failure=None): + def notify(error): + if on_child_failure is not None: + on_child_failure(error) + category_failed.set() + + return await gather(*args, on_child_failure=notify) + + # Observe the existing category-failure boundary without replacing its cancellation. + monkeypatch.setattr(tool_planning, "gather_with_cancel", observe_failure) + + async def extract(_context): + if phase == "extractor": + entered.set() + await release.wait() + return {"finalized": True} + + native, call = _native_tool(kind, extract, effects) + + @tool(failure_error_function=None) + async def failed_tool() -> str: + await entered.wait() + raise ValueError("synthetic sibling failure") + + class Hooks(RunHooks): + async def on_tool_end(self, context, agent, tool, result): + if tool is native: + if phase == "hook": + entered.set() + await release.wait() + finalized.set() + if hook_raises: + raise ValueError("synthetic native end-hook failure") + + calls = [call, function_call("failed_tool", {}, call_id="failed")] + agent = Agent( + name="native-agent", + model=ScriptedModel([get_exact_output_stream_step(calls) if streaming else calls]), + tools=[native, failed_tool], + ) + session = SQLiteSession("native") + streamed = None + output_events: list[Any] = [] + + async def run(): + nonlocal streamed + if streaming: + streamed = Runner.run_streamed(agent, "go", session=session, hooks=Hooks()) + async for event in streamed.stream_events(): + if ( + event.type == "run_item_stream_event" + and event.item.type == "tool_call_output_item" + ): + output_events.append(event.item) + else: + await Runner.run(agent, "go", session=session, hooks=Hooks()) + + task = asyncio.create_task(run()) + try: + await asyncio.wait_for(category_failed.wait(), timeout=5) + assert effects == [kind] + assert not finalized.is_set() + release.set() + with pytest.raises(UserError, match="synthetic sibling failure") as caught: + await task + assert finalized.is_set() + assert caught.value.run_data is not None + items = caught.value.run_data.new_items + assert len(items) == 2 + assert items[0].type == "tool_call_item" + assert items[1].type == "tool_call_output_item" + if kind in {"computer", "custom", "patch"}: + assert items[1].custom_data == {"finalized": True} + assert [i.get("call_id") for i in await session.get_items()] == [None, "native", "native"] + if streamed is not None: + assert len(output_events) == 1 + assert output_events[0].custom_data == items[1].custom_data + state = await RunState.from_json(agent, streamed.to_state().to_json()) + assert state._generated_items[-1].custom_data == items[1].custom_data + finally: + release.set() + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stop", ["parent_cancel", "drain_timeout"]) +async def test_native_finalization_does_not_delay_parent_or_publish_unfinished_output( + monkeypatch, stop +): + from agents.run_internal import tool_execution + + entered = asyncio.Event() + release = asyncio.Event() + exited = asyncio.Event() + effects: list[str] = [] + + async def extract(_context): + entered.set() + try: + await release.wait() + return {"finalized": True} + finally: + exited.set() + + native, call = _native_tool("custom", extract, effects) + + @tool(failure_error_function=None) + async def sibling() -> str: + await entered.wait() + if stop == "parent_cancel": + await release.wait() + raise ValueError("synthetic sibling failure") + + # This test exercises the bound, rather than spending the production drain budget. + monkeypatch.setattr(tool_execution, "_FUNCTION_TOOL_POST_INVOKE_WAIT_SECONDS", 0.001) + agent = Agent( + name="native-agent", + model=ScriptedModel([[call, function_call("sibling", {}, call_id="sibling")]]), + tools=[native, sibling], + ) + session = SQLiteSession("native-stop") + task = asyncio.create_task(Runner.run(agent, "go", session=session)) + try: + await asyncio.wait_for(entered.wait(), timeout=5) + if stop == "parent_cancel": + task.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, timeout=5) + await asyncio.wait_for(exited.wait(), timeout=5) + else: + with pytest.raises(UserError, match="synthetic sibling failure") as caught: + await asyncio.wait_for(task, timeout=5) + assert not exited.is_set() + assert caught.value.run_data is not None + assert caught.value.run_data.new_items == [] + release.set() + await asyncio.wait_for(exited.wait(), timeout=5) + assert caught.value.run_data.new_items == [] + assert effects == ["custom"] + assert [item.get("role") for item in await session.get_items()] == ["user"] + finally: + release.set() + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + await asyncio.wait_for(exited.wait(), timeout=5) + session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("verdict", ["pass", "reject", "error"]) +@pytest.mark.parametrize("kind", ["custom", "function"]) +async def test_input_verdict_during_tool_drain_preserves_selected_tool_error( + monkeypatch, verdict, kind +): + from agents import GuardrailFunctionOutput, input_guardrail + from agents.run_internal.tool_execution import _FunctionToolBatchExecutor + + entered = asyncio.Event() + category_failed = asyncio.Event() + release = asyncio.Event() + effects: list[str] = [] + gather = tool_planning.gather_with_cancel + + async def observe_failure(*args, on_child_failure=None): + def notify(error): + if on_child_failure is not None: + on_child_failure(error) + category_failed.set() + + return await gather(*args, on_child_failure=notify) + + monkeypatch.setattr(tool_planning, "gather_with_cancel", observe_failure) + + async def extract(_context): + entered.set() + await release.wait() + return {"finalized": True} + + @tool + async def completed() -> str: + effects.append("function") + return "completed" + + class Hooks(RunHooks): + async def on_tool_end(self, context, agent, tool, result): + if tool is completed: + await extract(None) + + if kind == "function": + original_drain = _FunctionToolBatchExecutor._raise_failure_after_draining_siblings + + async def observe_function_drain(self, failure): + category_failed.set() + return await original_drain(self, failure) + + monkeypatch.setattr( + _FunctionToolBatchExecutor, + "_raise_failure_after_draining_siblings", + observe_function_drain, + ) + native, call = completed, function_call("completed", {}, call_id="native") + else: + native, call = _native_tool("custom", extract, effects) + + @tool(failure_error_function=None) + async def fail() -> str: + await entered.wait() + raise ValueError("primary tool failure") + + @input_guardrail + async def input_verdict(context, agent, input): + await category_failed.wait() + release.set() + if verdict == "error": + raise ValueError("late input failure") + return GuardrailFunctionOutput(output_info=None, tripwire_triggered=verdict == "reject") + + calls = [call, function_call("fail", {}, call_id="fail")] + agent = Agent( + name="native-agent", + tools=[native, fail], + model=ScriptedModel([get_exact_output_stream_step(calls)]), + input_guardrails=[input_verdict], + ) + session = SQLiteSession("verdict-during-drain") + result = Runner.run_streamed(agent, "go", session=session, hooks=Hooks()) + try: + with pytest.raises(UserError, match="primary tool failure"): + async for _ in result.stream_events(): + pass + assert result.run_loop_task is not None and not result.run_loop_task.cancelled() + assert effects == [kind] + assert len(result.new_items) == (2 if verdict == "pass" else 0) + assert len(await session.get_items()) == (3 if verdict == "pass" else 1) + if verdict == "pass" and kind == "custom": + assert result.new_items[-1].custom_data == {"finalized": True} + finally: + release.set() + session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +async def test_approved_native_output_survives_custom_data_failure(streaming): + from agents.testing import assistant_message + + effects: list[str] = [] + + async def extract(_context): + raise ValueError("synthetic extractor failure") + + native, call = _native_tool("custom", extract, effects) + native.needs_approval = True + agent = Agent( + name="approved-native", + tools=[native], + model=ScriptedModel([[call], [assistant_message("done")]]), + ) + session = SQLiteSession("approved-native-extractor-failure") + try: + interrupted = await Runner.run(agent, "go", session=session) + state = interrupted.to_state() + state.approve(interrupted.interruptions[0]) + with pytest.raises(ValueError, match="synthetic extractor failure"): + if streaming: + result = Runner.run_streamed(agent, state, session=session) + async for _ in result.stream_events(): + pass + else: + await Runner.run(agent, state, session=session) + assert effects == ["custom"] + outputs = [item for item in state._generated_items if item.type == "tool_call_output_item"] + assert len(outputs) == 1 + assert outputs[0].output == "completed" + # Recovery keeps the invocation checkpoint, not the live metadata-finalization flag. + serialized = state.to_json() + assert "_custom_data_pending" not in json.dumps(serialized) + restored = await RunState.from_json(agent, serialized) + resumed = await Runner.run(agent, restored, session=session) + assert resumed.final_output == "done" + assert effects == ["custom"] + saved_outputs = [ + item + for item in await session.get_items() + if item.get("type") == "custom_tool_call_output" + ] + assert len(saved_outputs) == 1 + finally: + session.close() diff --git a/tests/test_tool_batch_failure_history.py b/tests/test_tool_batch_failure_history.py index bc473220c9..667fd95b99 100644 --- a/tests/test_tool_batch_failure_history.py +++ b/tests/test_tool_batch_failure_history.py @@ -431,3 +431,102 @@ async def on_tool_end(self, context, agent, tool, result): assert state._pending_session_write is None finally: session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("verdict", ["pass", "reject", "error", "cancel"]) +async def test_streamed_partial_history_waits_for_input_verdict(monkeypatch, verdict: str): + from agents.run_internal import run_loop + + finished = asyncio.Event() + release_verdict = asyncio.Event() + waiting_for_verdict = asyncio.Event() + effects: list[str] = [] + original_wait = run_loop.input_guardrail_tripwire_triggered_for_stream + + async def observe_verdict_wait(*args, **kwargs): + # Control the ordering at the existing verdict wait; all assertions below + # exercise the public run result and Session, not the helper's call shape. + waiting_for_verdict.set() + return await original_wait(*args, **kwargs) + + monkeypatch.setattr( + run_loop, "input_guardrail_tripwire_triggered_for_stream", observe_verdict_wait + ) + + @input_guardrail + async def delayed_verdict(context, agent, input): + await release_verdict.wait() + if verdict == "error": + raise ValueError("synthetic input verdict failure") + return GuardrailFunctionOutput( + output_info="checked", tripwire_triggered=verdict == "reject" + ) + + @tool + async def completed_tool() -> str: + effects.append("done") + return "completed" + + @tool(failure_error_function=None) + async def failed_tool() -> str: + await finished.wait() + raise ValueError("synthetic sibling failure") + + class Hooks(RunHooks): + async def on_tool_end(self, context, agent, tool, result): + finished.set() + + agent = Agent( + name="support", + model=ScriptedModel( + [ + [ + function_call("completed_tool", {}, call_id="done"), + function_call("failed_tool", {}, call_id="failed"), + ] + ] + ), + tools=[completed_tool, failed_tool], + input_guardrails=[delayed_verdict], + ) + session = SQLiteSession("delayed") + result = Runner.run_streamed(agent, "go", session=session, hooks=Hooks()) + + async def consume(): + async for _ in result.stream_events(): + pass + + consumer = asyncio.create_task(consume()) + try: + await asyncio.wait_for(waiting_for_verdict.wait(), timeout=5) + assert effects == ["done"] + assert result.new_items == [] + assert _shape(await session.get_items()) == ["user"] + assert result.run_loop_task is not None + if verdict == "cancel": + result.run_loop_task.cancel() + with pytest.raises(asyncio.CancelledError): + await result.run_loop_task + assert result._input_guardrails_task is not None + assert result._input_guardrails_task.cancelled() + else: + release_verdict.set() + with pytest.raises(UserError, match="synthetic sibling failure"): + await result.run_loop_task + with pytest.raises(UserError, match="synthetic sibling failure"): + await consumer + expected = ["function_call:done", "function_call_output:done"] if verdict == "pass" else [] + assert _shape(result.to_input_list()) == ["user", *expected] + assert _shape(await session.get_items()) == ["user", *expected] + state = result.to_state() + assert _shape([item.to_input_item() for item in state._generated_items]) == expected + finally: + release_verdict.set() + if result.run_loop_task is not None and not result.run_loop_task.done(): + result.run_loop_task.cancel() + await asyncio.gather(result.run_loop_task, return_exceptions=True) + if not consumer.done(): + consumer.cancel() + await asyncio.gather(consumer, return_exceptions=True) + session.close() diff --git a/tests/test_tool_failure_boundaries.py b/tests/test_tool_failure_boundaries.py new file mode 100644 index 0000000000..e82134925d --- /dev/null +++ b/tests/test_tool_failure_boundaries.py @@ -0,0 +1,146 @@ +"""Completed work survives supported failures outside ordinary tool exceptions.""" + +from __future__ import annotations + +import asyncio + +import pytest +from openai.types.responses import ResponseCustomToolCall, ResponseFunctionWebSearch + +from agents import Agent, CustomTool, RunHooks, Runner, RunState, SQLiteSession, UserError, handoff +from agents.decorators import tool +from agents.testing import ScriptedModel, assistant_message, function_call + +from .model_test_helpers import get_exact_output_stream_step +from .test_tool_batch_failure_history import _shape + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize( + "boundary", ["provider", "handoff", "tool_cancel", "native_cancel", "parent_cancel"] +) +async def test_completed_history_at_failure_boundaries(streaming: bool, boundary: str): + completed = asyncio.Event() + blocked = asyncio.Event() + effects: list[str] = [] + + @tool + async def create_ticket() -> str: + effects.append("ticket") + return "ticket T-1" + + @tool(failure_error_function=None) + async def fail() -> str: + if boundary != "provider": + await completed.wait() + if boundary == "parent_cancel": + blocked.set() + await asyncio.Future() + if boundary == "tool_cancel": + raise asyncio.CancelledError("tool-local cancellation") + raise ValueError("synthetic failure") + + async def cancel_native(context, value): + await completed.wait() + raise asyncio.CancelledError("native tool-local cancellation") + + native = CustomTool(name="native", description="synthetic", on_invoke_tool=cancel_native) + + class Hooks(RunHooks): + async def on_tool_end(self, context, agent, tool, result): + if tool.name == "create_ticket": + completed.set() + + async def fail_handoff(context): + assert effects == ["ticket"] + raise UserError("handoff failure") + + target = Agent(name="target", model=ScriptedModel([[assistant_message("done")]])) + transfer = handoff(target, on_handoff=fail_handoff) + calls = ( + [ + ResponseFunctionWebSearch( + id="ws_done", + type="web_search_call", + status="completed", + action={"type": "search", "query": "synthetic query"}, + ), + function_call("fail", {}, call_id="fail"), + ] + if boundary == "provider" + else [ + function_call("create_ticket", {}, call_id="ticket"), + function_call( + transfer.tool_name if boundary == "handoff" else "fail", {}, call_id="fail" + ), + ] + ) + if boundary == "native_cancel": + calls[-1] = ResponseCustomToolCall( + type="custom_tool_call", name="native", call_id="fail", input="synthetic" + ) + model = ScriptedModel( + [get_exact_output_stream_step(calls) if streaming else calls, [assistant_message("done")]] + ) + agent = Agent( + name="support", model=model, tools=[create_ticket, fail, native], handoffs=[transfer] + ) + session = SQLiteSession("failure-boundaries") + result = None + caught = None + try: + if streaming: + result = Runner.run_streamed(agent, "go", session=session, hooks=Hooks()) + + async def consume(): + async for _ in result.stream_events(): + pass + + consumer = asyncio.create_task(consume()) + if boundary == "parent_cancel": + await asyncio.wait_for(blocked.wait(), 2) + result.cancel() + if boundary in ("tool_cancel", "native_cancel", "parent_cancel"): + await asyncio.wait_for(consumer, 2) + else: + with pytest.raises(UserError) as exc: + await asyncio.wait_for(consumer, 2) + caught = exc.value + else: + task = asyncio.create_task(Runner.run(agent, "go", session=session, hooks=Hooks())) + if boundary == "parent_cancel": + await asyncio.wait_for(blocked.wait(), 2) + task.cancel() + error_type = ( + asyncio.CancelledError + if boundary in ("tool_cancel", "native_cancel", "parent_cancel") + else UserError + ) + with pytest.raises(error_type) as exc: + await asyncio.wait_for(task, 2) + caught = exc.value + expected = ( + [] + if boundary == "parent_cancel" + else ["web_search_call"] + if boundary == "provider" + else ["function_call:ticket", "function_call_output:ticket"] + ) + assert _shape(await session.get_items())[1:] == expected + if isinstance(caught, UserError): + assert caught.run_data is not None + assert _shape([item.to_input_item() for item in caught.run_data.new_items]) == expected + if result is not None: + assert _shape(result.to_input_list())[1:] == expected + if boundary != "parent_cancel": + restored = await RunState.from_json(agent, result.to_state().to_json()) + assert ( + _shape([item.to_input_item() for item in restored._generated_items]) == expected + ) + if boundary != "parent_cancel": + await Runner.run(agent, "continue", session=session) + assert _shape(model.calls[-1].input)[1:-1] == expected + assert effects == ([] if boundary == "provider" else ["ticket"]) + finally: + session.close() From 959ebbb16a1760654a28a14bbd4ba538a6b30e68 Mon Sep 17 00:00:00 2001 From: Justin Beckwith Date: Wed, 7 Oct 2026 09:51:07 -0700 Subject: [PATCH 4/8] fix: finish tool failure cleanup and preserve hosted results --- src/agents/run_internal/run_loop.py | 2 + src/agents/run_internal/tool_execution.py | 8 +- src/agents/run_internal/turn_resolution.py | 31 +++--- tests/test_native_tool_failure_history.py | 6 +- tests/test_tool_failure_boundaries.py | 104 +++++++++++++++++++-- 5 files changed, 124 insertions(+), 27 deletions(-) diff --git a/src/agents/run_internal/run_loop.py b/src/agents/run_internal/run_loop.py index 97e536fcf5..bb03ce5869 100644 --- a/src/agents/run_internal/run_loop.py +++ b/src/agents/run_internal/run_loop.py @@ -2062,6 +2062,8 @@ def _record_max_turns_handler_output( item for item in partial_result.new_step_items if item.type == "tool_call_output_item" + # Model-provided outputs were already emitted before execution. + and _stream_event_item_occurrence_key(item) is None ], streamed_result._event_queue, ) diff --git a/src/agents/run_internal/tool_execution.py b/src/agents/run_internal/tool_execution.py index 597a1fca3f..26b1a24826 100644 --- a/src/agents/run_internal/tool_execution.py +++ b/src/agents/run_internal/tool_execution.py @@ -2387,7 +2387,13 @@ async def run_native_tool_post_invoke( except asyncio.CancelledError: if sibling_category_failure.is_set(): try: - await asyncio.wait((task,), timeout=_FUNCTION_TOOL_POST_INVOKE_WAIT_SECONDS) + _, pending = await asyncio.wait( + (task,), timeout=_FUNCTION_TOOL_POST_INVOKE_WAIT_SECONDS + ) + if pending: + task.cancel() + # Bound cancellation cleanup too: application hooks may suppress cancellation. + await asyncio.wait((task,), timeout=_FUNCTION_TOOL_POST_INVOKE_WAIT_SECONDS) except BaseException: task.cancel() raise diff --git a/src/agents/run_internal/turn_resolution.py b/src/agents/run_internal/turn_resolution.py index d588854d15..15869fb923 100644 --- a/src/agents/run_internal/turn_resolution.py +++ b/src/agents/run_internal/turn_resolution.py @@ -898,7 +898,12 @@ async def execute_tools_and_side_effects( skipped_raw_item_ids=skipped_raw_item_ids, ) - completed_outputs: list[RunItem] = [] + completed_outputs: list[RunItem] = [ + item + for item in new_step_items + if isinstance(item, ToolCallOutputItem) + and get_mapping_or_attr(item.raw_item, "status") == "completed" + ] tool_input_guardrail_results: list[ToolInputGuardrailResult] = [] tool_output_guardrail_results: list[ToolOutputGuardrailResult] = [] prior_input_results = list(run_state._tool_input_guardrail_results) if run_state else [] @@ -1016,16 +1021,16 @@ def _publish_completed_tools() -> None: processed_response=processed_response, ) - await _append_mcp_callback_results( - agent=public_agent, - requests=plan.mcp_requests_with_callback, - context_wrapper=context_wrapper, - append_item=new_step_items.append, - ) - _register_tool_call_items(context_wrapper, new_step_items) + try: + await _append_mcp_callback_results( + agent=public_agent, + requests=plan.mcp_requests_with_callback, + context_wrapper=context_wrapper, + append_item=new_step_items.append, + ) + _register_tool_call_items(context_wrapper, new_step_items) - if run_handoffs := processed_response.handoffs: - try: + if run_handoffs := processed_response.handoffs: return await execute_handoffs_call( public_agent=public_agent, original_input=original_input, @@ -1040,9 +1045,9 @@ def _publish_completed_tools() -> None: tool_input_guardrail_results=tool_input_guardrail_results, tool_output_guardrail_results=tool_output_guardrail_results, ) - except Exception: - _publish_completed_tools() - raise + except Exception: + _publish_completed_tools() + raise tool_final_output = await _maybe_finalize_from_tool_results( public_agent=public_agent, diff --git a/tests/test_native_tool_failure_history.py b/tests/test_native_tool_failure_history.py index bbf2ad4a3b..d2eeb6a9ba 100644 --- a/tests/test_native_tool_failure_history.py +++ b/tests/test_native_tool_failure_history.py @@ -256,12 +256,10 @@ async def sibling() -> str: else: with pytest.raises(UserError, match="synthetic sibling failure") as caught: await asyncio.wait_for(task, timeout=5) - assert not exited.is_set() + assert exited.is_set() + assert not release.is_set() assert caught.value.run_data is not None assert caught.value.run_data.new_items == [] - release.set() - await asyncio.wait_for(exited.wait(), timeout=5) - assert caught.value.run_data.new_items == [] assert effects == ["custom"] assert [item.get("role") for item in await session.get_items()] == ["user"] finally: diff --git a/tests/test_tool_failure_boundaries.py b/tests/test_tool_failure_boundaries.py index e82134925d..51cf26ebb3 100644 --- a/tests/test_tool_failure_boundaries.py +++ b/tests/test_tool_failure_boundaries.py @@ -5,9 +5,26 @@ import asyncio import pytest -from openai.types.responses import ResponseCustomToolCall, ResponseFunctionWebSearch - -from agents import Agent, CustomTool, RunHooks, Runner, RunState, SQLiteSession, UserError, handoff +from openai.types.responses import ( + ResponseCustomToolCall, + ResponseFunctionShellToolCall, + ResponseFunctionShellToolCallOutput, + ResponseFunctionWebSearch, +) +from openai.types.responses.response_output_item import McpApprovalRequest + +from agents import ( + Agent, + CustomTool, + HostedMCPTool, + RunHooks, + Runner, + RunState, + ShellTool, + SQLiteSession, + UserError, + handoff, +) from agents.decorators import tool from agents.testing import ScriptedModel, assistant_message, function_call @@ -18,7 +35,16 @@ @pytest.mark.asyncio @pytest.mark.parametrize("streaming", [False, True]) @pytest.mark.parametrize( - "boundary", ["provider", "handoff", "tool_cancel", "native_cancel", "parent_cancel"] + "boundary", + [ + "provider", + "provider_shell", + "handoff", + "mcp_callback", + "tool_cancel", + "native_cancel", + "parent_cancel", + ], ) async def test_completed_history_at_failure_boundaries(streaming: bool, boundary: str): completed = asyncio.Event() @@ -32,7 +58,7 @@ async def create_ticket() -> str: @tool(failure_error_function=None) async def fail() -> str: - if boundary != "provider": + if boundary not in ("provider", "provider_shell"): await completed.wait() if boundary == "parent_cancel": blocked.set() @@ -80,22 +106,79 @@ async def fail_handoff(context): calls[-1] = ResponseCustomToolCall( type="custom_tool_call", name="native", call_id="fail", input="synthetic" ) + + async def fail_approval(request): + assert effects == ["ticket"] + raise UserError("approval callback failure") + + hosted_mcp = HostedMCPTool( + tool_config={ + "type": "mcp", + "server_label": "synthetic", + "server_url": "https://example.com", + "require_approval": "always", + }, + on_approval_request=fail_approval, + ) + hosted_shell = ShellTool( + environment={"type": "container_reference", "container_id": "cntr_synthetic"} + ) + if boundary == "provider_shell": + calls = [ + ResponseFunctionShellToolCall( + id="sh_done", + type="shell_call", + call_id="shell", + status="completed", + action={"commands": ["echo synthetic"]}, + ), + ResponseFunctionShellToolCallOutput( + id="sh_output", + type="shell_call_output", + call_id="shell", + status="completed", + output=[ + { + "stdout": "synthetic", + "stderr": "", + "outcome": {"type": "exit", "exit_code": 0}, + } + ], + ), + function_call("fail", {}, call_id="fail"), + ] + elif boundary == "mcp_callback": + calls[-1] = McpApprovalRequest( + id="approval", + type="mcp_approval_request", + server_label="synthetic", + arguments="{}", + name="synthetic", + ) model = ScriptedModel( [get_exact_output_stream_step(calls) if streaming else calls, [assistant_message("done")]] ) agent = Agent( - name="support", model=model, tools=[create_ticket, fail, native], handoffs=[transfer] + name="support", + model=model, + tools=[create_ticket, fail, native, hosted_mcp, hosted_shell], + handoffs=[transfer], ) session = SQLiteSession("failure-boundaries") result = None caught = None + output_events = [] try: if streaming: result = Runner.run_streamed(agent, "go", session=session, hooks=Hooks()) async def consume(): - async for _ in result.stream_events(): - pass + async for event in result.stream_events(): + if ( + event.type == "run_item_stream_event" + and event.item.type == "tool_call_output_item" + ): + output_events.append(event.item.to_input_item()) consumer = asyncio.create_task(consume()) if boundary == "parent_cancel": @@ -125,6 +208,8 @@ async def consume(): if boundary == "parent_cancel" else ["web_search_call"] if boundary == "provider" + else ["shell_call:shell", "shell_call_output:shell"] + if boundary == "provider_shell" else ["function_call:ticket", "function_call_output:ticket"] ) assert _shape(await session.get_items())[1:] == expected @@ -132,6 +217,7 @@ async def consume(): assert caught.run_data is not None assert _shape([item.to_input_item() for item in caught.run_data.new_items]) == expected if result is not None: + assert _shape(output_events) == [item for item in expected if "output:" in item] assert _shape(result.to_input_list())[1:] == expected if boundary != "parent_cancel": restored = await RunState.from_json(agent, result.to_state().to_json()) @@ -141,6 +227,6 @@ async def consume(): if boundary != "parent_cancel": await Runner.run(agent, "continue", session=session) assert _shape(model.calls[-1].input)[1:-1] == expected - assert effects == ([] if boundary == "provider" else ["ticket"]) + assert effects == ([] if boundary in ("provider", "provider_shell") else ["ticket"]) finally: session.close() From 5a2bc6993de4ec87eda510a4790e3fcda2588379 Mon Sep 17 00:00:00 2001 From: Justin Beckwith Date: Wed, 7 Oct 2026 10:26:07 -0700 Subject: [PATCH 5/8] fix: retain completed tools through finalization and input admission --- src/agents/run.py | 36 +-- src/agents/run_internal/tool_execution.py | 3 + src/agents/run_internal/turn_resolution.py | 249 ++++++++++----------- tests/test_tool_batch_failure_history.py | 169 ++++++++++++++ tests/test_tool_failure_boundaries.py | 32 ++- 5 files changed, 348 insertions(+), 141 deletions(-) diff --git a/src/agents/run.py b/src/agents/run.py index d819380afa..5c1da0a7e2 100644 --- a/src/agents/run.py +++ b/src/agents/run.py @@ -23,6 +23,7 @@ _clear_data_redacted_error_traceback, _detach_data_redacted_error_traceback, _is_error_data_redacted, + _is_tool_local_cancellation, _prepare_data_redacted_error, _raise_data_redacted_error, ) @@ -1846,18 +1847,23 @@ async def _save_max_turns_handler_output( ) ) raise - except BaseException: - # A non-tripwire failure (the model turn raising, or a - # guardrail raising a non-tripwire error) propagates from - # gather without cancelling the sibling task. Cancel and drain - # whichever side is still pending so it is not left running - # after the run has failed and its exception is not swallowed. - for pending_task in (guardrail_task, model_task): - if not pending_task.done(): - pending_task.cancel() - await asyncio.gather( - guardrail_task, model_task, return_exceptions=True - ) + except BaseException as error: + try: + if partial_tool_results and ( + not isinstance(error, asyncio.CancelledError) + or _is_tool_local_cancellation(error) + ): + # Settle admission without replacing the selected tool + # error. Only successful verdicts admit partial history. + await asyncio.wait((guardrail_task,)) + finally: + # Parent cancellation still cancels and drains both tasks. + for pending_task in (guardrail_task, model_task): + if not pending_task.done(): + pending_task.cancel() + await asyncio.gather( + guardrail_task, model_task, return_exceptions=True + ) raise else: turn_result = await model_task @@ -1887,7 +1893,11 @@ async def _save_max_turns_handler_output( run_state=run_state, on_tool_execution_error=partial_tool_results.append, ) - except (Exception, asyncio.CancelledError): + except (Exception, asyncio.CancelledError) as error: + if isinstance( + error, asyncio.CancelledError + ) and not _is_tool_local_cancellation(error): + raise if not partial_tool_results: raise input_accepted = len(_attempt_input_guardrail_results()) >= len( diff --git a/src/agents/run_internal/tool_execution.py b/src/agents/run_internal/tool_execution.py index 26b1a24826..ff1c6f6020 100644 --- a/src/agents/run_internal/tool_execution.py +++ b/src/agents/run_internal/tool_execution.py @@ -2171,6 +2171,7 @@ async def _invoke_tool_and_run_post_invoke( agent=self.public_agent, tool_origin=get_function_tool_origin(func_tool), ) + output_item._custom_data_pending = True self.output_items_by_tool_run[id(task_state.tool_run)] = output_item if self.tool_output_committer is not None: self.tool_output_committer(output_item) @@ -2188,6 +2189,8 @@ async def _invoke_tool_and_run_post_invoke( self.custom_data_by_tool_run[id(task_state.tool_run)] = custom_data if output_item is not None: output_item.custom_data = custom_data + if output_item is not None: + output_item._custom_data_pending = False await gather_with_cancel( self.hooks.on_tool_end(tool_context, self.public_agent, func_tool, final_result), diff --git a/src/agents/run_internal/turn_resolution.py b/src/agents/run_internal/turn_resolution.py index 15869fb923..db5d184f19 100644 --- a/src/agents/run_internal/turn_resolution.py +++ b/src/agents/run_internal/turn_resolution.py @@ -1045,84 +1045,95 @@ def _publish_completed_tools() -> None: tool_input_guardrail_results=tool_input_guardrail_results, tool_output_guardrail_results=tool_output_guardrail_results, ) - except Exception: - _publish_completed_tools() - raise - - tool_final_output = await _maybe_finalize_from_tool_results( - public_agent=public_agent, - original_input=original_input, - new_response=new_response, - pre_step_items=pre_step_items, - new_step_items=new_step_items, - function_results=function_results, - hooks=hooks, - context_wrapper=context_wrapper, - tool_input_guardrail_results=tool_input_guardrail_results, - tool_output_guardrail_results=tool_output_guardrail_results, - ) - if tool_final_output is not None: - return tool_final_output + tool_final_output = await _maybe_finalize_from_tool_results( + public_agent=public_agent, + original_input=original_input, + new_response=new_response, + pre_step_items=pre_step_items, + new_step_items=new_step_items, + function_results=function_results, + hooks=hooks, + context_wrapper=context_wrapper, + tool_input_guardrail_results=tool_input_guardrail_results, + tool_output_guardrail_results=tool_output_guardrail_results, + ) + if tool_final_output is not None: + return tool_final_output + message_items = [item for item in new_step_items if isinstance(item, MessageOutputItem)] + refusal = ItemHelpers.extract_refusal(message_items[-1].raw_item) if message_items else None + potential_final_output_text = ( + ItemHelpers.extract_text(message_items[-1].raw_item) if message_items else None + ) - message_items = [item for item in new_step_items if isinstance(item, MessageOutputItem)] - refusal = ItemHelpers.extract_refusal(message_items[-1].raw_item) if message_items else None - potential_final_output_text = ( - ItemHelpers.extract_text(message_items[-1].raw_item) if message_items else None - ) + if not processed_response.has_tools_or_approvals_to_run(): + has_tool_activity_without_message = not message_items and bool( + processed_response.tools_used or skipped_raw_item_ids + ) + if not has_tool_activity_without_message: + if refusal: + refusal_error = ModelRefusalError(refusal) + run_error_data = build_run_error_data( + input=original_input, + new_items=pre_step_items + new_step_items, + raw_responses=[new_response], + last_agent=public_agent, + ) + handler_result = await resolve_run_error_handler_result( + error_handlers=error_handlers, + error_kind="model_refusal", + error=refusal_error, + context_wrapper=context_wrapper, + run_data=run_error_data, + ) + if handler_result is None: + raise refusal_error - if not processed_response.has_tools_or_approvals_to_run(): - has_tool_activity_without_message = not message_items and bool( - processed_response.tools_used or skipped_raw_item_ids - ) - if not has_tool_activity_without_message: - if refusal: - refusal_error = ModelRefusalError(refusal) - run_error_data = build_run_error_data( - input=original_input, - new_items=pre_step_items + new_step_items, - raw_responses=[new_response], - last_agent=public_agent, - ) - handler_result = await resolve_run_error_handler_result( - error_handlers=error_handlers, - error_kind="model_refusal", - error=refusal_error, - context_wrapper=context_wrapper, - run_data=run_error_data, - ) - if handler_result is None: - raise refusal_error - - final_output = validate_handler_final_output( - public_agent, handler_result.final_output - ) - if handler_result.include_in_history: - output_text = format_final_output_text(public_agent, final_output) - new_step_items.append(create_message_output_item(public_agent, output_text)) - return await execute_final_output_call( - public_agent=public_agent, - original_input=original_input, - new_response=new_response, - pre_step_items=pre_step_items, - new_step_items=new_step_items, - final_output=final_output, - hooks=hooks, - context_wrapper=context_wrapper, - tool_input_guardrail_results=tool_input_guardrail_results, - tool_output_guardrail_results=tool_output_guardrail_results, - ) - if output_schema is not None and not output_schema.is_plain_text(): - if potential_final_output_text: - validation_error: ModelBehaviorError | None = None - try: - final_output = output_schema.validate_json(potential_final_output_text) - except ModelBehaviorError as error: - if _is_error_data_redacted(error): - validation_error = error - else: + final_output = validate_handler_final_output( + public_agent, handler_result.final_output + ) + if handler_result.include_in_history: + output_text = format_final_output_text(public_agent, final_output) + new_step_items.append(create_message_output_item(public_agent, output_text)) + return await execute_final_output_call( + public_agent=public_agent, + original_input=original_input, + new_response=new_response, + pre_step_items=pre_step_items, + new_step_items=new_step_items, + final_output=final_output, + hooks=hooks, + context_wrapper=context_wrapper, + tool_input_guardrail_results=tool_input_guardrail_results, + tool_output_guardrail_results=tool_output_guardrail_results, + ) + if output_schema is not None and not output_schema.is_plain_text(): + if potential_final_output_text: + validation_error: ModelBehaviorError | None = None + try: + final_output = output_schema.validate_json(potential_final_output_text) + except ModelBehaviorError as error: + if _is_error_data_redacted(error): + validation_error = error + else: + resolved_handler_output = await _resolve_invalid_final_output( + error_handlers=error_handlers, + error=error, + public_agent=public_agent, + original_input=original_input, + new_response=new_response, + new_items=pre_step_items + new_step_items, + context_wrapper=context_wrapper, + ) + if resolved_handler_output is None: + raise + final_output, message_item = resolved_handler_output + if message_item is not None: + new_step_items.append(message_item) + + if validation_error is not None: resolved_handler_output = await _resolve_invalid_final_output( error_handlers=error_handlers, - error=error, + error=validation_error, public_agent=public_agent, original_input=original_input, new_response=new_response, @@ -1130,15 +1141,16 @@ def _publish_completed_tools() -> None: context_wrapper=context_wrapper, ) if resolved_handler_output is None: - raise + raise validation_error final_output, message_item = resolved_handler_output if message_item is not None: new_step_items.append(message_item) - - if validation_error is not None: + else: resolved_handler_output = await _resolve_invalid_final_output( error_handlers=error_handlers, - error=validation_error, + error=ModelBehaviorError( + "Model returned no final output for the structured output type." + ), public_agent=public_agent, original_input=original_input, new_response=new_response, @@ -1146,61 +1158,48 @@ def _publish_completed_tools() -> None: context_wrapper=context_wrapper, ) if resolved_handler_output is None: - raise validation_error + return SingleStepResult( + original_input=original_input, + model_response=new_response, + pre_step_items=pre_step_items, + new_step_items=new_step_items, + next_step=NextStepRunAgain(), + tool_input_guardrail_results=tool_input_guardrail_results, + tool_output_guardrail_results=tool_output_guardrail_results, + ) final_output, message_item = resolved_handler_output if message_item is not None: new_step_items.append(message_item) - else: - resolved_handler_output = await _resolve_invalid_final_output( - error_handlers=error_handlers, - error=ModelBehaviorError( - "Model returned no final output for the structured output type." - ), + + return await execute_final_output_call( public_agent=public_agent, original_input=original_input, new_response=new_response, - new_items=pre_step_items + new_step_items, + pre_step_items=pre_step_items, + new_step_items=new_step_items, + final_output=final_output, + hooks=hooks, + context_wrapper=context_wrapper, + tool_input_guardrail_results=tool_input_guardrail_results, + tool_output_guardrail_results=tool_output_guardrail_results, + ) + if output_schema is None or output_schema.is_plain_text(): + return await execute_final_output_call( + public_agent=public_agent, + original_input=original_input, + new_response=new_response, + pre_step_items=pre_step_items, + new_step_items=new_step_items, + final_output=potential_final_output_text or "", + hooks=hooks, context_wrapper=context_wrapper, + tool_input_guardrail_results=tool_input_guardrail_results, + tool_output_guardrail_results=tool_output_guardrail_results, ) - if resolved_handler_output is None: - return SingleStepResult( - original_input=original_input, - model_response=new_response, - pre_step_items=pre_step_items, - new_step_items=new_step_items, - next_step=NextStepRunAgain(), - tool_input_guardrail_results=tool_input_guardrail_results, - tool_output_guardrail_results=tool_output_guardrail_results, - ) - final_output, message_item = resolved_handler_output - if message_item is not None: - new_step_items.append(message_item) - - return await execute_final_output_call( - public_agent=public_agent, - original_input=original_input, - new_response=new_response, - pre_step_items=pre_step_items, - new_step_items=new_step_items, - final_output=final_output, - hooks=hooks, - context_wrapper=context_wrapper, - tool_input_guardrail_results=tool_input_guardrail_results, - tool_output_guardrail_results=tool_output_guardrail_results, - ) - if output_schema is None or output_schema.is_plain_text(): - return await execute_final_output_call( - public_agent=public_agent, - original_input=original_input, - new_response=new_response, - pre_step_items=pre_step_items, - new_step_items=new_step_items, - final_output=potential_final_output_text or "", - hooks=hooks, - context_wrapper=context_wrapper, - tool_input_guardrail_results=tool_input_guardrail_results, - tool_output_guardrail_results=tool_output_guardrail_results, - ) + + except Exception: + _publish_completed_tools() + raise return SingleStepResult( original_input=original_input, diff --git a/tests/test_tool_batch_failure_history.py b/tests/test_tool_batch_failure_history.py index 667fd95b99..3c4f6a96cf 100644 --- a/tests/test_tool_batch_failure_history.py +++ b/tests/test_tool_batch_failure_history.py @@ -530,3 +530,172 @@ async def consume(): consumer.cancel() await asyncio.gather(consumer, return_exceptions=True) session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +async def test_function_partial_history_excludes_pending_custom_data(monkeypatch, streaming): + from agents.run_internal import tool_execution + + from .model_test_helpers import get_exact_output_stream_step + + extracting = asyncio.Event() + release = asyncio.Event() + settled = asyncio.Event() + effects = [] + + async def extract(context): + extracting.set() + try: + await release.wait() + return {"ticket": "T-1"} + finally: + settled.set() + + @tool(custom_data_extractor=extract) + async def create_ticket() -> str: + effects.append("ticket") + return "created" + + @tool(failure_error_function=None) + async def fail() -> str: + await extracting.wait() + raise ValueError("synthetic sibling failure") + + monkeypatch.setattr(tool_execution, "_FUNCTION_TOOL_POST_INVOKE_WAIT_SECONDS", 0.001) + calls = [ + function_call("create_ticket", {}, call_id="ticket"), + function_call("fail", {}, call_id="fail"), + ] + model = ScriptedModel([get_exact_output_stream_step(calls) if streaming else calls]) + agent = Agent(name="metadata", model=model, tools=[create_ticket, fail]) + session = SQLiteSession("pending-function-metadata") + result = None + outputs = [] + try: + with pytest.raises(UserError, match="synthetic sibling failure") as caught: + if streaming: + result = Runner.run_streamed(agent, "go", session=session) + async for event in result.stream_events(): + if ( + event.type == "run_item_stream_event" + and event.item.type == "tool_call_output_item" + ): + outputs.append(event.item) + else: + await Runner.run(agent, "go", session=session) + assert effects == ["ticket"] + assert not settled.is_set() + assert caught.value.run_data is not None + assert caught.value.run_data.new_items == [] + assert _shape(await session.get_items()) == ["user"] + if result is not None: + assert outputs == [] + restored = await RunState.from_json(agent, result.to_state().to_json()) + assert restored._generated_items == [] + release.set() + await asyncio.wait_for(settled.wait(), 2) + assert caught.value.run_data.new_items == [] + finally: + release.set() + await asyncio.wait_for(settled.wait(), 2) + session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("verdict", ["pass", "reject", "error", "cancel", "cancel_at_verdict"]) +async def test_nonstreamed_partial_history_waits_for_input_verdict(monkeypatch, verdict): + completed = asyncio.Event() + release = asyncio.Event() + waiting = asyncio.Event() + guardrail_exited = asyncio.Event() + original_wait = asyncio.wait + + async def observe_wait(tasks, *args, **kwargs): + if len(tasks) == 1 and any( + isinstance(task, asyncio.Task) and task.get_coro().__name__ == "run_input_guardrails" + for task in tasks + ): + waiting.set() + return await original_wait(tasks, *args, **kwargs) + + monkeypatch.setattr(asyncio, "wait", observe_wait) + + @input_guardrail + async def delayed(context, agent, input): + try: + await release.wait() + if verdict == "error": + raise ValueError("synthetic verdict failure") + return GuardrailFunctionOutput( + output_info="checked", tripwire_triggered=verdict == "reject" + ) + finally: + guardrail_exited.set() + + @tool + async def done() -> str: + return "completed" + + @tool(failure_error_function=None) + async def fail() -> str: + await completed.wait() + raise ValueError("synthetic sibling failure") + + class Hooks(RunHooks): + async def on_tool_end(self, context, agent, tool, result): + completed.set() + + agent = Agent( + name="admission", + tools=[done, fail], + input_guardrails=[delayed], + model=ScriptedModel( + [[function_call("done", {}, call_id="done"), function_call("fail", {}, call_id="fail")]] + ), + ) + session = SQLiteSession("nonstreamed-delayed-verdict") + run_task = asyncio.create_task(Runner.run(agent, "go", session=session, hooks=Hooks())) + waiting_task = asyncio.create_task(waiting.wait()) + try: + await original_wait( + (run_task, waiting_task), timeout=5, return_when=asyncio.FIRST_COMPLETED + ) + assert waiting.is_set() + assert not run_task.done() + assert _shape(await session.get_items()) == ["user"] + if verdict in ("cancel", "cancel_at_verdict"): + if verdict == "cancel_at_verdict": + verdict_task = next( + task + for task in asyncio.all_tasks() + if task.get_coro().__name__ == "run_input_guardrails" + ) + verdict_task.add_done_callback(lambda _: run_task.cancel()) + release.set() + else: + run_task.cancel() + with pytest.raises(asyncio.CancelledError): + await run_task + else: + release.set() + with pytest.raises(UserError, match="synthetic sibling failure") as caught: + await run_task + assert caught.value.run_data is not None + expected = ( + ["function_call:done", "function_call_output:done"] if verdict == "pass" else [] + ) + assert ( + _shape([item.to_input_item() for item in caught.value.run_data.new_items]) + == expected + ) + assert guardrail_exited.is_set() + expected = ["function_call:done", "function_call_output:done"] if verdict == "pass" else [] + assert _shape(await session.get_items()) == ["user", *expected] + finally: + release.set() + for task in (run_task, waiting_task): + if not task.done(): + task.cancel() + await asyncio.gather(run_task, waiting_task, return_exceptions=True) + session.close() diff --git a/tests/test_tool_failure_boundaries.py b/tests/test_tool_failure_boundaries.py index 51cf26ebb3..1333b3a34f 100644 --- a/tests/test_tool_failure_boundaries.py +++ b/tests/test_tool_failure_boundaries.py @@ -38,9 +38,12 @@ "boundary", [ "provider", + "provider_final_hook", "provider_shell", "handoff", "mcp_callback", + "tool_behavior", + "final_hook", "tool_cancel", "native_cancel", "parent_cancel", @@ -78,6 +81,19 @@ async def on_tool_end(self, context, agent, tool, result): if tool.name == "create_ticket": completed.set() + async def on_agent_end(self, context, agent, output): + if boundary == "provider_final_hook" or ( + boundary == "final_hook" and effects == ["ticket"] + ): + raise UserError("final hook failure") + + async def fail_tool_behavior(context, results): + if results: + raise UserError("tool behavior failure") + from agents import ToolsToFinalOutputResult + + return ToolsToFinalOutputResult(is_final_output=False, final_output=None) + async def fail_handoff(context): assert effects == ["ticket"] raise UserError("handoff failure") @@ -94,7 +110,7 @@ async def fail_handoff(context): ), function_call("fail", {}, call_id="fail"), ] - if boundary == "provider" + if boundary in ("provider", "provider_final_hook") else [ function_call("create_ticket", {}, call_id="ticket"), function_call( @@ -102,6 +118,8 @@ async def fail_handoff(context): ), ] ) + if boundary == "provider_final_hook": + calls[-1] = assistant_message("done") if boundary == "native_cancel": calls[-1] = ResponseCustomToolCall( type="custom_tool_call", name="native", call_id="fail", input="synthetic" @@ -155,6 +173,8 @@ async def fail_approval(request): arguments="{}", name="synthetic", ) + if boundary in ("tool_behavior", "final_hook"): + calls = calls[:1] model = ScriptedModel( [get_exact_output_stream_step(calls) if streaming else calls, [assistant_message("done")]] ) @@ -164,6 +184,10 @@ async def fail_approval(request): tools=[create_ticket, fail, native, hosted_mcp, hosted_shell], handoffs=[transfer], ) + if boundary == "tool_behavior": + agent.tool_use_behavior = fail_tool_behavior + elif boundary == "final_hook": + agent.tool_use_behavior = "stop_on_first_tool" session = SQLiteSession("failure-boundaries") result = None caught = None @@ -207,7 +231,7 @@ async def consume(): [] if boundary == "parent_cancel" else ["web_search_call"] - if boundary == "provider" + if boundary in ("provider", "provider_final_hook") else ["shell_call:shell", "shell_call_output:shell"] if boundary == "provider_shell" else ["function_call:ticket", "function_call_output:ticket"] @@ -227,6 +251,8 @@ async def consume(): if boundary != "parent_cancel": await Runner.run(agent, "continue", session=session) assert _shape(model.calls[-1].input)[1:-1] == expected - assert effects == ([] if boundary in ("provider", "provider_shell") else ["ticket"]) + assert effects == ( + [] if boundary in ("provider", "provider_final_hook", "provider_shell") else ["ticket"] + ) finally: session.close() From 6c177ae6856d59206238534b9880f43dcb0813d6 Mon Sep 17 00:00:00 2001 From: Justin Beckwith Date: Wed, 7 Oct 2026 10:45:00 -0700 Subject: [PATCH 6/8] fix: preserve tool cancellation across Python 3.10 tasks --- src/agents/run.py | 56 ++++++++++++++---------- src/agents/run_internal/run_loop.py | 18 ++++++++ tests/test_tool_batch_failure_history.py | 32 ++++++++++---- 3 files changed, 73 insertions(+), 33 deletions(-) diff --git a/src/agents/run.py b/src/agents/run.py index 5c1da0a7e2..cec00cf879 100644 --- a/src/agents/run.py +++ b/src/agents/run.py @@ -121,9 +121,11 @@ from .run_internal.run_grouping import resolve_run_grouping_id from .run_internal.run_loop import ( _safe_redacted_persistence_error, + _ToolTaskCancellation, cleanup_models_after_run, finalize_max_turns_handler_output, get_output_schema, + preserve_tool_task_cancellation, resolve_interrupted_turn, run_input_guardrails, run_output_guardrails, @@ -1788,30 +1790,32 @@ async def _save_max_turns_handler_output( raise model_task = asyncio.create_task( - run_single_turn( - bindings=current_bindings, - original_input=original_input, - generated_items=items_for_model, - hooks=hooks, - context_wrapper=context_wrapper, - run_config=run_config, - should_run_agent_start_hooks=should_run_agent_start_hooks, - tool_use_tracker=tool_use_tracker, - server_conversation_tracker=server_conversation_tracker, - session=session, - session_items_to_rewind=( - last_saved_input_snapshot_for_rewind - if not is_resumed_state and session_persistence_enabled - else None - ), - reasoning_item_id_policy=resolved_reasoning_item_id_policy, - prompt_cache_key_resolver=prompt_cache_key_resolver, - error_handlers=error_handlers, - agent_span=current_span, - on_response_accepted=_commit_pending_server_response, - on_response_hooks_started=_mark_response_hooks_started, - run_state=run_state, - on_tool_execution_error=partial_tool_results.append, + preserve_tool_task_cancellation( + run_single_turn( + bindings=current_bindings, + original_input=original_input, + generated_items=items_for_model, + hooks=hooks, + context_wrapper=context_wrapper, + run_config=run_config, + should_run_agent_start_hooks=should_run_agent_start_hooks, + tool_use_tracker=tool_use_tracker, + server_conversation_tracker=server_conversation_tracker, + session=session, + session_items_to_rewind=( + last_saved_input_snapshot_for_rewind + if not is_resumed_state and session_persistence_enabled + else None + ), + reasoning_item_id_policy=resolved_reasoning_item_id_policy, + prompt_cache_key_resolver=prompt_cache_key_resolver, + error_handlers=error_handlers, + agent_span=current_span, + on_response_accepted=_commit_pending_server_response, + on_response_hooks_started=_mark_response_hooks_started, + run_state=run_state, + on_tool_execution_error=partial_tool_results.append, + ) ) ) @@ -1899,6 +1903,8 @@ async def _save_max_turns_handler_output( ) and not _is_tool_local_cancellation(error): raise if not partial_tool_results: + if isinstance(error, _ToolTaskCancellation): + raise error.error from None raise input_accepted = len(_attempt_input_guardrail_results()) >= len( all_input_guardrails @@ -1942,6 +1948,8 @@ async def _save_max_turns_handler_output( ) except Exception: logger.warning("Failed to save completed tools after a tool error") + if isinstance(error, _ToolTaskCancellation): + raise error.error from None raise finally: if current_turn_span is not None: diff --git a/src/agents/run_internal/run_loop.py b/src/agents/run_internal/run_loop.py index bb03ce5869..6f3492b5fe 100644 --- a/src/agents/run_internal/run_loop.py +++ b/src/agents/run_internal/run_loop.py @@ -42,6 +42,7 @@ _copy_data_redacted_process_control_error, _detach_data_redacted_error_traceback, _is_error_data_redacted, + _is_tool_local_cancellation, _mark_error_data_redacted, _mark_error_to_drain_stream_events, _prepare_data_redacted_error, @@ -966,6 +967,23 @@ async def _finalize_streamed_interruption( T = TypeVar("T") +class _ToolTaskCancellation(Exception): + """Carry selected tool cancellation across Task on Python versions that replace it.""" + + def __init__(self, error: asyncio.CancelledError) -> None: + super().__init__() + self.error = error + + +async def preserve_tool_task_cancellation(awaitable: Awaitable[T]) -> T: + try: + return await awaitable + except asyncio.CancelledError as error: + if _is_tool_local_cancellation(error): + raise _ToolTaskCancellation(error) from None + raise + + async def start_streaming( starting_input: str | list[TResponseInputItem], streamed_result: RunResultStreaming, diff --git a/tests/test_tool_batch_failure_history.py b/tests/test_tool_batch_failure_history.py index 3c4f6a96cf..3327887bde 100644 --- a/tests/test_tool_batch_failure_history.py +++ b/tests/test_tool_batch_failure_history.py @@ -603,7 +603,9 @@ async def fail() -> str: @pytest.mark.asyncio -@pytest.mark.parametrize("verdict", ["pass", "reject", "error", "cancel", "cancel_at_verdict"]) +@pytest.mark.parametrize( + "verdict", ["pass", "reject", "error", "cancel", "cancel_at_verdict", "tool_cancel"] +) async def test_nonstreamed_partial_history_waits_for_input_verdict(monkeypatch, verdict): completed = asyncio.Event() release = asyncio.Event() @@ -640,6 +642,8 @@ async def done() -> str: @tool(failure_error_function=None) async def fail() -> str: await completed.wait() + if verdict == "tool_cancel": + raise asyncio.CancelledError("synthetic sibling failure") raise ValueError("synthetic sibling failure") class Hooks(RunHooks): @@ -679,18 +683,28 @@ async def on_tool_end(self, context, agent, tool, result): await run_task else: release.set() - with pytest.raises(UserError, match="synthetic sibling failure") as caught: + error_type = asyncio.CancelledError if verdict == "tool_cancel" else UserError + with pytest.raises( + error_type, match=None if verdict == "tool_cancel" else "synthetic sibling failure" + ) as caught: await run_task - assert caught.value.run_data is not None expected = ( - ["function_call:done", "function_call_output:done"] if verdict == "pass" else [] - ) - assert ( - _shape([item.to_input_item() for item in caught.value.run_data.new_items]) - == expected + ["function_call:done", "function_call_output:done"] + if verdict in ("pass", "tool_cancel") + else [] ) + if isinstance(caught.value, UserError): + assert caught.value.run_data is not None + assert ( + _shape([item.to_input_item() for item in caught.value.run_data.new_items]) + == expected + ) assert guardrail_exited.is_set() - expected = ["function_call:done", "function_call_output:done"] if verdict == "pass" else [] + expected = ( + ["function_call:done", "function_call_output:done"] + if verdict in ("pass", "tool_cancel") + else [] + ) assert _shape(await session.get_items()) == ["user", *expected] finally: release.set() From c727ff53cdce68c5b17f1552dc2f8e4959636768 Mon Sep 17 00:00:00 2001 From: Justin Beckwith Date: Wed, 7 Oct 2026 11:20:55 -0700 Subject: [PATCH 7/8] fix: retain replay ownership in partial tool history --- src/agents/run_internal/turn_resolution.py | 48 +++++- tests/test_tool_batch_failure_history.py | 20 ++- tests/test_tool_failure_boundaries.py | 189 +++++++++++++++++++++ 3 files changed, 247 insertions(+), 10 deletions(-) diff --git a/src/agents/run_internal/turn_resolution.py b/src/agents/run_internal/turn_resolution.py index db5d184f19..2dc4cae675 100644 --- a/src/agents/run_internal/turn_resolution.py +++ b/src/agents/run_internal/turn_resolution.py @@ -140,9 +140,11 @@ REJECTION_MESSAGE, NestedHistoryOwnedItem, apply_patch_rejection_item, + drop_orphan_function_calls, extract_mcp_request_id_from_run, function_rejection_item, order_current_turn_tool_outputs, + run_item_to_input_item, shell_rejection_item, ) from .run_steps import ( @@ -803,19 +805,34 @@ async def check_for_final_output_from_tools( raise UserError(f"Invalid tool_use_behavior: {agent.tool_use_behavior}") -def _completed_tool_step_items(model_items: list[RunItem], outputs: list[RunItem]) -> list[RunItem]: +def _completed_tool_step_items( + pre_step_items: list[RunItem], model_items: list[RunItem], outputs: list[RunItem] +) -> list[RunItem]: """Keep accepted call/output pairs in model order and their preceding reasoning.""" outputs_by_call_id = { extract_tool_call_id(item.raw_item): item for item in outputs if not (isinstance(item, ToolCallOutputItem) and item._custom_data_pending) } + finalized_output_ids = {id(item) for item in outputs_by_call_id.values()} + model_output_ids = {id(item) for item in model_items if isinstance(item, ToolCallOutputItem)} retained: list[RunItem] = [] ordered_outputs: list[RunItem] = [] reasoning: list[RunItem] = [] for item in model_items: if isinstance(item, ReasoningItem): reasoning.append(item) + continue + if ( + isinstance(item, (ToolSearchCallItem, ToolSearchOutputItem)) + and get_mapping_or_attr(item.raw_item, "execution") == "server" + and get_mapping_or_attr(item.raw_item, "status") == "completed" + ): + retained.extend(reasoning) + retained.append(item) + elif isinstance(item, ToolCallOutputItem) and id(item) in finalized_output_ids: + retained.extend(reasoning) + retained.append(item) elif isinstance(item, ToolCallItem): output = outputs_by_call_id.get(extract_tool_call_id(item.raw_item)) provider_completed = ( @@ -831,13 +848,30 @@ def _completed_tool_step_items(model_items: list[RunItem], outputs: list[RunItem ) and item.raw_item.status == "completed" ) - if output is not None or provider_completed: + if output is not None or provider_completed or isinstance(item.raw_item, Program): retained.extend(reasoning) - reasoning.clear() retained.append(item) - if output is not None: + if output is not None and id(output) not in model_output_ids: ordered_outputs.append(output) - return [*retained, *ordered_outputs] + # Reasoning belongs to the next model item, even when that item is omitted. + reasoning.clear() + candidates = [*retained, *ordered_outputs] + inputs = [item.to_input_item() for item in candidates] + prior_inputs = [ + payload for item in pre_step_items if (payload := run_item_to_input_item(item)) is not None + ] + replayable_inputs = { + id(item) + for item in drop_orphan_function_calls( + [*prior_inputs, *inputs], + output_pruning_indexes=set(range(len(prior_inputs), len(prior_inputs) + len(inputs))), + ) + } + return [ + item + for item, payload in zip(candidates, inputs, strict=True) + if id(payload) in replayable_inputs + ] async def execute_tools_and_side_effects( @@ -934,7 +968,9 @@ def _commit_accepted_response_tool_output(item: RunItem) -> None: def _publish_completed_tools() -> None: # Accepted server responses already have their own resumable checkpoint. if on_tool_execution_error is not None and not server_manages_conversation: - retained_items = _completed_tool_step_items(new_step_items, completed_outputs) + retained_items = _completed_tool_step_items( + pre_step_items, new_step_items, completed_outputs + ) if retained_items: on_tool_execution_error( SingleStepResult( diff --git a/tests/test_tool_batch_failure_history.py b/tests/test_tool_batch_failure_history.py index 3327887bde..6bc1306562 100644 --- a/tests/test_tool_batch_failure_history.py +++ b/tests/test_tool_batch_failure_history.py @@ -91,6 +91,7 @@ async def on_tool_end(self, context, agent, tool, result): [ ResponseReasoningItem(id="rs_before", type="reasoning", summary=[]), function_call("send_email", {}, call_id="email"), + ResponseReasoningItem(id="rs_ticket", type="reasoning", summary=[]), function_call("create_ticket", {}, call_id="ticket"), ResponseReasoningItem(id="rs_after", type="reasoning", summary=[]), ], @@ -125,11 +126,11 @@ async def on_tool_end(self, context, agent, tool, result): else: await Runner.run(agent, "go", session=session, hooks=Hooks()) assert effects == ["ticket"] - expected = ["reasoning"] + expected = [] if failure == "hook": # An end hook runs after the output has passed its tool guardrails. - expected += ["function_call:email"] - expected += ["function_call:ticket"] + expected += ["reasoning", "function_call:email"] + expected += ["reasoning", "function_call:ticket"] if failure == "hook": expected += ["function_call_output:email"] expected += ["function_call_output:ticket"] @@ -146,7 +147,10 @@ async def on_tool_end(self, context, agent, tool, result): assert _shape([i.to_input_item() for i in caught.value.run_data.new_items]) == expected history = await session.get_items() assert _shape(history) == ["user", *expected] - assert history[1]["id"] == "rs_before" + expected_reasoning = ["rs_before", "rs_ticket"] if failure == "hook" else ["rs_ticket"] + assert [ + item["id"] for item in history if item.get("type") == "reasoning" + ] == expected_reasoning assert history[-1]["output"] == "ticket T-1" if result is not None: assert _shape(output_events) == [ @@ -164,11 +168,19 @@ async def on_tool_end(self, context, agent, tool, result): assert replay_model.last_call is not None assert replay_model.last_call.model_settings.tool_choice is None assert _shape(replay_model.last_call.input) == ["user", *expected] + assert [ + item["id"] + for item in replay_model.last_call.input + if item.get("type") == "reasoning" + ] == expected_reasoning assert effects == ["ticket"] agent.model = model await Runner.run(agent, "finish", session=session) assert model.last_call is not None assert _shape(model.last_call.input) == ["user", *expected, "user"] + assert [ + item["id"] for item in model.last_call.input if item.get("type") == "reasoning" + ] == expected_reasoning assert effects == ["ticket"] finally: session.close() diff --git a/tests/test_tool_failure_boundaries.py b/tests/test_tool_failure_boundaries.py index 1333b3a34f..85a9894e92 100644 --- a/tests/test_tool_failure_boundaries.py +++ b/tests/test_tool_failure_boundaries.py @@ -256,3 +256,192 @@ async def consume(): ) finally: session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("scenario", ["completed", "anonymous", "incomplete"]) +async def test_completed_provider_tool_search_survives_local_failure(streaming, scenario): + from openai.types.responses import ResponseToolSearchCall, ResponseToolSearchOutputItem + + @tool(failure_error_function=None) + async def fail() -> str: + raise ValueError("synthetic sibling failure") + + call_id = None if scenario == "anonymous" else "search" + search_call = ResponseToolSearchCall( + id="search_call", + call_id=call_id, + type="tool_search_call", + execution="server", + status="completed", + arguments={"query": "synthetic lookup"}, + ) + search_output = ResponseToolSearchOutputItem( + id="search_output", + call_id=call_id, + type="tool_search_output", + execution="server", + status="incomplete" if scenario == "incomplete" else "completed", + tools=[ + { + "type": "function", + "name": "lookup", + "parameters": {"type": "object", "properties": {}}, + } + ], + created_by="synthetic_provider", + ) + calls = [search_call, search_output, function_call("fail", {}, call_id="fail")] + model = ScriptedModel( + [get_exact_output_stream_step(calls) if streaming else calls, [assistant_message("done")]] + ) + agent = Agent(name="search", model=model, tools=[fail]) + session = SQLiteSession("search-failure") + result = None + events = [] + expected = ( + [] + if scenario == "incomplete" + else [ + search_call.model_dump(exclude_unset=True), + { + key: value + for key, value in search_output.model_dump(exclude_unset=True).items() + if key != "created_by" + }, + ] + ) + try: + with pytest.raises(UserError, match="synthetic sibling failure") as caught: + if streaming: + result = Runner.run_streamed(agent, "go", session=session) + async for event in result.stream_events(): + if event.type == "run_item_stream_event" and event.name in ( + "tool_search_called", + "tool_search_output_created", + ): + events.append(event.name) + else: + await Runner.run(agent, "go", session=session) + assert caught.value.run_data is not None + assert [item.to_input_item() for item in caught.value.run_data.new_items] == expected + assert (await session.get_items())[1:] == expected + if result is not None: + # Provider events already emitted before local execution must not be emitted twice. + assert events == ["tool_search_called", "tool_search_output_created"] + assert result.to_input_list()[1:] == expected + restored = await RunState.from_json(agent, result.to_state().to_json()) + assert [item.to_input_item() for item in restored._generated_items] == expected + await Runner.run(agent, restored) + assert model.calls[-1].input[1:] == expected + else: + await Runner.run(agent, "continue", session=session) + assert model.calls[-1].input[1:-1] == expected + finally: + session.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("completed_on_followup", [False, True]) +async def test_completed_program_child_retains_parent_after_sibling_failure( + streaming, completed_on_followup +): + from openai.types.responses import ResponseFunctionToolCall, ResponseReasoningItem + from openai.types.responses.response_function_tool_call import CallerProgram + from openai.types.responses.response_output_item import Program, ProgramOutput + + from agents import ProgrammaticToolCallingTool + + completed = asyncio.Event() + effects = [] + + @tool(allowed_callers=["programmatic"]) + async def lookup() -> str: + effects.append("lookup") + return "found" + + @tool(failure_error_function=None) + async def fail() -> str: + await completed.wait() + raise ValueError("synthetic sibling failure") + + class Hooks(RunHooks): + async def on_tool_end(self, context, agent, tool, result): + completed.set() + + reasoning = ResponseReasoningItem(id="program_reasoning", type="reasoning", summary=[]) + program = Program( + id="program_item", + call_id="program", + code="lookup()", + fingerprint="synthetic", + type="program", + ) + caller = CallerProgram(type="program", caller_id="program") + child = ResponseFunctionToolCall( + id="child", + call_id="lookup", + name="lookup", + arguments="{}", + caller=caller, + type="function_call", + ) + program_output = ProgramOutput( + id="program_output", + call_id="program", + result="found", + status="completed", + type="program_output", + ) + fail_call = function_call("fail", {}, call_id="fail") + first = [reasoning, program, child] + failure_step = [program_output, fail_call] if completed_on_followup else [*first, fail_call] + steps = [first, failure_step] if completed_on_followup else [failure_step] + continuation = ( + [assistant_message("done")] + if completed_on_followup + else [program_output, assistant_message("done")] + ) + model = ScriptedModel( + [ + *(get_exact_output_stream_step(step) if streaming else step for step in steps), + continuation, + ] + ) + agent = Agent(name="program", model=model, tools=[ProgrammaticToolCallingTool(), lookup, fail]) + session = SQLiteSession("program-failure") + result = None + try: + with pytest.raises(UserError, match="synthetic sibling failure") as caught: + if streaming: + result = Runner.run_streamed(agent, "go", session=session, hooks=Hooks()) + async for _ in result.stream_events(): + pass + else: + await Runner.run(agent, "go", session=session, hooks=Hooks()) + expected = [ + reasoning.model_dump(exclude_unset=True), + program.model_dump(exclude_unset=True), + child.model_dump(exclude_unset=True), + { + "type": "function_call_output", + "call_id": "lookup", + "output": "found", + "caller": caller.model_dump(exclude_unset=True), + }, + ] + if completed_on_followup: + expected.append(program_output.model_dump(exclude_unset=True)) + assert caught.value.run_data is not None + assert [item.to_input_item() for item in caught.value.run_data.new_items] == expected + assert (await session.get_items())[1:] == expected + if result is not None: + restored = await RunState.from_json(agent, result.to_state().to_json()) + continued = await Runner.run(agent, restored) + assert model.calls[-1].input[1:] == expected + assert continued.final_output == "done" + assert effects == ["lookup"] + finally: + session.close() From 57a2c7122ef193b671e093ab3eeefd26875c9f98 Mon Sep 17 00:00:00 2001 From: Justin Beckwith Date: Wed, 7 Oct 2026 13:12:23 -0700 Subject: [PATCH 8/8] fix: preserve accepted outputs while metadata is pending --- src/agents/items.py | 5 +- src/agents/run_internal/tool_actions.py | 12 +-- src/agents/run_internal/tool_execution.py | 4 - src/agents/run_internal/turn_resolution.py | 94 ++++++++++++------ tests/test_native_tool_failure_history.py | 13 ++- tests/test_tool_batch_failure_history.py | 105 +++++++++++++++------ 6 files changed, 154 insertions(+), 79 deletions(-) diff --git a/src/agents/items.py b/src/agents/items.py index 7a350e62ad..4e8c600709 100644 --- a/src/agents/items.py +++ b/src/agents/items.py @@ -448,12 +448,9 @@ class ToolCallOutputItem(RunItemBase[Any]): """SDK-only custom data attached to this tool output. This data is not part of ``raw_item`` and is not sent back to the model when the output item is - replayed as input. + replayed as input. On a failed run, unfinished custom-data extraction may leave this unset. """ - _custom_data_pending: bool = field(default=False, init=False, repr=False, compare=False) - """Live finalization state; excluded from fresh partial history, not serialized to RunState.""" - @property def call_id(self) -> str | None: """Return the call identifier from the raw item, if available.""" diff --git a/src/agents/run_internal/tool_actions.py b/src/agents/run_internal/tool_actions.py index 3b938bd27b..21a97af88b 100644 --- a/src/agents/run_internal/tool_actions.py +++ b/src/agents/run_internal/tool_actions.py @@ -181,8 +181,7 @@ async def _run_action(span: Any | None) -> RunItem: raw_item=raw_item, ) - # Preserve the pre-metadata recovery checkpoint for approved/server-owned calls. - output_item._custom_data_pending = True + # Record accepted output before optional metadata and end hooks. if tool_output_committer is not None: tool_output_committer(output_item) @@ -198,7 +197,6 @@ async def finalize_output() -> None: ), ) output_item.custom_data = custom_data - output_item._custom_data_pending = False await gather_with_cancel( hooks.on_tool_end(context_wrapper, agent, action.computer_tool, output), @@ -829,8 +827,7 @@ async def _run_call(span: Any | None) -> RunItem: raw_item=raw_item, ) - # Preserve the pre-metadata recovery checkpoint for approved/server-owned calls. - output_item._custom_data_pending = True + # Record accepted output before optional metadata and end hooks. if tool_output_committer is not None: tool_output_committer(output_item) @@ -846,7 +843,6 @@ async def finalize_output() -> None: ), ) output_item.custom_data = custom_data - output_item._custom_data_pending = False await gather_with_cancel( hooks.on_tool_end(tool_context, agent, custom_tool, output_text), @@ -1074,8 +1070,7 @@ async def _run_call(span: Any | None) -> RunItem: raw_item=raw_item, ) - # Preserve the pre-metadata recovery checkpoint for approved/server-owned calls. - output_item._custom_data_pending = True + # Record accepted output before optional metadata and end hooks. if tool_output_committer is not None: tool_output_committer(output_item) @@ -1092,7 +1087,6 @@ async def finalize_output() -> None: ), ) output_item.custom_data = custom_data - output_item._custom_data_pending = False await gather_with_cancel( hooks.on_tool_end(context_wrapper, agent, apply_patch_tool, output_text), diff --git a/src/agents/run_internal/tool_execution.py b/src/agents/run_internal/tool_execution.py index ff1c6f6020..7f8e86ce54 100644 --- a/src/agents/run_internal/tool_execution.py +++ b/src/agents/run_internal/tool_execution.py @@ -2171,7 +2171,6 @@ async def _invoke_tool_and_run_post_invoke( agent=self.public_agent, tool_origin=get_function_tool_origin(func_tool), ) - output_item._custom_data_pending = True self.output_items_by_tool_run[id(task_state.tool_run)] = output_item if self.tool_output_committer is not None: self.tool_output_committer(output_item) @@ -2189,9 +2188,6 @@ async def _invoke_tool_and_run_post_invoke( self.custom_data_by_tool_run[id(task_state.tool_run)] = custom_data if output_item is not None: output_item.custom_data = custom_data - if output_item is not None: - output_item._custom_data_pending = False - await gather_with_cancel( self.hooks.on_tool_end(tool_context, self.public_agent, func_tool, final_result), ( diff --git a/src/agents/run_internal/turn_resolution.py b/src/agents/run_internal/turn_resolution.py index 2dc4cae675..af13faf1e9 100644 --- a/src/agents/run_internal/turn_resolution.py +++ b/src/agents/run_internal/turn_resolution.py @@ -3,7 +3,7 @@ import asyncio import inspect from collections.abc import Awaitable, Callable, Container, Mapping, Sequence -from copy import deepcopy +from copy import copy, deepcopy from dataclasses import replace from typing import Any, Literal, cast @@ -805,16 +805,30 @@ async def check_for_final_output_from_tools( raise UserError(f"Invalid tool_use_behavior: {agent.tool_use_behavior}") +def _snapshot_tool_outputs(items: list[RunItem], live_outputs: list[RunItem]) -> list[RunItem]: + """Detach failed-run replay data from finalizers that may still be running.""" + # Provider-owned and prior-turn items retain their established response identities. + live_output_ids = {id(item) for item in live_outputs} + snapshots: list[RunItem] = [] + for item in items: + if isinstance(item, ToolCallOutputItem) and id(item) in live_output_ids: + # A shallow copy preserves occurrence markers and opaque application output. + # Only provider payloads and SDK metadata belong to this recovery snapshot. + snapshot = copy(item) + snapshot.raw_item = deepcopy(item.raw_item) + snapshot.custom_data = deepcopy(item.custom_data) + snapshots.append(snapshot) + else: + snapshots.append(item) + return snapshots + + def _completed_tool_step_items( pre_step_items: list[RunItem], model_items: list[RunItem], outputs: list[RunItem] ) -> list[RunItem]: """Keep accepted call/output pairs in model order and their preceding reasoning.""" - outputs_by_call_id = { - extract_tool_call_id(item.raw_item): item - for item in outputs - if not (isinstance(item, ToolCallOutputItem) and item._custom_data_pending) - } - finalized_output_ids = {id(item) for item in outputs_by_call_id.values()} + outputs_by_call_id = {extract_tool_call_id(item.raw_item): item for item in outputs} + accepted_output_ids = {id(item) for item in outputs_by_call_id.values()} model_output_ids = {id(item) for item in model_items if isinstance(item, ToolCallOutputItem)} retained: list[RunItem] = [] ordered_outputs: list[RunItem] = [] @@ -830,7 +844,7 @@ def _completed_tool_step_items( ): retained.extend(reasoning) retained.append(item) - elif isinstance(item, ToolCallOutputItem) and id(item) in finalized_output_ids: + elif isinstance(item, ToolCallOutputItem) and id(item) in accepted_output_ids: retained.extend(reasoning) retained.append(item) elif isinstance(item, ToolCallItem): @@ -943,7 +957,11 @@ async def execute_tools_and_side_effects( prior_input_results = list(run_state._tool_input_guardrail_results) if run_state else [] prior_output_results = list(run_state._tool_output_guardrail_results) if run_state else [] + accepting_outputs = True + def _commit_accepted_response_tool_output(item: RunItem) -> None: + if not accepting_outputs: + return completed_outputs.append(item) if run_state is None or not isinstance(run_state._current_step, NextStepInterruption): return @@ -966,10 +984,14 @@ def _commit_accepted_response_tool_output(item: RunItem) -> None: ] def _publish_completed_tools() -> None: - # Accepted server responses already have their own resumable checkpoint. + nonlocal accepting_outputs + accepting_outputs = False + # Server-managed continuation retains its separate recovery path. if on_tool_execution_error is not None and not server_manages_conversation: - retained_items = _completed_tool_step_items( - pre_step_items, new_step_items, completed_outputs + model_item_ids = {id(item) for item in new_step_items} + retained_items = _snapshot_tool_outputs( + _completed_tool_step_items(pre_step_items, new_step_items, completed_outputs), + [item for item in completed_outputs if id(item) not in model_item_ids], ) if retained_items: on_tool_execution_error( @@ -2710,7 +2732,11 @@ def _rebind_function_run( prior_input_results = list(run_state._tool_input_guardrail_results) if run_state else [] prior_output_results = list(run_state._tool_output_guardrail_results) if run_state else [] + accepting_outputs = True + def _commit_tool_output(item: RunItem) -> None: + if not accepting_outputs: + return if any(existing is item for existing in committed_tool_outputs): return committed_tool_outputs.append(item) @@ -2754,25 +2780,33 @@ def _commit_tool_output(item: RunItem) -> None: ) _register_tool_call_items(context_wrapper, [item]) - ( - function_results, - tool_input_guardrail_results, - tool_output_guardrail_results, - computer_results, - custom_tool_results, - shell_results, - apply_patch_results, - _local_shell_results, - ) = await _execute_tool_plan( - plan=plan, - bindings=bindings, - hooks=hooks, - context_wrapper=context_wrapper, - run_config=run_config, - tool_output_committer=_commit_tool_output, - tool_input_guardrail_results=tool_input_guardrail_results, - tool_output_guardrail_results=tool_output_guardrail_results, - ) + try: + ( + function_results, + tool_input_guardrail_results, + tool_output_guardrail_results, + computer_results, + custom_tool_results, + shell_results, + apply_patch_results, + _local_shell_results, + ) = await _execute_tool_plan( + plan=plan, + bindings=bindings, + hooks=hooks, + context_wrapper=context_wrapper, + run_config=run_config, + tool_output_committer=_commit_tool_output, + tool_input_guardrail_results=tool_input_guardrail_results, + tool_output_guardrail_results=tool_output_guardrail_results, + ) + except BaseException: + accepting_outputs = False + if run_state is not None and not server_manages_conversation: + run_state._generated_items = _snapshot_tool_outputs( + run_state._generated_items, committed_tool_outputs + ) + raise for interruption in _collect_tool_interruptions( function_results=function_results, diff --git a/tests/test_native_tool_failure_history.py b/tests/test_native_tool_failure_history.py index d2eeb6a9ba..68953fe604 100644 --- a/tests/test_native_tool_failure_history.py +++ b/tests/test_native_tool_failure_history.py @@ -210,9 +210,7 @@ async def run(): @pytest.mark.asyncio @pytest.mark.parametrize("stop", ["parent_cancel", "drain_timeout"]) -async def test_native_finalization_does_not_delay_parent_or_publish_unfinished_output( - monkeypatch, stop -): +async def test_native_finalization_preserves_output_without_delaying_parent(monkeypatch, stop): from agents.run_internal import tool_execution entered = asyncio.Event() @@ -259,9 +257,14 @@ async def sibling() -> str: assert exited.is_set() assert not release.is_set() assert caught.value.run_data is not None - assert caught.value.run_data.new_items == [] + assert caught.value.run_data.new_items[-1].raw_item["output"] == "completed" + assert caught.value.run_data.new_items[-1].custom_data is None assert effects == ["custom"] - assert [item.get("role") for item in await session.get_items()] == ["user"] + saved = await session.get_items() + if stop == "parent_cancel": + assert [item.get("role") for item in saved] == ["user"] + else: + assert saved[-1]["output"] == "completed" finally: release.set() if not task.done(): diff --git a/tests/test_tool_batch_failure_history.py b/tests/test_tool_batch_failure_history.py index 6bc1306562..47665a2661 100644 --- a/tests/test_tool_batch_failure_history.py +++ b/tests/test_tool_batch_failure_history.py @@ -546,10 +546,11 @@ async def consume(): @pytest.mark.asyncio @pytest.mark.parametrize("streaming", [False, True]) -async def test_function_partial_history_excludes_pending_custom_data(monkeypatch, streaming): - from agents.run_internal import tool_execution - +@pytest.mark.parametrize("kind", ["function", "native"]) +@pytest.mark.parametrize("approved", [False, True]) +async def test_partial_history_freezes_accepted_output_before_metadata(streaming, kind, approved): from .model_test_helpers import get_exact_output_stream_step + from .test_native_tool_failure_history import _native_tool extracting = asyncio.Event() release = asyncio.Event() @@ -559,35 +560,58 @@ async def test_function_partial_history_excludes_pending_custom_data(monkeypatch async def extract(context): extracting.set() try: - await release.wait() + # Deliberately finish after cancellation and failed-run publication. + while not release.is_set(): + try: + await release.wait() + except asyncio.CancelledError: + pass + context.raw_item["output"] = "extractor-local mutation" return {"ticket": "T-1"} finally: settled.set() - @tool(custom_data_extractor=extract) + @tool(custom_data_extractor=extract, needs_approval=approved) async def create_ticket() -> str: effects.append("ticket") return "created" - @tool(failure_error_function=None) + @tool(failure_error_function=None, needs_approval=approved) async def fail() -> str: await extracting.wait() raise ValueError("synthetic sibling failure") - monkeypatch.setattr(tool_execution, "_FUNCTION_TOOL_POST_INVOKE_WAIT_SECONDS", 0.001) - calls = [ - function_call("create_ticket", {}, call_id="ticket"), - function_call("fail", {}, call_id="fail"), - ] - model = ScriptedModel([get_exact_output_stream_step(calls) if streaming else calls]) - agent = Agent(name="metadata", model=model, tools=[create_ticket, fail]) - session = SQLiteSession("pending-function-metadata") + if kind == "native": + completed_tool, completed_call = _native_tool("custom", extract, effects) + completed_tool.needs_approval = approved + expected_output = "completed" + else: + completed_tool = create_ticket + completed_call = function_call("create_ticket", {}, call_id="ticket") + expected_output = "created" + calls = [completed_call, function_call("fail", {}, call_id="fail")] + model = ScriptedModel( + [ + get_exact_output_stream_step(calls) if streaming and not approved else calls, + [assistant_message("done")], + ] + ) + agent = Agent(name="metadata", model=model, tools=[completed_tool, fail]) + session = SQLiteSession("pending-metadata") result = None + state = None outputs = [] + run_input = "go" try: - with pytest.raises(UserError, match="synthetic sibling failure") as caught: + if approved: + interrupted = await Runner.run(agent, run_input, session=session) + state = interrupted.to_state() + for item in interrupted.interruptions: + state.approve(item) + run_input = state + with pytest.raises((UserError, ValueError), match="synthetic sibling failure") as caught: if streaming: - result = Runner.run_streamed(agent, "go", session=session) + result = Runner.run_streamed(agent, run_input, session=session) async for event in result.stream_events(): if ( event.type == "run_item_stream_event" @@ -595,22 +619,49 @@ async def fail() -> str: ): outputs.append(event.item) else: - await Runner.run(agent, "go", session=session) - assert effects == ["ticket"] + await Runner.run(agent, run_input, session=session) + assert len(effects) == 1 assert not settled.is_set() - assert caught.value.run_data is not None - assert caught.value.run_data.new_items == [] - assert _shape(await session.get_items()) == ["user"] - if result is not None: - assert outputs == [] - restored = await RunState.from_json(agent, result.to_state().to_json()) - assert restored._generated_items == [] + if state is not None: + items = state._generated_items + else: + assert caught.value.run_data is not None + items = caught.value.run_data.new_items + accepted = [item for item in items if item.type == "tool_call_output_item"] + assert len(accepted) == 1 + assert accepted[0].raw_item["output"] == expected_output + assert accepted[0].custom_data is None + if not approved: + saved = await session.get_items() + assert saved[-1]["output"] == expected_output + if result is not None: + assert len(outputs) == 1 + state = result.to_state() + serialized = state.to_json() if state is not None else None release.set() await asyncio.wait_for(settled.wait(), 2) - assert caught.value.run_data.new_items == [] + # Let the extractor's owner perform its assignment after extract() returns. + await asyncio.sleep(0) + assert accepted[0].custom_data is None + assert accepted[0].raw_item["output"] == expected_output + if state is not None: + assert state.to_json() == serialized + restored = await RunState.from_json(agent, serialized) + restored_output = next( + item for item in restored._generated_items if item.type == "tool_call_output_item" + ) + assert restored_output.custom_data is None + if not approved: + assert await session.get_items() == saved + if outputs: + assert outputs[0].custom_data is None + continuation = await Runner.run(agent, "continue", session=session) + assert continuation.final_output == "done" + assert len(effects) == 1 finally: release.set() - await asyncio.wait_for(settled.wait(), 2) + if extracting.is_set(): + await asyncio.wait_for(settled.wait(), 2) session.close()