From 211fe74d4446454e461ba6df2531d298107d812e Mon Sep 17 00:00:00 2001 From: Justin Beckwith Date: Mon, 28 Sep 2026 09:48:11 -0700 Subject: [PATCH 1/2] fix: preserve streamed cancellation and pending guardrail cleanup --- src/agents/result.py | 16 ++- src/agents/run_internal/guardrails.py | 9 +- tests/test_cancel_streaming.py | 134 +++++++++++++++++++++++++- 3 files changed, 148 insertions(+), 11 deletions(-) diff --git a/src/agents/result.py b/src/agents/result.py index 70d48fe3ef..3dd36a6837 100644 --- a/src/agents/result.py +++ b/src/agents/result.py @@ -1166,20 +1166,18 @@ def __str__(self) -> str: return pretty_print_run_result_streaming(self) async def _await_task_safely(self, task: asyncio.Task[Any] | None) -> None: - """Await a task if present, ignoring cancellation and storing exceptions elsewhere. + """Wait for a task, leaving its cancellation and exceptions to _check_errors(). - This ensures we do not lose late guardrail exceptions while not surfacing - CancelledError to callers of stream_events. + This ensures we do not lose late guardrail exceptions while not surfacing the task's + CancelledError to callers of stream_events. Consumer cancellation still propagates. """ if task and not task.done(): try: - await task + # Waiting directly would conflate child cancellation with consumer cancellation. + await asyncio.wait((task,)) except asyncio.CancelledError: - # Task was cancelled (e.g., due to result.cancel()). Nothing to do here. - pass - except Exception: - # The exception will be surfaced via _check_errors() if needed. - pass + task.cancel() + raise def _drain_event_queue(self) -> None: """Remove any pending items from the event queue and mark them done.""" diff --git a/src/agents/run_internal/guardrails.py b/src/agents/run_internal/guardrails.py index dc09f0cd7a..5e5910b4ab 100644 --- a/src/agents/run_internal/guardrails.py +++ b/src/agents/run_internal/guardrails.py @@ -233,7 +233,14 @@ async def input_guardrail_tripwire_triggered_for_stream( return False if not task.done(): - await task + # A cancelled child must not abort the caller's remaining run cleanup. + try: + await asyncio.wait((task,)) + except asyncio.CancelledError: + task.cancel() + raise + if not task.cancelled(): + task.result() return any( guardrail_result.output.tripwire_triggered diff --git a/tests/test_cancel_streaming.py b/tests/test_cancel_streaming.py index ea7cb85635..e9076eb505 100644 --- a/tests/test_cancel_streaming.py +++ b/tests/test_cancel_streaming.py @@ -6,13 +6,17 @@ import pytest from openai.types.responses import ResponseCompletedEvent -from agents import Agent, Runner +from agents import Agent, ComputerProvider, ComputerTool, GuardrailFunctionOutput, Runner +from agents.decorators import tool from agents.guardrail import input_guardrail from agents.models.multi_provider import MultiProvider +from agents.result import RunResultStreaming from agents.stream_events import RawResponsesStreamEvent from agents.testing import ScriptedModel +from .test_computer_tool_lifecycle import FakeComputer from .test_responses import get_function_tool, get_function_tool_call, get_text_message +from .testing_processor import fetch_events class SlowCompleteScriptedModel(ScriptedModel): @@ -328,3 +332,131 @@ async def raising_guardrail(context, agent, input): with pytest.raises(FalsyRuntimeError, match="falsy guardrail boom"): async for _ in result.stream_events(): pass + + +@pytest.mark.asyncio +async def test_cancel_with_pending_parallel_input_guardrail_finishes_cleanup() -> None: + guardrail_started = asyncio.Event() + tool_started = asyncio.Event() + disposed: list[FakeComputer] = [] + + @input_guardrail + async def slow_guardrail(context, agent, input): + guardrail_started.set() + await asyncio.Event().wait() + return GuardrailFunctionOutput(output_info=None, tripwire_triggered=False) + + @tool + async def slow_tool() -> str: + await guardrail_started.wait() + tool_started.set() + await asyncio.Event().wait() + return "unreachable" + + computer_tool = ComputerTool( + computer=ComputerProvider[FakeComputer]( + create=lambda *, run_context: FakeComputer(), + dispose=lambda *, run_context, computer: disposed.append(computer), + ) + ) + agent = Agent( + name="A", + model=ScriptedModel([[get_function_tool_call("slow_tool", "{}", "call_1")]]), + tools=[slow_tool, computer_tool], + input_guardrails=[slow_guardrail], + ) + result = Runner.run_streamed(agent, input="hi") + + async def consume() -> None: + async for _ in result.stream_events(): + pass + + consumer = asyncio.create_task(consume()) + try: + await asyncio.wait_for(tool_started.wait(), timeout=2) + result.cancel() + await asyncio.wait_for(consumer, timeout=2) + assert len(disposed) == 1 + events = fetch_events() + assert events.count("trace_start") == events.count("trace_end") == 1 + assert events.count("span_start") == events.count("span_end") + finally: + result.cancel() + assert result.run_loop_task is not None + await asyncio.gather(consumer, result.run_loop_task, return_exceptions=True) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("wait_target", ["guardrail", "run_loop"]) +async def test_consumer_cancellation_at_terminal_wait_propagates( + monkeypatch: pytest.MonkeyPatch, wait_target: str +) -> None: + wait_started = asyncio.Event() + guardrail_started = asyncio.Event() + disposal_started = asyncio.Event() + child_cancelled = asyncio.Event() + + @input_guardrail + async def slow_guardrail(context, agent, input): + guardrail_started.set() + try: + await asyncio.Event().wait() + finally: + child_cancelled.set() + return GuardrailFunctionOutput(output_info=None, tripwire_triggered=False) + + async def dispose(**kwargs) -> None: + disposal_started.set() + try: + await asyncio.Event().wait() + finally: + child_cancelled.set() + + computer_tool = ComputerTool( + computer=ComputerProvider[FakeComputer]( + create=lambda *, run_context: FakeComputer(), dispose=dispose + ) + ) + result = Runner.run_streamed( + Agent( + name="A", + model=ScriptedModel([[get_text_message("done")]]), + input_guardrails=[slow_guardrail] if wait_target == "guardrail" else [], + tools=[computer_tool] if wait_target == "run_loop" else [], + ), + input="hi", + ) + original_wait = RunResultStreaming._await_task_safely + + async def observe_wait(self, task) -> None: + target = self._input_guardrails_task if wait_target == "guardrail" else self.run_loop_task + if self is result and task is target and task is not None and not task.done(): + wait_started.set() + await original_wait(self, task) + + # Observe entry to the terminal wait without changing its behavior. Public events do not + # expose this boundary, and the consumer must already be suspended here before cancellation. + monkeypatch.setattr(RunResultStreaming, "_await_task_safely", observe_wait) + + async def consume() -> None: + async for _ in result.stream_events(): + pass + + consumer = asyncio.create_task(consume()) + try: + child_started = guardrail_started if wait_target == "guardrail" else disposal_started + await asyncio.wait_for(child_started.wait(), timeout=2) + await asyncio.wait_for(wait_started.wait(), timeout=2) + assert result.final_output == "done" + consumer.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(consumer, timeout=2) + await asyncio.wait_for(child_cancelled.wait(), timeout=2) + if wait_target == "guardrail": + events = fetch_events() + assert events.count("trace_start") == events.count("trace_end") == 1 + assert events.count("span_start") == events.count("span_end") + finally: + result.cancel() + assert result.run_loop_task is not None + await asyncio.gather(consumer, result.run_loop_task, return_exceptions=True) From 846e167cfb121c7c004719e570750c0af8d1670e Mon Sep 17 00:00:00 2001 From: Justin Beckwith Date: Mon, 28 Sep 2026 10:12:54 -0700 Subject: [PATCH 2/2] fix: preserve terminal cleanup and guardrail cancellation boundaries --- src/agents/result.py | 40 +++---- src/agents/run_internal/guardrails.py | 6 +- src/agents/run_internal/run_loop.py | 4 +- tests/test_agent_runner_streamed.py | 2 +- tests/test_cancel_streaming.py | 146 +++++++++++++++++++++++++- 5 files changed, 171 insertions(+), 27 deletions(-) diff --git a/src/agents/result.py b/src/agents/result.py index 3dd36a6837..0dad43e5f2 100644 --- a/src/agents/result.py +++ b/src/agents/result.py @@ -1045,25 +1045,25 @@ def register_current_consumer() -> None: registered_consumer_task.remove_done_callback(consumer_task_done) unregister_consumer() try: - if cancelled: - # Cancellation should return promptly, so avoid waiting on long-running tasks. - # Tasks have already been cancelled above. - self._cleanup_tasks() - else: - # Ensure main execution completes before cleanup to avoid race conditions - # with session operations. - await self._await_task_safely(self.run_loop_task) - # Re-check for exceptions now that the run loop has fully settled. - # _await_task_safely swallows exceptions; without this call, a run-loop - # failure that races past the sentinel (e.g. early sandbox failures) would - # be silently lost instead of surfaced via _stored_exception. - self._check_errors() - # Safely terminate all background tasks after main execution has finished. - self._cleanup_tasks() - - if not cancelled: - await self._await_model_provider_cleanup() - await self._run_sandbox_cleanup() + try: + if cancelled: + # Queue-wait cancellation should return without awaiting long-running tasks. + # Tasks have already been cancelled above. + self._cleanup_tasks() + else: + # Ensure main execution completes before cleanup to avoid race conditions + # with session operations. + await self._await_task_safely(self.run_loop_task) + # Re-check for exceptions now that the run loop has fully settled. + # _await_task_safely leaves errors to _check_errors(); without this call, + # a run-loop failure racing past the sentinel would be silently lost. + self._check_errors() + # Safely terminate background tasks after main execution has finished. + self._cleanup_tasks() + finally: + if not cancelled: + await self._await_model_provider_cleanup() + await self._run_sandbox_cleanup() finally: # Allow any pending callbacks (e.g., cancellation handlers) to enqueue their # completion sentinels before we clear the queues for observability. @@ -1177,6 +1177,8 @@ async def _await_task_safely(self, task: asyncio.Task[Any] | None) -> None: await asyncio.wait((task,)) except asyncio.CancelledError: task.cancel() + # Preserve direct-await settlement before the consumer releases owned resources. + await asyncio.wait((task,)) raise def _drain_event_queue(self) -> None: diff --git a/src/agents/run_internal/guardrails.py b/src/agents/run_internal/guardrails.py index 5e5910b4ab..d1a1c24d2c 100644 --- a/src/agents/run_internal/guardrails.py +++ b/src/agents/run_internal/guardrails.py @@ -226,8 +226,10 @@ def record(result: OutputGuardrailResult) -> None: async def input_guardrail_tripwire_triggered_for_stream( streamed_result: RunResultStreaming, + *, + ignore_cancelled: bool = False, ) -> bool: - """Return True if any input guardrail triggered during a streamed run.""" + """Wait for input verdicts, optionally tolerating child cancellation during final cleanup.""" task = streamed_result._input_guardrails_task if task is None: return False @@ -239,7 +241,7 @@ async def input_guardrail_tripwire_triggered_for_stream( except asyncio.CancelledError: task.cancel() raise - if not task.cancelled(): + if not ignore_cancelled or not task.cancelled(): task.result() return any( diff --git a/src/agents/run_internal/run_loop.py b/src/agents/run_internal/run_loop.py index da527bb5f9..bad24a8357 100644 --- a/src/agents/run_internal/run_loop.py +++ b/src/agents/run_internal/run_loop.py @@ -2047,7 +2047,9 @@ def _record_max_turns_handler_output( _sync_conversation_tracking_from_tracker() if streamed_result._input_guardrails_task: try: - triggered = await input_guardrail_tripwire_triggered_for_stream(streamed_result) + triggered = await input_guardrail_tripwire_triggered_for_stream( + streamed_result, ignore_cancelled=True + ) if triggered: first_trigger = next( ( diff --git a/tests/test_agent_runner_streamed.py b/tests/test_agent_runner_streamed.py index e2518e93b6..addef5e659 100644 --- a/tests/test_agent_runner_streamed.py +++ b/tests/test_agent_runner_streamed.py @@ -1580,7 +1580,7 @@ async def safe_guardrail( error = RuntimeError("SECRET_STREAM_FINALIZER_ERROR") - async def fail_finalizer(_result: Any) -> bool: + async def fail_finalizer(_result: Any, *, ignore_cancelled: bool = False) -> bool: raise error monkeypatch.setattr( diff --git a/tests/test_cancel_streaming.py b/tests/test_cancel_streaming.py index e9076eb505..bfe26d20a3 100644 --- a/tests/test_cancel_streaming.py +++ b/tests/test_cancel_streaming.py @@ -11,12 +11,14 @@ from agents.guardrail import input_guardrail from agents.models.multi_provider import MultiProvider from agents.result import RunResultStreaming +from agents.run_internal import run_loop from agents.stream_events import RawResponsesStreamEvent from agents.testing import ScriptedModel from .test_computer_tool_lifecycle import FakeComputer from .test_responses import get_function_tool, get_function_tool_call, get_text_message from .testing_processor import fetch_events +from .utils.simple_session import SimpleListSession class SlowCompleteScriptedModel(ScriptedModel): @@ -353,9 +355,12 @@ async def slow_tool() -> str: await asyncio.Event().wait() return "unreachable" + def create_fake_computer(*, run_context) -> FakeComputer: + return FakeComputer() + computer_tool = ComputerTool( computer=ComputerProvider[FakeComputer]( - create=lambda *, run_context: FakeComputer(), + create=create_fake_computer, dispose=lambda *, run_context, computer: disposed.append(computer), ) ) @@ -412,10 +417,11 @@ async def dispose(**kwargs) -> None: finally: child_cancelled.set() + def create_fake_computer(*, run_context) -> FakeComputer: + return FakeComputer() + computer_tool = ComputerTool( - computer=ComputerProvider[FakeComputer]( - create=lambda *, run_context: FakeComputer(), dispose=dispose - ) + computer=ComputerProvider[FakeComputer](create=create_fake_computer, dispose=dispose) ) result = Runner.run_streamed( Agent( @@ -460,3 +466,135 @@ async def consume() -> None: result.cancel() assert result.run_loop_task is not None await asyncio.gather(consumer, result.run_loop_task, return_exceptions=True) + + +@pytest.mark.asyncio +async def test_pending_guardrail_cancellation_does_not_accept_or_persist_output( + monkeypatch: pytest.MonkeyPatch, +) -> None: + verdict_wait_started = asyncio.Event() + dependency = asyncio.create_task(asyncio.Event().wait()) + session = SimpleListSession() + original_verdict = run_loop.input_guardrail_tripwire_triggered_for_stream + + @input_guardrail + async def guardrail(context, agent, input): + await dependency + return GuardrailFunctionOutput(output_info=None, tripwire_triggered=False) + + async def observe_verdict(result, **kwargs): + verdict_wait_started.set() + return await original_verdict(result, **kwargs) + + monkeypatch.setattr(run_loop, "input_guardrail_tripwire_triggered_for_stream", observe_verdict) + result = Runner.run_streamed( + Agent( + name="A", + model=ScriptedModel([[get_text_message("done")]]), + input_guardrails=[guardrail], + ), + "hi", + session=session, + ) + + async def consume() -> None: + async for _ in result.stream_events(): + pass + + consumer = asyncio.create_task(consume()) + try: + await asyncio.wait_for(verdict_wait_started.wait(), timeout=2) + dependency.cancel() + await asyncio.wait_for(consumer, timeout=2) + assert result.final_output is None + assert await session.get_items() == [{"content": "hi", "role": "user"}] + assert result.run_loop_task is not None + assert result.run_loop_task.cancelled() + events = fetch_events() + assert events.count("trace_start") == events.count("trace_end") == 1 + assert events.count("span_start") == events.count("span_end") + finally: + dependency.cancel() + result.cancel() + assert result.run_loop_task is not None + await asyncio.gather(dependency, consumer, result.run_loop_task, return_exceptions=True) + + +@pytest.mark.asyncio +async def test_terminal_consumer_cancellation_waits_for_registered_cleanup( + monkeypatch: pytest.MonkeyPatch, +) -> None: + wait_started = asyncio.Event() + model_cleanup_started = asyncio.Event() + provider_cleanup_started = asyncio.Event() + cleanup_wait_started = asyncio.Event() + cleanup_release = asyncio.Event() + completed: list[str] = [] + + class CleanupModel(ScriptedModel): + async def _cleanup_on_run_end(self, owner) -> None: + model_cleanup_started.set() + await asyncio.Event().wait() + + async def close_provider(provider: MultiProvider) -> None: + provider_cleanup_started.set() + await cleanup_release.wait() + completed.append("provider") + + async def cleanup_sandbox() -> None: + await cleanup_release.wait() + completed.append("sandbox") + + monkeypatch.setattr(MultiProvider, "aclose", close_provider) + result = Runner.run_streamed( + Agent(name="A", model=CleanupModel([[get_text_message("done")]])), "hi" + ) + await asyncio.wait_for(model_cleanup_started.wait(), timeout=2) + # Register the same cleanup wrapper used by SandboxRuntime without a provider backend. + result._sandbox_cleanup = cleanup_sandbox + result.ensure_sandbox_cleanup_on_completion() + original_wait = RunResultStreaming._await_task_safely + original_provider_wait = RunResultStreaming._await_model_provider_cleanup + + async def observe_wait(self, task) -> None: + if self is result and task is self.run_loop_task: + wait_started.set() + await original_wait(self, task) + + monkeypatch.setattr(RunResultStreaming, "_await_task_safely", observe_wait) + + async def observe_provider_wait(self) -> None: + if self is result: + cleanup_wait_started.set() + await original_provider_wait(self) + + monkeypatch.setattr(RunResultStreaming, "_await_model_provider_cleanup", observe_provider_wait) + + async def consume() -> None: + async for _ in result.stream_events(): + pass + + consumer = asyncio.create_task(consume()) + cleanup_waiter = asyncio.create_task(cleanup_wait_started.wait()) + try: + await asyncio.wait_for(wait_started.wait(), timeout=2) + consumer.cancel() + await asyncio.wait_for(provider_cleanup_started.wait(), timeout=2) + # The consumer must await registered cleanup rather than return while callbacks run. + await asyncio.wait_for( + asyncio.wait((consumer, cleanup_waiter), return_when=asyncio.FIRST_COMPLETED), timeout=2 + ) + assert cleanup_wait_started.is_set() + assert not consumer.done() + cleanup_release.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(consumer, timeout=2) + assert sorted(completed) == ["provider", "sandbox"] + finally: + cleanup_release.set() + cleanup_waiter.cancel() + result.cancel() + assert result.run_loop_task is not None + await asyncio.gather(consumer, cleanup_waiter, result.run_loop_task, return_exceptions=True) + await result._await_model_provider_cleanup() + await result._run_sandbox_cleanup()