diff --git a/agent/codex_runtime.py b/agent/codex_runtime.py index cdce3f87c1..44a35410f5 100644 --- a/agent/codex_runtime.py +++ b/agent/codex_runtime.py @@ -8,6 +8,7 @@ import contextvars import json import logging import os +import threading import time from contextlib import suppress from types import SimpleNamespace @@ -23,6 +24,8 @@ _codex_watchdog_state_var: contextvars.ContextVar[Any | None] = contextvars.Cont "codex_watchdog_state", default=None ) +_CODEX_POST_TERMINAL_DRAIN_TIMEOUT_SECONDS = 2.0 + def _call_guarded(fn: Callable | None, fail_msg: str, *fail_args: Any, args: tuple = (), kwargs: dict | None = None): """Invoke an optional display/debug callback; a buggy hook must never tear down the turn.""" @@ -981,15 +984,32 @@ def run_codex_stream(agent, api_kwargs: dict, client: Any = None, on_first_delta def _drain_for_finalizer(event_stream: Any) -> None: # ``final`` is already assembled; draining only lets Relay run its finalizer. A transport error # here must NOT discard the completed, already-billed response. - try: - for _ignored in event_stream: - pass - except (*transport_errors, _APIConnectionError) as exc: - if not isinstance(exc, transport_errors): - _log_failure(exc) - logger.warning("Codex Responses stream transport finalization failed after a terminal response was already " - "received; returning the completed response instead of retrying. %s error=%s", - agent._client_log_context(), exc) + drained = threading.Event() + + def _drain() -> None: + try: + for _ignored in event_stream: + pass + except (*transport_errors, _APIConnectionError) as exc: + if not isinstance(exc, transport_errors): + _log_failure(exc) + logger.warning("Codex Responses stream transport finalization failed after a terminal response was already " + "received; returning the completed response instead of retrying. %s error=%s", + agent._client_log_context(), exc) + except Exception: + logger.debug("Codex Responses stream finalization failed after a terminal response", exc_info=True) + finally: + drained.set() + + threading.Thread(target=_drain, name="codex-post-terminal-drain", daemon=True).start() + if drained.wait(_CODEX_POST_TERMINAL_DRAIN_TIMEOUT_SECONDS): + return + logger.warning( + "Codex Responses stream remained open %.1fs after a terminal response; closing it and returning the " + "completed response instead of retrying. %s", + _CODEX_POST_TERMINAL_DRAIN_TIMEOUT_SECONDS, agent._client_log_context(), + ) + _close_event_stream(event_stream) def _close_event_stream(event_stream: Any) -> None: close_fn = getattr(event_stream, "close", None) # None while connect never succeeded diff --git a/tests/agent/test_run_agent_codex_responses.py b/tests/agent/test_run_agent_codex_responses.py index 647a7b2c56..9ad2717662 100644 --- a/tests/agent/test_run_agent_codex_responses.py +++ b/tests/agent/test_run_agent_codex_responses.py @@ -1080,6 +1080,68 @@ def test_run_codex_stream_returns_terminal_response_when_post_terminal_drain_fai ) +def test_run_codex_stream_bounds_post_terminal_drain(monkeypatch): + """A relay that keeps SSE open after completion cannot discard the billed response.""" + import threading + import time + + import agent.codex_runtime as codex_runtime + + agent = _build_agent(monkeypatch) + message_item = SimpleNamespace( + type="message", + status="completed", + content=[SimpleNamespace(type="output_text", text="All done.")], + ) + usage = SimpleNamespace(input_tokens=10, output_tokens=6, total_tokens=16) + closed = threading.Event() + + class _HeldOpenAfterTerminalStream: + def __init__(self): + self._events = iter([ + SimpleNamespace(type="response.output_item.done", item=message_item), + SimpleNamespace( + type="response.completed", + response=SimpleNamespace( + status="completed", usage=usage, id="resp_held_open", + ), + ), + ]) + + def __iter__(self): + return self + + def __next__(self): + try: + return next(self._events) + except StopIteration: + closed.wait(3.0) + raise + + def close(self): + closed.set() + + calls = {"count": 0} + + def _fake_create(**kwargs): + calls["count"] += 1 + return _HeldOpenAfterTerminalStream() + + agent.client = SimpleNamespace(responses=SimpleNamespace(create=_fake_create)) + monkeypatch.setattr(codex_runtime, "_CODEX_POST_TERMINAL_DRAIN_TIMEOUT_SECONDS", 0.01) + + started = time.monotonic() + response = agent._run_codex_stream(_codex_request_kwargs()) + elapsed = time.monotonic() - started + + assert elapsed < 2.0 + assert calls["count"] == 1 + assert response.status == "completed" + assert response.usage is usage + assert response.id == "resp_held_open" + assert closed.wait(1.0) + + def test_run_conversation_codex_plain_text(monkeypatch): agent = _build_agent(monkeypatch) monkeypatch.setattr(agent, "_interruptible_api_call", lambda api_kwargs: _codex_message_response("OK"))