Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
62 changes: 61 additions & 1 deletion src/agents/result.py
Original file line number Diff line number Diff line change
Expand Up @@ -693,6 +693,21 @@ class RunResultStreaming(RunResultBase):
)
_sandbox_cleanup_task: asyncio.Task[None] | None = field(default=None, init=False, repr=False)
_sandbox_cleanup_callback_registered: bool = field(default=False, init=False, repr=False)
_sandbox_wrapped_run_loop_task: asyncio.Task[Any] | None = field(
default=None,
init=False,
repr=False,
)
_model_provider_cleanup: Callable[[], Awaitable[None]] | None = field(
default=None,
init=False,
repr=False,
)
_model_provider_cleanup_task: asyncio.Task[None] | None = field(
default=None,
init=False,
repr=False,
)

def __post_init__(self, _run_impl_task: asyncio.Task[Any] | None) -> None:
self._current_agent_ref = weakref.ref(self.current_agent)
Expand Down Expand Up @@ -757,6 +772,7 @@ def ensure_sandbox_cleanup_on_completion(self) -> None:
return

original_task = self.run_loop_task
self._sandbox_wrapped_run_loop_task = original_task
self._sandbox_cleanup_callback_registered = True
original_task.add_done_callback(
lambda _task: asyncio.create_task(self._run_sandbox_cleanup())
Expand Down Expand Up @@ -784,10 +800,50 @@ async def _await_run_and_cleanup() -> Any:
await self._run_sandbox_cleanup()
return result

self.run_loop_task = asyncio.create_task(
cleanup_wrapper_task = asyncio.create_task(
_await_data_redacted_error_boundary(_await_run_and_cleanup)
)

def cancel_original_if_wrapper_cancelled(task: asyncio.Task[Any]) -> None:
# A task cancelled before its first event-loop step cannot enter its coroutine body.
if task.cancelled() and not original_task.done():
original_task.cancel()

cleanup_wrapper_task.add_done_callback(cancel_original_if_wrapper_cancelled)
self.run_loop_task = cleanup_wrapper_task

def _ensure_model_provider_cleanup_on_completion(
self,
cleanup: Callable[[], Awaitable[None]],
) -> None:
"""Register one cleanup task that also starts if the run task never enters its body."""
self._model_provider_cleanup = cleanup
if self.run_loop_task is not None:
self.run_loop_task.add_done_callback(lambda _task: self._start_model_provider_cleanup())

def _start_model_provider_cleanup(self) -> asyncio.Task[None] | None:
task = self._model_provider_cleanup_task
if task is not None:
return task

cleanup = self._model_provider_cleanup
if cleanup is None:
return None

self._model_provider_cleanup = None

async def run_cleanup() -> None:
await cleanup()

task = asyncio.create_task(run_cleanup())
self._model_provider_cleanup_task = task
return task

async def _await_model_provider_cleanup(self) -> None:
task = self._start_model_provider_cleanup()
if task is not None:
await asyncio.shield(task)

@property
def run_loop_exception(self) -> BaseException | None:
"""The exception raised by the background run loop, if any.
Expand Down Expand Up @@ -1003,6 +1059,7 @@ def register_current_consumer() -> None:
self._cleanup_tasks()

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
Expand Down Expand Up @@ -1093,6 +1150,9 @@ def _cleanup_tasks(self):
if self.run_loop_task and not self.run_loop_task.done():
self.run_loop_task.cancel()

if self._sandbox_wrapped_run_loop_task and not self._sandbox_wrapped_run_loop_task.done():
self._sandbox_wrapped_run_loop_task.cancel()

if self._input_guardrails_task and not self._input_guardrails_task.done():
self._input_guardrails_task.cancel()

Expand Down
40 changes: 28 additions & 12 deletions src/agents/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,7 +55,6 @@
ToolExecutionConfig,
ToolNameCollisionPolicy as ToolNameCollisionPolicy,
ToolNotFoundBehavior,
_coerce_run_config,
)
from .run_context import RunContextWrapper, TContext
from .run_error_handlers import RunErrorHandlers
Expand Down Expand Up @@ -108,6 +107,10 @@
normalize_resumed_input,
reconcile_nested_history_owned_input_after_rewrite,
)
from .run_internal.model_provider_lifecycle import (
_close_runner_owned_model_provider,
_normalize_run_config_for_runner,
)
from .run_internal.oai_conversation import OpenAIServerConversationTracker
from .run_internal.prompt_cache_key import PromptCacheKeyResolver
from .run_internal.run_grouping import resolve_run_grouping_id
Expand Down Expand Up @@ -553,18 +556,25 @@ async def run(
input: str | list[TResponseInputItem] | RunState[TContext],
**kwargs: Unpack[RunOptions[TContext]],
) -> RunResult:
run_config, owns_model_provider = _normalize_run_config_for_runner(kwargs.get("run_config"))
cast(dict[str, Any], kwargs)["run_config"] = run_config
redacted_error: BaseException | None = None
try:
return await self._run_impl(starting_agent, input, **kwargs)
except BaseException as error:
if not _is_error_data_redacted(error):
raise
_detach_data_redacted_error_traceback(error)
redacted_error = error
try:
return await self._run_impl(starting_agent, input, **kwargs)
except BaseException as error:
if not _is_error_data_redacted(error):
raise
_detach_data_redacted_error_traceback(error)
redacted_error = error
finally:
if owns_model_provider:
await _close_runner_owned_model_provider(run_config.model_provider)

self = cast(Any, None)
starting_agent = cast(Any, None)
input = cast(Any, None)
run_config = cast(Any, None)
cast(dict[str, Any], kwargs).clear()
assert redacted_error is not None
_detach_data_redacted_error_traceback(redacted_error)
Expand All @@ -586,7 +596,7 @@ async def _run_impl(
conversation_id = kwargs.get("conversation_id")
session = kwargs.get("session")

run_config = RunConfig() if run_config is None else _coerce_run_config(run_config)
run_config = cast(RunConfig, run_config)

is_resumed_state = isinstance(input, RunState)
run_state: RunState[TContext] | None = (
Expand Down Expand Up @@ -2347,7 +2357,7 @@ def run_streamed(
conversation_id = kwargs.get("conversation_id")
session = kwargs.get("session")

run_config = RunConfig() if run_config is None else _coerce_run_config(run_config)
run_config, owns_model_provider = _normalize_run_config_for_runner(run_config)

# Handle RunState input
is_resumed_state = isinstance(input, RunState)
Expand Down Expand Up @@ -2565,8 +2575,8 @@ def run_streamed(
sandbox_runtime.apply_result_metadata(streamed_result)

# Kick off the actual agent loop in the background and return the streamed result object.
streamed_result.run_loop_task = asyncio.create_task(
_await_data_redacted_error_boundary(
async def run_loop() -> None:
await _await_data_redacted_error_boundary(
lambda: start_streaming(
starting_input=input_for_result,
streamed_result=streamed_result,
Expand All @@ -2586,7 +2596,13 @@ def run_streamed(
sandbox_runtime=sandbox_runtime,
)
)
)

streamed_result.run_loop_task = asyncio.create_task(run_loop())
if owns_model_provider:
model_provider = run_config.model_provider
streamed_result._ensure_model_provider_cleanup_on_completion(
lambda: _close_runner_owned_model_provider(model_provider)
)
if sandbox_runtime.enabled:
streamed_result.ensure_sandbox_cleanup_on_completion()
return streamed_result
Expand Down
46 changes: 46 additions & 0 deletions src/agents/run_internal/model_provider_lifecycle.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
from __future__ import annotations

import asyncio
from typing import Any

from ..logger import log_model_action_warning, logger
from ..models.interface import ModelProvider
from ..run_config import RunConfig, _coerce_run_config


def _normalize_run_config_for_runner(
value: RunConfig | dict[str, Any] | None,
) -> tuple[RunConfig, bool]:
"""Normalize a run config and report whether Runner created its model provider."""
owns_model_provider = value is None or (
isinstance(value, dict) and "model_provider" not in value
)
run_config = RunConfig() if value is None else _coerce_run_config(value)
return run_config, owns_model_provider


async def _close_runner_owned_model_provider(model_provider: ModelProvider) -> None:
"""Finish provider cleanup despite repeated cancellation, then restore cancellation."""

async def close() -> None:
try:
await model_provider.aclose()
except Exception as error:
log_model_action_warning(
logger,
"Failed to close model provider created for run",
error,
)

close_task = asyncio.create_task(close())
try:
await asyncio.shield(close_task)
except asyncio.CancelledError:
while not close_task.done():
try:
await asyncio.shield(close_task)
except asyncio.CancelledError:
continue
if not close_task.cancelled():
close_task.result()
raise
20 changes: 19 additions & 1 deletion tests/test_cancel_streaming.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import asyncio
import gc
import json
import time

Expand All @@ -7,6 +8,7 @@

from agents import Agent, Runner
from agents.guardrail import input_guardrail
from agents.models.multi_provider import MultiProvider
from agents.stream_events import RawResponsesStreamEvent
from agents.testing import ScriptedModel

Expand Down Expand Up @@ -109,13 +111,29 @@ async def test_cancel_is_idempotent():


@pytest.mark.asyncio
async def test_cancel_before_streaming():
async def test_cancel_before_streaming(
monkeypatch: pytest.MonkeyPatch,
recwarn: pytest.WarningsRecorder,
) -> None:
closed: list[MultiProvider] = []

async def record_close(provider: MultiProvider) -> None:
closed.append(provider)

monkeypatch.setattr(MultiProvider, "aclose", record_close)
model = ScriptedModel()
agent = Agent(name="Joker", model=model)
result = Runner.run_streamed(agent, input="Please tell me 5 jokes.")
result.cancel() # Cancel before streaming
events = [e async for e in result.stream_events()]
gc.collect()

assert events == [], "No events should be yielded if cancel() is called before streaming."
assert len(closed) == 1
assert not any(
warning.category is RuntimeWarning and "was never awaited" in str(warning.message)
for warning in recwarn
)


@pytest.mark.asyncio
Expand Down
Loading