diff --git a/src/agents/models/openai_responses.py b/src/agents/models/openai_responses.py index b9c6c75013..33de03b69f 100644 --- a/src/agents/models/openai_responses.py +++ b/src/agents/models/openai_responses.py @@ -53,7 +53,7 @@ from ..items import ItemHelpers, ModelResponse, TResponseInputItem from ..logger import log_model_action_debug, log_model_action_error, logger from ..model_settings import MCPToolChoice -from ..retry import ModelRetryAdvice, ModelRetryAdviceRequest +from ..retry import ModelRetryAdvice, ModelRetryAdviceRequest, ModelRetryNormalizedError from ..tool import ( ApplyPatchTool, CodeInterpreterTool, @@ -444,15 +444,41 @@ def _get_wrapped_websocket_replay_safety(error: Exception) -> str | None: return replay_safety if replay_safety in {"safe", "unsafe"} else None +def _mark_websocket_close_invalidation(error: Exception) -> None: + setattr(error, "_openai_agents_ws_close_invalidated", True) # noqa: B010 + + +def _did_websocket_close_invalidate(error: Exception) -> bool: + return any( + getattr(candidate, "_openai_agents_ws_close_invalidated", False) + for candidate in _iter_retry_error_chain(error) + ) + + +def _websocket_close_invalidation_error() -> RuntimeError: + error = RuntimeError("Responses websocket connection closed while establishing a connection.") + _mark_websocket_close_invalidation(error) + return error + + def _did_start_websocket_response(error: Exception) -> bool: return bool(getattr(error, "_openai_agents_ws_response_started", False)) +def _is_websocket_disconnect_error(error: Exception) -> bool: + exc_module = error.__class__.__module__ + exc_name = error.__class__.__name__ + # websockets reports a peer closing before a valid HTTP upgrade as InvalidMessage. Only an + # InvalidMessage caused by EOFError is transient according to websockets' retry policy. + return exc_module.startswith("websockets") and ( + exc_name.startswith("ConnectionClosed") + or (exc_name == "InvalidMessage" and isinstance(error.__cause__, EOFError)) + ) + + def _is_never_sent_websocket_error(error: Exception) -> bool: for candidate in _iter_retry_error_chain(error): - if candidate.__class__.__module__.startswith( - "websockets" - ) and candidate.__class__.__name__.startswith("ConnectionClosed"): + if _is_websocket_disconnect_error(candidate): if "client closed" not in str(candidate).lower(): return True return False @@ -1153,6 +1179,13 @@ def _supports_default_prompt_cache_key(self) -> bool: return super()._supports_default_prompt_cache_key() def get_retry_advice(self, request: ModelRetryAdviceRequest) -> ModelRetryAdvice | None: + if _did_websocket_close_invalidate(request.error): + return ModelRetryAdvice( + suggested=False, + reason=str(request.error), + normalized=ModelRetryNormalizedError(is_abort=True), + ) + stateful_request = bool(request.previous_response_id or request.conversation_id) wrapped_replay_safety = _get_wrapped_websocket_replay_safety(request.error) if wrapped_replay_safety == "unsafe": @@ -1343,17 +1376,26 @@ async def _iter_websocket_response_events( ) retry_pre_event_disconnect = _should_retry_pre_event_websocket_disconnect() while True: - connection = await self._await_websocket_with_timeout( - self._ensure_websocket_connection( - ws_url, request_headers, connect_timeout=request_timeouts.connect - ), - request_timeouts.connect, - "connect", - ) + connection: Any = None received_any_event = False yielded_terminal_event = False sent_request_frame = False try: + connection = await self._await_websocket_with_timeout( + self._ensure_websocket_connection( + ws_url, + request_headers, + connect_timeout=request_timeouts.connect, + request_close_generation=request_close_generation, + ), + request_timeouts.connect, + "connect", + ) + if self._ws_client_close_generation != request_close_generation: + await self._drop_websocket_connection() + connection = None + raise _websocket_close_invalidation_error() + # Once we begin awaiting `send()`, treat the request as potentially # transmitted to avoid replaying it on send/close races. sent_request_frame = True @@ -1410,11 +1452,15 @@ async def _iter_websocket_response_events( is_non_terminal_generator_exit = ( isinstance(exc, GeneratorExit) and not yielded_terminal_event ) - if isinstance(exc, asyncio.CancelledError) or is_non_terminal_generator_exit: - self._force_abort_websocket_connection(connection) - self._clear_websocket_connection_state() - elif not (yielded_terminal_event and isinstance(exc, GeneratorExit)): - await self._drop_websocket_connection() + if connection is not None: + if ( + isinstance(exc, asyncio.CancelledError) + or is_non_terminal_generator_exit + ): + self._force_abort_websocket_connection(connection) + self._clear_websocket_connection_state() + elif not (yielded_terminal_event and isinstance(exc, GeneratorExit)): + await self._drop_websocket_connection() if ( isinstance(exc, Exception) @@ -1435,10 +1481,12 @@ async def _iter_websocket_response_events( is_pre_event_disconnect and not sent_request_frame ) if ( - is_pre_event_disconnect + isinstance(exc, Exception) and self._ws_client_close_generation != request_close_generation ): - raise + _mark_websocket_close_invalidation(exc) + if is_pre_event_disconnect: + raise if retry_pre_event_disconnect and is_retryable_pre_event_disconnect: retry_pre_event_disconnect = False continue @@ -1472,9 +1520,7 @@ def _should_wrap_pre_event_websocket_disconnect(self, exc: Exception) -> bool: "Responses websocket connection closed before a terminal response event." ) - exc_module = exc.__class__.__module__ - exc_name = exc.__class__.__name__ - return exc_module.startswith("websockets") and exc_name.startswith("ConnectionClosed") + return _is_websocket_disconnect_error(exc) def _get_websocket_request_timeouts(self, timeout: Any) -> _WebsocketRequestTimeouts: if timeout is None or _is_openai_omitted_value(timeout): @@ -1593,6 +1639,7 @@ async def _ensure_websocket_connection( headers: Mapping[str, str], *, connect_timeout: float | None, + request_close_generation: int | None = None, ) -> Any: running_loop = asyncio.get_running_loop() identity = ( @@ -1600,6 +1647,12 @@ async def _ensure_websocket_connection( tuple(sorted((str(key).lower(), str(value)) for key, value in headers.items())), ) + if ( + request_close_generation is not None + and self._ws_client_close_generation != request_close_generation + ): + raise _websocket_close_invalidation_error() + if self._ws_connection is not None and self._ws_connection_identity == identity: if ( self._ws_connection_loop_ref is not None @@ -1609,6 +1662,11 @@ async def _ensure_websocket_connection( return self._ws_connection if self._ws_connection is not None: await self._drop_websocket_connection() + if ( + request_close_generation is not None + and self._ws_client_close_generation != request_close_generation + ): + raise _websocket_close_invalidation_error() self._ws_connection = await self._open_websocket_connection( ws_url, headers, diff --git a/tests/models/test_model_retry.py b/tests/models/test_model_retry.py index c8079ab8a7..bac1565e87 100644 --- a/tests/models/test_model_retry.py +++ b/tests/models/test_model_retry.py @@ -18,6 +18,7 @@ should_disable_provider_managed_retries, should_disable_websocket_pre_event_retries, ) +from agents.models.openai_responses import OpenAIResponsesWSModel from agents.retry import ( ModelRetryAdvice, ModelRetryAdviceRequest, @@ -1630,6 +1631,82 @@ async def get_response() -> ModelResponse: assert calls == 1 +def _ws_close_invalidated_error() -> RuntimeError: + error = RuntimeError("Responses websocket connection closed while establishing a connection.") + setattr(error, "_openai_agents_ws_close_invalidated", True) # noqa: B010 + return error + + +@pytest.mark.asyncio +async def test_get_response_with_retry_does_not_replay_websocket_close_invalidated_request() -> ( + None +): + model = OpenAIResponsesWSModel(model="gpt-4", openai_client=cast(Any, object())) + calls = 0 + + async def get_response() -> ModelResponse: + nonlocal calls + calls += 1 + raise _ws_close_invalidated_error() + + async def rewind() -> None: + raise AssertionError("A close-invalidated request must not be rewound for retry") + + with pytest.raises(RuntimeError, match="closed while establishing"): + await get_response_with_retry( + get_response=get_response, + rewind=rewind, + retry_settings=ModelRetrySettings( + max_retries=1, + backoff={"initial_delay": 0}, + policy=retry_policies.network_error(), + ), + get_retry_advice=model.get_retry_advice, + previous_response_id=None, + conversation_id=None, + ) + + assert calls == 1 + + +@pytest.mark.asyncio +async def test_stream_response_with_retry_does_not_replay_websocket_close_invalidated_request() -> ( + None +): + model = OpenAIResponsesWSModel(model="gpt-4", openai_client=cast(Any, object())) + calls = 0 + + def get_stream() -> AsyncIterator[TResponseStreamEvent]: + nonlocal calls + calls += 1 + + async def iterator() -> AsyncIterator[TResponseStreamEvent]: + raise _ws_close_invalidated_error() + yield # pragma: no cover + + return iterator() + + async def rewind() -> None: + raise AssertionError("A close-invalidated request must not be rewound for retry") + + with pytest.raises(RuntimeError, match="closed while establishing"): + async for _event in stream_response_with_retry( + get_stream=get_stream, + rewind=rewind, + retry_settings=ModelRetrySettings( + max_retries=1, + backoff={"initial_delay": 0}, + policy=retry_policies.network_error(), + ), + get_retry_advice=model.get_retry_advice, + previous_response_id=None, + conversation_id=None, + ): + pass + + assert calls == 1 + + @pytest.mark.asyncio async def test_get_response_with_retry_allows_custom_policy_to_override_provider_veto( monkeypatch, diff --git a/tests/models/test_openai_responses.py b/tests/models/test_openai_responses.py index d6b787df9b..cfc7a7e45f 100644 --- a/tests/models/test_openai_responses.py +++ b/tests/models/test_openai_responses.py @@ -244,6 +244,17 @@ class ConnectionClosedError(Exception): return ConnectionClosedError(message) +def _invalid_message_error(message: str, *, cause: BaseException | None = None) -> Exception: + class InvalidMessage(Exception): + pass + + InvalidMessage.__module__ = "websockets.exceptions" + error = InvalidMessage(message) + if cause is not None: + error.__cause__ = cause + return error + + @pytest.mark.parametrize("parallel_tool_calls", [True, False, None]) @pytest.mark.parametrize("tool_source", ["none", "function", "handoff"]) def test_parallel_tool_calls_follow_converted_responses_tools( @@ -2925,6 +2936,292 @@ async def fake_open( assert model._ws_connection is ws2 +@pytest.mark.allow_call_model_methods +@pytest.mark.asyncio +async def test_websocket_model_retries_if_handshake_fails_before_request(monkeypatch): + client = DummyWSClient() + + ws = DummyWSConnection([_response_completed_frame("resp-retried", 1)]) + model = OpenAIResponsesWSModel(model="gpt-4", openai_client=client) # type: ignore[arg-type] + open_calls = 0 + + async def fake_open( + ws_url: str, headers: dict[str, str], *, connect_timeout: float | None = None + ) -> DummyWSConnection: + nonlocal open_calls + open_calls += 1 + if open_calls == 1: + raise _invalid_message_error( + "did not receive a valid HTTP response", + cause=EOFError("connection closed while reading HTTP status line"), + ) + return ws + + monkeypatch.setattr(model, "_open_websocket_connection", fake_open) + + response = await model.get_response( + system_instructions=None, + input="hi", + model_settings=ModelSettings(), + tools=[], + output_schema=None, + handoffs=[], + tracing=ModelTracing.DISABLED, + ) + + assert response.response_id == "resp-retried" + assert open_calls == 2 + assert len(ws.sent_messages) == 1 + + +@pytest.mark.allow_call_model_methods +@pytest.mark.asyncio +async def test_websocket_model_does_not_retry_malformed_handshake(monkeypatch): + client = DummyWSClient() + error = _invalid_message_error("malformed HTTP status line") + model = OpenAIResponsesWSModel(model="gpt-4", openai_client=client) # type: ignore[arg-type] + open_calls = 0 + + async def fake_open( + ws_url: str, headers: dict[str, str], *, connect_timeout: float | None = None + ) -> DummyWSConnection: + nonlocal open_calls + open_calls += 1 + raise error + + monkeypatch.setattr(model, "_open_websocket_connection", fake_open) + + with pytest.raises(type(error)) as exc_info: + await model.get_response( + system_instructions=None, + input="hi", + model_settings=ModelSettings(), + tools=[], + output_schema=None, + handoffs=[], + tracing=ModelTracing.DISABLED, + ) + + assert exc_info.value is error + assert open_calls == 1 + + +@pytest.mark.allow_call_model_methods +@pytest.mark.asyncio +async def test_websocket_model_exhausts_one_retry_after_repeated_eof_handshake_failures( + monkeypatch, +): + client = DummyWSClient() + model = OpenAIResponsesWSModel(model="gpt-4", openai_client=client) # type: ignore[arg-type] + open_calls = 0 + + async def fake_open( + ws_url: str, headers: dict[str, str], *, connect_timeout: float | None = None + ) -> DummyWSConnection: + nonlocal open_calls + open_calls += 1 + raise _invalid_message_error( + "did not receive a valid HTTP response", + cause=EOFError("connection closed while reading HTTP status line"), + ) + + monkeypatch.setattr(model, "_open_websocket_connection", fake_open) + + with pytest.raises(RuntimeError, match="before any response events were received"): + await model.get_response( + system_instructions=None, + input="hi", + model_settings=ModelSettings(), + tools=[], + output_schema=None, + handoffs=[], + tracing=ModelTracing.DISABLED, + ) + + assert open_calls == 2 + assert model._ws_connection is None + + +@pytest.mark.allow_call_model_methods +@pytest.mark.asyncio +async def test_websocket_model_close_during_failing_handshake_prevents_retry(monkeypatch): + client = DummyWSClient() + model = OpenAIResponsesWSModel(model="gpt-4", openai_client=client) # type: ignore[arg-type] + handshake_started = asyncio.Event() + release_handshake = asyncio.Event() + error = _invalid_message_error( + "did not receive a valid HTTP response", + cause=EOFError("connection closed while reading HTTP status line"), + ) + open_calls = 0 + + async def fake_open( + ws_url: str, headers: dict[str, str], *, connect_timeout: float | None = None + ) -> DummyWSConnection: + nonlocal open_calls + open_calls += 1 + handshake_started.set() + await release_handshake.wait() + raise error + + monkeypatch.setattr(model, "_open_websocket_connection", fake_open) + + request_task = asyncio.create_task( + model.get_response( + system_instructions=None, + input="hi", + model_settings=ModelSettings(), + tools=[], + output_schema=None, + handoffs=[], + tracing=ModelTracing.DISABLED, + ) + ) + await handshake_started.wait() + await model.close() + release_handshake.set() + + with pytest.raises(type(error)) as exc_info: + await request_task + + assert exc_info.value is error + advice = model.get_retry_advice( + ModelRetryAdviceRequest(error=exc_info.value, attempt=1, stream=False) + ) + assert advice is not None + assert advice.normalized is not None + assert advice.normalized.is_abort is True + assert open_calls == 1 + + +@pytest.mark.allow_call_model_methods +@pytest.mark.asyncio +async def test_websocket_model_close_during_retry_handshake_discards_acquired_connection( + monkeypatch, +): + client = DummyWSClient() + ws = DummyWSConnection([_response_completed_frame("resp-after-close", 1)]) + first_handshake = True + second_handshake_started = asyncio.Event() + release_second_handshake = asyncio.Event() + error = _invalid_message_error( + "did not receive a valid HTTP response", + cause=EOFError("connection closed while reading HTTP status line"), + ) + open_calls = 0 + + async def fake_open( + ws_url: str, headers: dict[str, str], *, connect_timeout: float | None = None + ) -> DummyWSConnection: + nonlocal first_handshake, open_calls + open_calls += 1 + if first_handshake: + first_handshake = False + raise error + second_handshake_started.set() + await release_second_handshake.wait() + return ws + + model = OpenAIResponsesWSModel(model="gpt-4", openai_client=client) # type: ignore[arg-type] + monkeypatch.setattr(model, "_open_websocket_connection", fake_open) + + request_task = asyncio.create_task( + model.get_response( + system_instructions=None, + input="hi", + model_settings=ModelSettings(), + tools=[], + output_schema=None, + handoffs=[], + tracing=ModelTracing.DISABLED, + ) + ) + await second_handshake_started.wait() + await model.close() + release_second_handshake.set() + + with pytest.raises(RuntimeError) as exc_info: + await request_task + + advice = model.get_retry_advice( + ModelRetryAdviceRequest(error=exc_info.value, attempt=1, stream=False) + ) + assert advice is not None + assert advice.normalized is not None + assert advice.normalized.is_abort is True + assert open_calls == 2 + assert ws.sent_messages == [] + assert ws.close_calls == 1 + assert model._ws_connection is None + + +@pytest.mark.allow_call_model_methods +@pytest.mark.asyncio +async def test_websocket_model_close_during_cached_connection_drop_prevents_replacement( + monkeypatch, +): + client = DummyWSClient() + + class BlockingCloseWSConnection(DummyWSConnection): + def __init__(self, frames: list[str]): + super().__init__(frames) + self.close_started = asyncio.Event() + self.release_close = asyncio.Event() + + async def close(self) -> None: + self.close_started.set() + await self.release_close.wait() + await super().close() + + ws1 = BlockingCloseWSConnection([_response_completed_frame("resp-1", 1)]) + ws2 = DummyWSConnection([_response_completed_frame("resp-2", 1)]) + opened: list[DummyWSConnection] = [] + + async def fake_open( + ws_url: str, headers: dict[str, str], *, connect_timeout: float | None = None + ) -> DummyWSConnection: + connection = ws1 if not opened else ws2 + opened.append(connection) + return connection + + model = OpenAIResponsesWSModel(model="gpt-4", openai_client=client) # type: ignore[arg-type] + monkeypatch.setattr(model, "_open_websocket_connection", fake_open) + + first = await model.get_response( + system_instructions=None, + input="hi", + model_settings=ModelSettings(), + tools=[], + output_schema=None, + handoffs=[], + tracing=ModelTracing.DISABLED, + ) + assert first.response_id == "resp-1" + + ws1.close_code = 1001 + request_task = asyncio.create_task( + model.get_response( + system_instructions=None, + input="next", + model_settings=ModelSettings(), + tools=[], + output_schema=None, + handoffs=[], + tracing=ModelTracing.DISABLED, + ) + ) + await ws1.close_started.wait() + await model.close() + ws1.release_close.set() + + with pytest.raises(RuntimeError): + await request_task + + assert opened == [ws1] + assert ws2.sent_messages == [] + assert model._ws_connection is None + + @pytest.mark.allow_call_model_methods @pytest.mark.asyncio async def test_websocket_model_does_not_retry_if_send_raises_after_writing_on_reused_connection( @@ -4214,6 +4511,65 @@ def test_websocket_get_retry_advice_marks_connect_timeout_replay_safe() -> None: assert advice.replay_safety == "safe" +@pytest.mark.allow_call_model_methods +def test_websocket_get_retry_advice_marks_handshake_failure_replay_safe() -> None: + model = OpenAIResponsesWSModel(model="gpt-4", openai_client=cast(Any, DummyWSClient())) + error = _invalid_message_error( + "did not receive a valid HTTP response", + cause=EOFError("connection closed while reading HTTP status line"), + ) + + advice = model.get_retry_advice( + ModelRetryAdviceRequest( + error=error, + attempt=1, + stream=True, + previous_response_id="resp_prev", + ) + ) + + assert advice is not None + assert advice.suggested is True + assert advice.replay_safety == "safe" + + +@pytest.mark.allow_call_model_methods +def test_websocket_get_retry_advice_ignores_malformed_handshake() -> None: + model = OpenAIResponsesWSModel(model="gpt-4", openai_client=cast(Any, DummyWSClient())) + error = _invalid_message_error("malformed HTTP status line") + + advice = model.get_retry_advice( + ModelRetryAdviceRequest( + error=error, + attempt=1, + stream=True, + previous_response_id="resp_prev", + ) + ) + + assert advice is None + + +@pytest.mark.allow_call_model_methods +def test_websocket_get_retry_advice_marks_close_invalidation_non_retryable() -> None: + model = OpenAIResponsesWSModel(model="gpt-4", openai_client=cast(Any, DummyWSClient())) + error = RuntimeError("Responses websocket connection closed while establishing a connection.") + setattr(error, "_openai_agents_ws_close_invalidated", True) # noqa: B010 + + advice = model.get_retry_advice( + ModelRetryAdviceRequest( + error=error, + attempt=1, + stream=False, + ) + ) + + assert advice is not None + assert advice.suggested is False + assert advice.normalized is not None + assert advice.normalized.is_abort is True + + @pytest.mark.allow_call_model_methods def test_websocket_get_retry_advice_marks_request_lock_timeout_replay_safe() -> None: model = OpenAIResponsesWSModel(model="gpt-4", openai_client=cast(Any, DummyWSClient()))