feat(core): add model call timeouts (#4428)

This commit is contained in:
Kazuhiro Sera
2026-08-16 11:15:36 +09:00
committed by GitHub
parent 2588d154e4
commit b4faf7090c
12 changed files with 1101 additions and 32 deletions
+2
View File
@@ -25,6 +25,7 @@ from .exceptions import (
MCPToolCancellationError,
ModelBehaviorError,
ModelRefusalError,
ModelTimeoutError,
OutputGuardrailTripwireTriggered,
RunErrorDetails,
ToolInputGuardrailTripwireTriggered,
@@ -412,6 +413,7 @@ __all__ = [
"MCPToolCancellationError",
"ModelBehaviorError",
"ModelRefusalError",
"ModelTimeoutError",
"ToolTimeoutError",
"UserError",
"InputGuardrail",
+10
View File
@@ -474,6 +474,16 @@ class ModelRefusalError(AgentsException):
super().__init__(f"Model refused to produce output: {refusal}")
class ModelTimeoutError(AgentsException):
"""Exception raised when a model-call attempt exceeds its configured timeout."""
timeout_seconds: float
def __init__(self, timeout_seconds: float):
self.timeout_seconds = timeout_seconds
super().__init__(f"Model call timed out after {timeout_seconds:g} seconds.")
class UserError(AgentsException):
"""Exception raised when the user makes an error using the SDK."""
+11 -1
View File
@@ -9,7 +9,7 @@ from openai._types import Body, Query
from openai.types.responses import ResponseIncludable
from openai.types.responses.response_create_params import ContextManagement, PromptCacheOptions
from openai.types.shared import Reasoning
from pydantic import GetCoreSchemaHandler, TypeAdapter
from pydantic import Field, FiniteFloat, GetCoreSchemaHandler, TypeAdapter
from pydantic.dataclasses import dataclass
from pydantic_core import core_schema
@@ -81,6 +81,7 @@ _TRACEABLE_MODEL_SETTING_FIELDS = (
"retry",
"context_management",
"prompt_cache_options",
"timeout",
)
@@ -211,6 +212,14 @@ class ModelSettings:
usage from the provider; use ``include_usage`` separately when a streaming provider requires it.
"""
timeout: Annotated[FiniteFloat, Field(gt=0)] | None = None
"""Maximum duration in seconds for each model-call attempt.
The timeout is enforced cooperatively through normal asyncio cancellation. It bounds the
complete model attempt, including transport waits, but does not replace provider-specific
phase timeout configuration or bound the full run, tool calls, or retry backoff.
"""
if TYPE_CHECKING:
def __init__(
@@ -239,6 +248,7 @@ class ModelSettings:
context_management: list[ContextManagement] | None = None,
prompt_cache_options: PromptCacheOptions | None = None,
preserve_raw_usage: bool | None = None,
timeout: Annotated[FiniteFloat, Field(gt=0)] | None = None,
) -> None: ...
def resolve(self, override: ModelSettings | dict[str, Any] | None) -> ModelSettings:
+18 -1
View File
@@ -84,7 +84,10 @@ from ..usage import (
_response_usage_to_usage,
model_usage_to_span_usage,
)
from ..util._error_tracing import record_model_error_on_span
from ..util._error_tracing import (
record_current_task_model_timeout_on_span,
record_model_error_on_span,
)
from ..util._json import _to_dump_compatible
from ..version import __version__
from ._openai_retry import get_openai_retry_advice
@@ -597,6 +600,13 @@ class OpenAIResponsesModel(Model):
if tracing.include_data():
span_response.span_data.response = response
span_response.span_data.input = input
except asyncio.CancelledError:
record_current_task_model_timeout_on_span(
span_response,
message="Error getting response",
trace_include_sensitive_data=tracing.include_data(),
)
raise
except Exception as e:
span_response.set_error(
SpanError(
@@ -732,6 +742,13 @@ class OpenAIResponsesModel(Model):
if final_response is not None and tracing.include_data():
span_response.span_data.response = final_response
span_response.span_data.input = input
except asyncio.CancelledError:
record_current_task_model_timeout_on_span(
span_response,
message="Error streaming response",
trace_include_sensitive_data=tracing.include_data(),
)
raise
except Exception as e:
span_response.set_error(
SpanError(
+223 -28
View File
@@ -4,12 +4,13 @@ import asyncio
import random
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Mapping
from inspect import isawaitable
from typing import Any
from typing import Any, TypeVar
import httpx2
from openai import APIConnectionError, APITimeoutError, BadRequestError
from .._httpx_compat import is_legacy_httpx_instance
from ..exceptions import ModelTimeoutError
from ..items import ModelResponse, TResponseStreamEvent
from ..logger import log_model_action_debug, logger
from ..models._retry_runtime import (
@@ -34,11 +35,13 @@ from ..retry import (
retry_policy_retries_safe_transport_errors,
)
from ..usage import RequestUsage, Usage
from ..util._error_tracing import mark_model_timeout_task
GetResponseCallable = Callable[[], Awaitable[ModelResponse]]
GetStreamCallable = Callable[[], AsyncIterator[TResponseStreamEvent]]
RewindCallable = Callable[[], Awaitable[None]]
GetRetryAdviceCallable = Callable[[ModelRetryAdviceRequest], ModelRetryAdvice | None]
T = TypeVar("T")
DEFAULT_INITIAL_DELAY_SECONDS = 0.25
DEFAULT_MAX_DELAY_SECONDS = 2.0
@@ -69,6 +72,9 @@ def _is_conversation_locked_error(error: Exception) -> bool:
def _is_abort_like_error(error: Exception) -> bool:
if isinstance(error, ModelTimeoutError):
return False
if isinstance(error, asyncio.CancelledError):
return True
@@ -120,7 +126,7 @@ def _normalize_retry_error(
is_abort=_is_abort_like_error(error),
is_network_error=_is_network_like_error(error),
is_timeout=any(
isinstance(candidate, APITimeoutError | TimeoutError)
isinstance(candidate, APITimeoutError | TimeoutError | ModelTimeoutError)
for candidate in _iter_error_chain(error)
),
)
@@ -204,6 +210,111 @@ async def _sleep_for_retry(delay: float) -> None:
await asyncio.sleep(delay)
async def _drain_model_attempt_task(task: asyncio.Future[Any]) -> BaseException | None:
"""Wait for a cancelled model task without letting its outcome replace cancellation."""
try:
await task
except BaseException as error:
return error
return None
async def _await_cleanup_ignoring_cancellation(
cleanup_task: asyncio.Task[BaseException | None],
) -> BaseException | None:
"""Finish cooperative cleanup before restoring an already-received parent cancellation."""
while True:
try:
return await asyncio.shield(cleanup_task)
except asyncio.CancelledError:
if cleanup_task.done():
return cleanup_task.result()
async def _cancel_and_drain_model_attempt_task(
task: asyncio.Future[Any],
timeout_error: ModelTimeoutError | None = None,
*,
cancel_cleanup: bool = False,
) -> BaseException | None:
if timeout_error is not None:
mark_model_timeout_task(task, timeout_error)
task.cancel()
if cancel_cleanup:
# Give a cancelled model operation one event-loop turn to enter its owner-owned cleanup,
# then interrupt a cooperative cleanup wait that must not outlive the SDK deadline.
# Caller-owned cancellation and early consumer close do not opt into this second cancel.
await asyncio.sleep(0)
if not task.done():
task.cancel()
cleanup_task = asyncio.create_task(_drain_model_attempt_task(task))
try:
return await asyncio.shield(cleanup_task)
except asyncio.CancelledError:
await _await_cleanup_ignoring_cancellation(cleanup_task)
raise
async def _run_stream_attempt_in_one_task(
get_stream: GetStreamCallable,
requests: asyncio.Queue[None],
results: asyncio.Queue[tuple[str, TResponseStreamEvent | BaseException | None]],
) -> None:
"""Own stream construction, pulls, and cleanup in one task context."""
stream: AsyncIterator[TResponseStreamEvent] | None = None
terminal_result: tuple[str, TResponseStreamEvent | BaseException | None] | None = None
try:
stream = get_stream()
while True:
await requests.get()
try:
event = await stream.__anext__()
except StopAsyncIteration:
terminal_result = ("done", None)
break
except BaseException as error:
terminal_result = ("error", error)
break
results.put_nowait(("event", event))
except BaseException as error:
terminal_result = ("error", error)
finally:
if stream is not None:
await _close_async_iterator_quietly(stream)
if terminal_result is not None:
results.put_nowait(terminal_result)
async def _await_model_attempt(
awaitable: Awaitable[T],
timeout: float | None,
*,
timeout_error_seconds: float | None = None,
) -> T:
"""Await one model operation and turn only this deadline into a timeout error."""
if timeout is None:
return await awaitable
task = asyncio.ensure_future(awaitable)
try:
done, pending = await asyncio.wait({task}, timeout=timeout)
except asyncio.CancelledError:
await _cancel_and_drain_model_attempt_task(task)
raise
if task in done:
return await task
timeout_error = ModelTimeoutError(
timeout_error_seconds if timeout_error_seconds is not None else timeout
)
await _cancel_and_drain_model_attempt_task(task, timeout_error, cancel_cleanup=True)
# A traceback retains this frame's locals. Discard every reference that can keep the
# cancelled provider task, its cleanup exception, or provider payload locals reachable.
del awaitable, task, done, pending
raise timeout_error from None
def _build_zero_request_usage_entry() -> RequestUsage:
return RequestUsage(
input_tokens=0,
@@ -468,6 +579,7 @@ async def get_response_with_retry(
get_retry_advice: GetRetryAdviceCallable,
previous_response_id: str | None,
conversation_id: str | None,
timeout: float | None = None,
replay_unsafe_request: bool = False,
) -> ModelResponse:
request_attempt = 1
@@ -497,7 +609,7 @@ async def get_response_with_retry(
),
websocket_pre_event_retries_disabled(disable_websocket_pre_event_retry),
):
response = await get_response()
response = await _await_model_attempt(get_response(), timeout)
response.usage = apply_retry_attempt_usage(
response.usage,
failed_policy_attempts + compatibility_retries_taken,
@@ -577,6 +689,7 @@ async def stream_response_with_retry(
get_retry_advice: GetRetryAdviceCallable,
previous_response_id: str | None,
conversation_id: str | None,
timeout: float | None = None,
failed_retry_attempts_out: list[int] | None = None,
replay_unsafe_request: bool = False,
) -> AsyncGenerator[TResponseStreamEvent, None]:
@@ -595,6 +708,8 @@ async def stream_response_with_retry(
while True:
emitted_retry_unsafe_event = False
stream: AsyncIterator[TResponseStreamEvent] | None = None
stream_owner: asyncio.Task[None] | None = None
deadline = asyncio.get_running_loop().time() + timeout if timeout is not None else None
try:
disable_provider_managed_retries = _should_disable_provider_managed_retries(
retry_settings,
@@ -602,33 +717,113 @@ async def stream_response_with_retry(
stateful_request=stateful_request,
replay_unsafe_request=replay_unsafe_request,
)
# Pull stream events under the retry-disable context, but yield them outside it so
# unrelated model calls made by the consumer do not inherit this setting.
with (
provider_managed_retries_disabled(disable_provider_managed_retries),
websocket_pre_event_retries_disabled(disable_websocket_pre_event_retry),
):
stream = get_stream()
while True:
try:
with (
provider_managed_retries_disabled(disable_provider_managed_retries),
websocket_pre_event_retries_disabled(disable_websocket_pre_event_retry),
):
event = await stream.__anext__()
except StopAsyncIteration:
await _close_async_iterator_quietly(stream)
return
if _stream_event_blocks_retry(event):
emitted_retry_unsafe_event = True
if failed_retry_attempts_out is not None:
failed_retry_attempts_out[:] = [
failed_policy_attempts + compatibility_retries_taken
]
yield event
stream_requests: asyncio.Queue[None] | None = None
stream_results: (
asyncio.Queue[tuple[str, TResponseStreamEvent | BaseException | None]] | None
) = None
if timeout is None:
# Pull stream events under the retry-disable context, but yield them outside it
# so unrelated model calls made by the consumer do not inherit this setting.
with (
provider_managed_retries_disabled(disable_provider_managed_retries),
websocket_pre_event_retries_disabled(disable_websocket_pre_event_retry),
):
stream = get_stream()
else:
stream_requests = asyncio.Queue()
stream_results = asyncio.Queue()
with (
provider_managed_retries_disabled(disable_provider_managed_retries),
websocket_pre_event_retries_disabled(disable_websocket_pre_event_retry),
):
stream_owner = asyncio.create_task(
_run_stream_attempt_in_one_task(
get_stream,
stream_requests,
stream_results,
)
)
try:
while True:
try:
if timeout is None:
assert stream is not None
with (
provider_managed_retries_disabled(disable_provider_managed_retries),
websocket_pre_event_retries_disabled(
disable_websocket_pre_event_retry
),
):
event = await stream.__anext__()
else:
assert stream_owner is not None
assert stream_requests is not None
assert stream_results is not None
assert deadline is not None
stream_requests.put_nowait(None)
remaining = max(deadline - asyncio.get_running_loop().time(), 0.0)
try:
result_kind, result_value = await _await_model_attempt(
stream_results.get(),
remaining,
timeout_error_seconds=timeout,
)
except ModelTimeoutError as error:
await _cancel_and_drain_model_attempt_task(
stream_owner,
error,
cancel_cleanup=True,
)
raise
except asyncio.CancelledError:
await _cancel_and_drain_model_attempt_task(stream_owner)
raise
if result_kind == "done":
remaining = max(deadline - asyncio.get_running_loop().time(), 0.0)
await _await_model_attempt(
stream_owner,
remaining,
timeout_error_seconds=timeout,
)
return
if result_kind == "error":
remaining = max(deadline - asyncio.get_running_loop().time(), 0.0)
await _await_model_attempt(
stream_owner,
remaining,
timeout_error_seconds=timeout,
)
assert isinstance(result_value, BaseException)
raise result_value
assert result_kind == "event"
assert result_value is not None
assert not isinstance(result_value, BaseException)
event = result_value
except StopAsyncIteration:
if stream_owner is None:
await _close_async_iterator_quietly(stream)
return
if _stream_event_blocks_retry(event):
emitted_retry_unsafe_event = True
if failed_retry_attempts_out is not None:
failed_retry_attempts_out[:] = [
failed_policy_attempts + compatibility_retries_taken
]
yield event
finally:
if stream_owner is not None and not stream_owner.done():
await _cancel_and_drain_model_attempt_task(stream_owner)
return
except BaseException as error:
await _close_async_iterator_quietly(stream)
if isinstance(error, ModelTimeoutError):
# The timed owner has already been cancelled and drained. Do not retain its
# task or result queues in the public timeout traceback frame.
stream = None
stream_owner = None
stream_requests = None
stream_results = None
if stream_owner is None:
await _close_async_iterator_quietly(stream)
if isinstance(error, asyncio.CancelledError | GeneratorExit):
raise
if not isinstance(error, Exception):
+2
View File
@@ -2081,6 +2081,7 @@ async def run_single_turn_streamed(
get_retry_advice=model.get_retry_advice,
previous_response_id=previous_response_id,
conversation_id=conversation_id,
timeout=model_settings.timeout,
failed_retry_attempts_out=stream_failed_retry_attempts,
replay_unsafe_request=any(
isinstance(tool, ProgrammaticToolCallingTool) for tool in all_tools
@@ -2463,6 +2464,7 @@ async def get_new_response(
get_retry_advice=model.get_retry_advice,
previous_response_id=previous_response_id,
conversation_id=conversation_id,
timeout=model_settings.timeout,
replay_unsafe_request=any(
isinstance(tool, ProgrammaticToolCallingTool) for tool in all_tools
),
+19 -1
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
import asyncio
import copy
import inspect
import json
@@ -80,7 +81,10 @@ from ..usage import (
_attach_raw_usage_snapshot,
_raw_usage_snapshot,
)
from ..util._error_tracing import REDACTED_TRACE_ERROR_MESSAGE
from ..util._error_tracing import (
REDACTED_TRACE_ERROR_MESSAGE,
record_current_task_model_timeout_on_span,
)
class ModelScriptError(Exception):
@@ -350,6 +354,13 @@ class ScriptedModel(Model):
retry_advice_synced = True
raise step.error
return self._model_response(step, call.model_settings)
except asyncio.CancelledError:
record_current_task_model_timeout_on_span(
span,
message="Error",
trace_include_sensitive_data=call.tracing.include_data(),
)
raise
except Exception as error:
if not retry_advice_synced:
self._forget_retry_advice(error)
@@ -426,6 +437,13 @@ class ScriptedModel(Model):
)
for event in events:
yield event
except asyncio.CancelledError:
record_current_task_model_timeout_on_span(
span,
message="Error",
trace_include_sensitive_data=call.tracing.include_data(),
)
raise
except Exception as error:
if not retry_advice_synced:
self._forget_retry_advice(error)
+41
View File
@@ -1,12 +1,46 @@
import asyncio
import contextlib
from collections.abc import Iterator
from typing import Any
from .. import _debug
from ..exceptions import ModelTimeoutError
from ..logger import logger
from ..tracing import Span, SpanError, get_current_span
REDACTED_TRACE_ERROR_MESSAGE = "Error details are redacted."
_MODEL_TIMEOUT_ERROR_ATTR = "_openai_agents_model_timeout_error"
def mark_model_timeout_task(task: asyncio.Future[Any], error: ModelTimeoutError) -> None:
"""Mark a model-owned task so cancellation-aware spans can record its timeout."""
setattr(task, _MODEL_TIMEOUT_ERROR_ATTR, error)
def get_current_task_model_timeout_error() -> ModelTimeoutError | None:
"""Return the timeout that is cancelling the current model-owned task, if any."""
task = asyncio.current_task()
error = getattr(task, _MODEL_TIMEOUT_ERROR_ATTR, None) if task is not None else None
return error if isinstance(error, ModelTimeoutError) else None
def record_current_task_model_timeout_on_span(
span: Span[Any],
*,
message: str,
trace_include_sensitive_data: bool,
) -> bool:
"""Record a marked timeout cancellation on the current model span."""
timeout_error = get_current_task_model_timeout_error()
if timeout_error is None:
return False
record_model_error_on_span(
span,
message=message,
error=timeout_error,
trace_include_sensitive_data=trace_include_sensitive_data,
)
return True
def get_trace_error(
@@ -102,6 +136,13 @@ def model_span_errors(
"""
try:
yield
except asyncio.CancelledError:
record_current_task_model_timeout_on_span(
span,
message=message,
trace_include_sensitive_data=trace_include_sensitive_data,
)
raise
except Exception as error:
record_model_error_on_span(
span,
+3 -1
View File
@@ -127,6 +127,7 @@ def test_all_fields_serialization() -> None:
context_management=[{"type": "compaction", "compact_threshold": 200000}],
prompt_cache_options={"mode": "explicit", "ttl": "30m"},
preserve_raw_usage=True,
timeout=1.25,
)
# Verify that every single field is set to a non-None value
@@ -158,10 +159,11 @@ def test_gpt_5_6_reasoning_and_prompt_cache_serialization() -> None:
def test_usage_preservation_is_appended_to_public_field_order() -> None:
field_names = [field.name for field in fields(ModelSettings)]
assert field_names[-3:] == [
assert field_names[-4:] == [
"context_management",
"prompt_cache_options",
"preserve_raw_usage",
"timeout",
]
+708
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator
from contextvars import ContextVar
from typing import Any, cast
import httpx2
@@ -9,7 +10,9 @@ import pytest
from openai import APIConnectionError, APIStatusError, BadRequestError
from pydantic import ValidationError
from agents.exceptions import ModelTimeoutError
from agents.items import ModelResponse, TResponseStreamEvent
from agents.model_settings import ModelSettings
from agents.models._openai_retry import get_openai_retry_advice
from agents.models._retry_runtime import (
should_disable_provider_managed_retries,
@@ -58,6 +61,19 @@ def test_model_retry_backoff_settings_allow_zero_values() -> None:
assert backoff.multiplier == 0
@pytest.mark.parametrize("timeout", [0, -0.1, float("inf"), float("nan")])
def test_model_settings_rejects_invalid_model_call_timeout(timeout: float) -> None:
with pytest.raises(ValidationError):
ModelSettings(timeout=timeout)
def test_model_settings_accepts_finite_model_call_timeout() -> None:
settings = ModelSettings(timeout=1.25)
assert settings.timeout == 1.25
assert settings.to_traceable_dict()["timeout"] == 1.25
def test_retry_capabilities_preserve_falsey_policy() -> None:
class FalseyPolicy:
_openai_agents_retries_safe_transport_errors: bool
@@ -130,6 +146,265 @@ def _status_error_without_code(status_code: int, body_code: str = "server_error"
)
@pytest.mark.asyncio
async def test_get_response_with_retry_times_out_and_retries_stateless_attempt() -> None:
calls = 0
async def get_response() -> ModelResponse:
nonlocal calls
calls += 1
if calls == 1:
await asyncio.Event().wait()
return ModelResponse(output=[get_text_message("ok")], usage=Usage(), response_id="resp")
result = await get_response_with_retry(
get_response=get_response,
rewind=lambda: asyncio.sleep(0),
retry_settings=ModelRetrySettings(
max_retries=1,
backoff={"initial_delay": 0},
policy=retry_policies.network_error(),
),
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=0.01,
)
assert calls == 2
assert result.response_id == "resp"
@pytest.mark.asyncio
async def test_get_response_with_retry_blocks_unsafe_stateful_timeout_replay() -> None:
calls = 0
async def get_response() -> ModelResponse:
nonlocal calls
calls += 1
await asyncio.Event().wait()
raise AssertionError("unreachable")
with pytest.raises(ModelTimeoutError, match="timed out after 0.01 seconds"):
await get_response_with_retry(
get_response=get_response,
rewind=lambda: asyncio.sleep(0),
retry_settings=ModelRetrySettings(
max_retries=1,
backoff={"initial_delay": 0},
policy=retry_policies.network_error(),
),
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id="conv_123",
timeout=0.01,
)
assert calls == 1
@pytest.mark.asyncio
async def test_model_timeout_discards_cleanup_exception_graph() -> None:
sensitive_payload = "sensitive provider payload"
async def get_response() -> ModelResponse:
try:
await asyncio.Event().wait()
except asyncio.CancelledError as exc:
raise RuntimeError(sensitive_payload) from exc
raise AssertionError("unreachable")
with pytest.raises(ModelTimeoutError) as exc_info:
await get_response_with_retry(
get_response=get_response,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=0.01,
)
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
assert sensitive_payload not in repr(exc_info.value)
traceback = exc_info.value.__traceback__
while traceback is not None:
if traceback.tb_frame.f_code.co_name == "_await_model_attempt":
frame_locals = traceback.tb_frame.f_locals
assert "awaitable" not in frame_locals
assert "task" not in frame_locals
assert "done" not in frame_locals
assert "pending" not in frame_locals
traceback = traceback.tb_next
@pytest.mark.asyncio
async def test_stream_timeout_discards_owner_task_from_traceback_locals() -> None:
sensitive_payload = "sensitive stream cleanup payload"
def get_stream() -> AsyncIterator[TResponseStreamEvent]:
async def iterator() -> AsyncIterator[TResponseStreamEvent]:
try:
await asyncio.Event().wait()
finally:
raise RuntimeError(sensitive_payload)
yield cast(TResponseStreamEvent, {"type": "response.created"})
return iterator()
with pytest.raises(ModelTimeoutError) as exc_info:
async for _event in stream_response_with_retry(
get_stream=get_stream,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=0.01,
):
pass
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
traceback = exc_info.value.__traceback__
while traceback is not None:
if traceback.tb_frame.f_code.co_name == "stream_response_with_retry":
frame_locals = traceback.tb_frame.f_locals
assert frame_locals.get("stream_owner") is None
assert frame_locals.get("stream_requests") is None
assert frame_locals.get("stream_results") is None
traceback = traceback.tb_next
@pytest.mark.asyncio
async def test_model_timeout_cancels_blocked_cleanup_after_deadline() -> None:
cleanup_started = asyncio.Event()
release_cleanup = asyncio.Event()
async def get_response() -> ModelResponse:
try:
await asyncio.Event().wait()
finally:
cleanup_started.set()
await release_cleanup.wait()
raise AssertionError("unreachable")
task = asyncio.create_task(
get_response_with_retry(
get_response=get_response,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=0.01,
)
)
await cleanup_started.wait()
done, _ = await asyncio.wait({task}, timeout=0.2)
if task not in done:
release_cleanup.set()
await task
pytest.fail("Non-streaming timeout cleanup did not receive a second cancellation.")
with pytest.raises(ModelTimeoutError) as exc_info:
await task
assert exc_info.value.timeout_seconds == 0.01
@pytest.mark.asyncio
async def test_get_response_with_retry_preserves_parent_cancellation() -> None:
started = asyncio.Event()
async def get_response() -> ModelResponse:
started.set()
await asyncio.Event().wait()
raise AssertionError("unreachable")
task = asyncio.create_task(
get_response_with_retry(
get_response=get_response,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=10,
)
)
await started.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
@pytest.mark.asyncio
async def test_get_response_with_retry_preserves_parent_cancellation_during_timeout_cleanup() -> (
None
):
cleanup_started = asyncio.Event()
release_cleanup = asyncio.Event()
async def get_response() -> ModelResponse:
try:
await asyncio.Event().wait()
except asyncio.CancelledError:
cleanup_started.set()
await release_cleanup.wait()
raise AssertionError("unreachable")
task = asyncio.create_task(
get_response_with_retry(
get_response=get_response,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=0.01,
)
)
await cleanup_started.wait()
task.cancel()
release_cleanup.set()
with pytest.raises(asyncio.CancelledError):
await task
@pytest.mark.asyncio
async def test_get_response_with_retry_preserves_parent_cancellation_over_cleanup_error() -> None:
started = asyncio.Event()
async def get_response() -> ModelResponse:
started.set()
try:
await asyncio.Event().wait()
except asyncio.CancelledError as exc:
raise RuntimeError("cleanup failed") from exc
raise AssertionError("unreachable")
task = asyncio.create_task(
get_response_with_retry(
get_response=get_response,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=10,
)
)
await started.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
@pytest.mark.asyncio
async def test_programmatic_request_disables_hidden_provider_retries() -> None:
provider_retry_flags: list[bool] = []
@@ -1516,6 +1791,439 @@ async def test_stream_response_with_retry_retries_before_first_event(monkeypatch
assert events == [cast(TResponseStreamEvent, {"type": "response.created"})]
@pytest.mark.asyncio
async def test_stream_response_with_retry_retries_timeout_before_output() -> None:
attempts = 0
def get_stream() -> AsyncIterator[TResponseStreamEvent]:
nonlocal attempts
attempts += 1
async def iterator() -> AsyncIterator[TResponseStreamEvent]:
if attempts == 1:
await asyncio.Event().wait()
yield cast(TResponseStreamEvent, {"type": "response.created"})
return iterator()
events = [
event
async for event in stream_response_with_retry(
get_stream=get_stream,
rewind=lambda: asyncio.sleep(0),
retry_settings=ModelRetrySettings(
max_retries=1,
backoff={"initial_delay": 0},
policy=retry_policies.network_error(),
),
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=0.01,
)
]
assert attempts == 2
assert events == [cast(TResponseStreamEvent, {"type": "response.created"})]
@pytest.mark.asyncio
async def test_stream_response_with_retry_keeps_one_context_for_timed_attempt() -> None:
context_value: ContextVar[str | None] = ContextVar("context_value", default=None)
def get_stream() -> AsyncIterator[TResponseStreamEvent]:
async def iterator() -> AsyncIterator[TResponseStreamEvent]:
token = context_value.set("active")
try:
yield cast(TResponseStreamEvent, {"type": "response.created"})
yield cast(TResponseStreamEvent, {"type": "response.in_progress"})
finally:
context_value.reset(token)
return iterator()
events = [
event
async for event in stream_response_with_retry(
get_stream=get_stream,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=1,
)
]
assert events == [
cast(TResponseStreamEvent, {"type": "response.created"}),
cast(TResponseStreamEvent, {"type": "response.in_progress"}),
]
assert context_value.get() is None
@pytest.mark.asyncio
async def test_timed_stream_constructs_and_closes_iterator_in_one_context() -> None:
context_value: ContextVar[str | None] = ContextVar("context_value", default=None)
class ConstructionScopedStream:
def __init__(self) -> None:
self.token = context_value.set("active")
self.emitted = False
def __aiter__(self) -> ConstructionScopedStream:
return self
async def __anext__(self) -> TResponseStreamEvent:
if self.emitted:
raise StopAsyncIteration
self.emitted = True
return cast(TResponseStreamEvent, {"type": "response.created"})
async def aclose(self) -> None:
context_value.reset(self.token)
events = [
event
async for event in stream_response_with_retry(
get_stream=ConstructionScopedStream,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=1,
)
]
assert events == [cast(TResponseStreamEvent, {"type": "response.created"})]
assert context_value.get() is None
@pytest.mark.asyncio
async def test_timed_stream_early_close_keeps_generator_cleanup_context() -> None:
context_value: ContextVar[str | None] = ContextVar("context_value", default=None)
cleanup_state: list[str] = []
def get_stream() -> AsyncIterator[TResponseStreamEvent]:
async def iterator() -> AsyncIterator[TResponseStreamEvent]:
token = context_value.set("active")
try:
yield cast(TResponseStreamEvent, {"type": "response.created"})
await asyncio.Event().wait()
finally:
context_value.reset(token)
cleanup_state.append("reset")
return iterator()
stream = stream_response_with_retry(
get_stream=get_stream,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=1,
)
assert await stream.__anext__() == cast(TResponseStreamEvent, {"type": "response.created"})
await stream.aclose()
assert cleanup_state == ["reset"]
assert context_value.get() is None
@pytest.mark.asyncio
async def test_timed_stream_early_close_allows_cooperative_cleanup() -> None:
cleanup_state: list[str] = []
class CooperativeCloseStream:
def __init__(self) -> None:
self.emitted = False
def __aiter__(self) -> CooperativeCloseStream:
return self
async def __anext__(self) -> TResponseStreamEvent:
if not self.emitted:
self.emitted = True
return cast(TResponseStreamEvent, {"type": "response.created"})
await asyncio.Event().wait()
raise AssertionError("unreachable")
async def aclose(self) -> None:
await asyncio.sleep(0)
cleanup_state.append("closed")
stream = stream_response_with_retry(
get_stream=CooperativeCloseStream,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=1,
)
assert await stream.__anext__() == cast(TResponseStreamEvent, {"type": "response.created"})
await stream.aclose()
assert cleanup_state == ["closed"]
@pytest.mark.asyncio
async def test_timed_stream_parent_cancellation_allows_cooperative_cleanup() -> None:
read_started = asyncio.Event()
cleanup_state: list[str] = []
class CooperativeCancelStream:
def __aiter__(self) -> CooperativeCancelStream:
return self
async def __anext__(self) -> TResponseStreamEvent:
read_started.set()
await asyncio.Event().wait()
raise AssertionError("unreachable")
async def aclose(self) -> None:
await asyncio.sleep(0)
cleanup_state.append("closed")
async def consume() -> None:
async for _event in stream_response_with_retry(
get_stream=CooperativeCancelStream,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=1,
):
pass
task = asyncio.create_task(consume())
await read_started.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert cleanup_state == ["closed"]
@pytest.mark.asyncio
async def test_timed_stream_closes_provider_iterator_once() -> None:
class CloseCountingStream:
def __init__(self) -> None:
self.emitted = False
self.close_calls = 0
def __aiter__(self) -> CloseCountingStream:
return self
async def __anext__(self) -> TResponseStreamEvent:
if self.emitted:
raise StopAsyncIteration
self.emitted = True
return cast(TResponseStreamEvent, {"type": "response.created"})
async def aclose(self) -> None:
self.close_calls += 1
provider_stream = CloseCountingStream()
events = [
event
async for event in stream_response_with_retry(
get_stream=lambda: provider_stream,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=1,
)
]
assert events == [cast(TResponseStreamEvent, {"type": "response.created"})]
assert provider_stream.close_calls == 1
@pytest.mark.asyncio
async def test_stream_response_with_retry_does_not_retry_timeout_after_output() -> None:
attempts = 0
def get_stream() -> AsyncIterator[TResponseStreamEvent]:
nonlocal attempts
attempts += 1
async def iterator() -> AsyncIterator[TResponseStreamEvent]:
yield cast(TResponseStreamEvent, {"type": "response.output_item.added"})
await asyncio.Event().wait()
return iterator()
with pytest.raises(ModelTimeoutError):
async for _event in stream_response_with_retry(
get_stream=get_stream,
rewind=lambda: asyncio.sleep(0),
retry_settings=ModelRetrySettings(
max_retries=1,
backoff={"initial_delay": 0},
policy=retry_policies.network_error(),
),
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=0.01,
):
pass
assert attempts == 1
@pytest.mark.asyncio
async def test_stream_timeout_reports_configured_attempt_timeout() -> None:
def get_stream() -> AsyncIterator[TResponseStreamEvent]:
async def iterator() -> AsyncIterator[TResponseStreamEvent]:
yield cast(TResponseStreamEvent, {"type": "response.created"})
await asyncio.Event().wait()
return iterator()
stream = stream_response_with_retry(
get_stream=get_stream,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=0.05,
)
assert await stream.__anext__() == cast(TResponseStreamEvent, {"type": "response.created"})
await asyncio.sleep(0.03)
with pytest.raises(ModelTimeoutError) as exc_info:
await stream.__anext__()
assert exc_info.value.timeout_seconds == 0.05
assert str(exc_info.value) == "Model call timed out after 0.05 seconds."
@pytest.mark.asyncio
async def test_stream_timeout_bounds_owner_cleanup_after_exhaustion() -> None:
class SlowCloseStream:
def __aiter__(self) -> SlowCloseStream:
return self
async def __anext__(self) -> TResponseStreamEvent:
raise StopAsyncIteration
async def aclose(self) -> None:
await asyncio.Event().wait()
with pytest.raises(ModelTimeoutError) as exc_info:
async for _event in stream_response_with_retry(
get_stream=SlowCloseStream,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=0.01,
):
pass
assert exc_info.value.timeout_seconds == 0.01
@pytest.mark.asyncio
async def test_stream_timeout_cancels_blocked_cleanup_after_blocked_read() -> None:
read_started = asyncio.Event()
close_started = asyncio.Event()
release_close = asyncio.Event()
class BlockingReadAndCloseStream:
def __aiter__(self) -> BlockingReadAndCloseStream:
return self
async def __anext__(self) -> TResponseStreamEvent:
read_started.set()
await asyncio.Event().wait()
raise AssertionError("unreachable")
async def aclose(self) -> None:
close_started.set()
await release_close.wait()
async def consume() -> None:
async for _event in stream_response_with_retry(
get_stream=BlockingReadAndCloseStream,
rewind=lambda: asyncio.sleep(0),
retry_settings=None,
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=0.01,
):
pass
task = asyncio.create_task(consume())
await read_started.wait()
done, _ = await asyncio.wait({task}, timeout=0.2)
if task not in done:
release_close.set()
await task
pytest.fail("Timed stream cleanup did not receive a second cancellation.")
with pytest.raises(ModelTimeoutError) as exc_info:
await task
assert close_started.is_set()
assert exc_info.value.timeout_seconds == 0.01
@pytest.mark.asyncio
async def test_stream_response_with_retry_preserves_parent_cancellation_over_cleanup_error() -> (
None
):
started = asyncio.Event()
def get_stream() -> AsyncIterator[TResponseStreamEvent]:
async def iterator() -> AsyncIterator[TResponseStreamEvent]:
started.set()
try:
await asyncio.Event().wait()
except asyncio.CancelledError as exc:
raise RuntimeError("stream cleanup failed") from exc
yield cast(TResponseStreamEvent, {"type": "response.created"})
return iterator()
async def consume() -> None:
async for _event in stream_response_with_retry(
get_stream=get_stream,
rewind=lambda: asyncio.sleep(0),
retry_settings=ModelRetrySettings(
max_retries=1,
backoff={"initial_delay": 0},
policy=retry_policies.network_error(),
),
get_retry_advice=lambda _request: None,
previous_response_id=None,
conversation_id=None,
timeout=10,
):
pass
task = asyncio.create_task(consume())
await started.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
@pytest.mark.asyncio
async def test_stream_response_with_retry_keeps_provider_retries_on_first_attempt(
monkeypatch,
+33
View File
@@ -8,6 +8,7 @@ tests pin the same behavior for the other providers.
from __future__ import annotations
import asyncio
from typing import Any
import pytest
@@ -347,6 +348,38 @@ def test_failing_span_recording_preserves_the_provider_exception(
assert exc_info.value is original
@pytest.mark.asyncio
async def test_marked_model_timeout_cancellation_records_span_error() -> None:
from agents.exceptions import ModelTimeoutError
from agents.tracing import generation_span
from agents.util._error_tracing import mark_model_timeout_task, model_span_errors
started = asyncio.Event()
async def run() -> None:
with trace(workflow_name="test"):
with generation_span() as span:
with model_span_errors(
span,
message="Error getting response",
trace_include_sensitive_data=True,
):
started.set()
await asyncio.Event().wait()
task = asyncio.create_task(run())
await started.wait()
mark_model_timeout_task(task, ModelTimeoutError(0.01))
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
error = _span_error("generation")
assert error is not None
assert error["data"]["error"] == "Model call timed out after 0.01 seconds."
class _TerminalFailureEvent:
"""A terminal `response.failed` event with no response payload attached."""
+31
View File
@@ -51,6 +51,7 @@ from agents import (
ModelRetryAdvice,
ModelRetryAdviceRequest,
ModelRetrySettings,
ModelTimeoutError,
RunConfig,
Runner,
handoff,
@@ -1326,6 +1327,36 @@ async def test_scripted_model_preserves_unformattable_responder_error(streamed:
]
@pytest.mark.asyncio
@pytest.mark.parametrize("streamed", [False, True])
async def test_scripted_model_timeout_records_generation_span_error(streamed: bool) -> None:
async def respond(_call: ModelCall) -> Any:
await asyncio.Event().wait()
raise AssertionError("unreachable")
model = ScriptedModel([ModelStep.respond(respond)], emit_traces=True)
agent = Agent(
name="test",
model=model,
model_settings=ModelSettings(timeout=0.01),
)
with pytest.raises(ModelTimeoutError):
if streamed:
result = Runner.run_streamed(agent, "hi")
async for _event in result.stream_events():
pass
else:
await Runner.run(agent, "hi")
assert fetch_span_errors("generation") == [
{
"message": "Error",
"data": {"error": "Model call timed out after 0.01 seconds."},
}
]
def test_scripted_model_ignores_span_attachment_failure() -> None:
class FailingSpan:
def set_error(self, _error: SpanError) -> None: