feat(core): add model call timeouts (#4428)
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
),
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user