diff --git a/agent/codex_runtime.py b/agent/codex_runtime.py index f6100b542f..b67fe5fbf3 100644 --- a/agent/codex_runtime.py +++ b/agent/codex_runtime.py @@ -991,6 +991,7 @@ def run_codex_stream(agent, api_kwargs: dict, client: Any = None, on_first_delta def _codex_stream_created(_raw_stream: Any) -> None: # Claim the delta sink for THIS attempt; a newer attempt supersedes this token. writer_token["value"] = claim_stream_writer(agent) + writer_token["raw_stream"] = _raw_stream def _accept_codex_chunk(_chunk: Any) -> bool: token = writer_token["value"] @@ -1031,6 +1032,11 @@ def run_codex_stream(agent, api_kwargs: dict, client: Any = None, on_first_delta "closing it and returning the completed response instead of retrying. %s", budget, agent._client_log_context(), ) + # Under a live Relay loop the managed wrapper's close() cannot reach the provider response + # (the loop is still running the drain); close the raw stream captured at stream creation too. + raw_stream = writer_token.get("raw_stream") + if raw_stream is not None and raw_stream is not event_stream: + _close_event_stream(raw_stream) _close_event_stream(event_stream) def _close_event_stream(event_stream: Any) -> None: @@ -1060,7 +1066,7 @@ def run_codex_stream(agent, api_kwargs: dict, client: Any = None, on_first_delta with watchdog_state.lock: watchdog_state.retry_started_ts = time.time() intercepted_events: list = [] - writer_token["value"] = event_stream = None + writer_token["value"] = writer_token["raw_stream"] = event_stream = None try: try: event_stream = relay_llm.stream( diff --git a/tests/agent/test_run_agent_codex_responses.py b/tests/agent/test_run_agent_codex_responses.py index 0825463432..8856916227 100644 --- a/tests/agent/test_run_agent_codex_responses.py +++ b/tests/agent/test_run_agent_codex_responses.py @@ -1142,6 +1142,68 @@ def test_run_codex_stream_bounds_post_terminal_drain(monkeypatch): assert closed.wait(1.0) +def test_run_codex_stream_drain_timeout_closes_raw_stream_when_managed_close_raises(monkeypatch): + """A Relay-managed wrapper whose close() raises (running loop) must not leak the provider stream.""" + import threading + + import agent.codex_runtime as codex_runtime + from agent import relay_llm + + 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) + raw_closed = threading.Event() + + class _HeldOpenRawStream: + 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_managed")), + ]) + + def __iter__(self): + return self + + def __next__(self): + try: + return next(self._events) + except StopIteration: + raw_closed.wait(3.0) + raise + + def close(self): + raw_closed.set() + + class _ManagedWrapper: + final_response = None + + def __init__(self, request, stream_factory, *, on_stream_created=None, **_kwargs): + raw = stream_factory(request) + on_stream_created(raw) + self._iter = iter(raw) + + def __iter__(self): + return self + + def __next__(self): + return next(self._iter) + + def close(self): + raise RuntimeError("Cannot close a running event loop") + + agent.client = SimpleNamespace(responses=SimpleNamespace(create=lambda **kwargs: _HeldOpenRawStream())) + monkeypatch.setattr(relay_llm, "stream", _ManagedWrapper) + monkeypatch.setattr(codex_runtime, "_stream_drain_timeout", lambda: 0.01) + + response = agent._run_codex_stream(_codex_request_kwargs()) + + assert response.id == "resp_managed" + assert raw_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"))