fix(codex): bound post-terminal stream drain

This commit is contained in:
fangliquanflq
2026-09-06 04:23:58 +08:00
committed by Teknium
parent 614f11d3f6
commit 39884274d6
2 changed files with 91 additions and 9 deletions

View File

@@ -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

View File

@@ -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"))