refactor(agent/codex_runtime): flatten stream-assembler dispatch tables, extra_body merge, and relay stream call

This commit is contained in:
Teknium
2026-09-02 22:38:16 -07:00
parent 894abab3d9
commit 30a1acba47
+20 -41
View File
@@ -498,11 +498,8 @@ def run_codex_app_server_turn(
)
return _turn_result(
interrupt, messages, api_calls=1, completed=not turn.interrupted and turn.error is None, error=turn.error,
final_response=turn.final_text,
# We flushed the projected rows ourselves; the gateway must skip its own DB write.
agent_persisted=True,
codex_thread_id=turn.thread_id,
codex_turn_id=turn.turn_id,
# We flushed the projected rows ourselves (agent_persisted); the gateway must skip its own DB write.
final_response=turn.final_text, agent_persisted=True, codex_thread_id=turn.thread_id, codex_turn_id=turn.turn_id,
**usage_result,
)
@@ -560,10 +557,9 @@ def _message_phase(item: Any) -> str | None:
def _output_text_of(item: Any) -> str:
"""Concatenated ``output_text`` parts of a message item ("" if content is not a list)."""
content_parts = _event_field(item, "content", [])
if not isinstance(content_parts, list):
return ""
parts = content_parts if isinstance(content_parts, list) else []
return "".join(
str(_event_field(part, "text", "") or "") for part in content_parts if _event_field(part, "type", "") == "output_text"
str(_event_field(part, "text", "") or "") for part in parts if _event_field(part, "type", "") == "output_text"
).strip()
@@ -656,8 +652,7 @@ class _CodexResponseAssembler:
elif event_type.endswith("function_call_arguments.done"):
# Authoritative for the accumulated string; an explicit "" (zero-arg
# call) counts, only a missing field keeps the streamed deltas.
done_args = _event_field(event, "arguments", None)
if done_args is not None:
if (done_args := _event_field(event, "arguments", None)) is not None:
pending["arguments"] = str(done_args)
def _on_reasoning_delta(self, event: Any, event_type: str) -> None:
@@ -716,8 +711,7 @@ class _CodexResponseAssembler:
"response.completed": _on_terminal, "response.incomplete": _on_terminal, "response.failed": _on_terminal,
}
_FUZZY_HANDLERS = (
(lambda t: "output_text.delta" in t, _on_text_delta),
(lambda t: "function_call" in t, _on_function_call),
(lambda t: "output_text.delta" in t, _on_text_delta), (lambda t: "function_call" in t, _on_function_call),
(lambda t: "reasoning" in t and "delta" in t, _on_reasoning_delta),
)
@@ -725,9 +719,7 @@ class _CodexResponseAssembler:
"""Process one event; True when the stream hit a terminal frame."""
event_type = _event_field(event, "type", "")
event_type = event_type if isinstance(event_type, str) else ""
handler = self._EXACT_HANDLERS.get(event_type) or next(
(h for matches, h in self._FUZZY_HANDLERS if matches(event_type)), None
)
handler = self._EXACT_HANDLERS.get(event_type) or next((h for m, h in self._FUZZY_HANDLERS if m(event_type)), None)
return bool(handler(self, event, event_type)) if handler is not None else False
def _settled_output(self) -> List[Any]:
@@ -857,20 +849,15 @@ def _bypass_sdk_request_transform(stream_kwargs: dict) -> dict:
walk. HERMES_CODEX_SDK_TRANSFORM=1 disables."""
if os.environ.get("HERMES_CODEX_SDK_TRANSFORM", "").strip().lower() in {"1", "true", "yes", "on"}:
return stream_kwargs
moved = {
field: stream_kwargs[field]
for field in _SDK_TRANSFORM_BYPASS_FIELDS
if isinstance(stream_kwargs.get(field), (dict, list)) and _is_plain_json_data(stream_kwargs[field])
}
moved = {f: stream_kwargs[f] for f in _SDK_TRANSFORM_BYPASS_FIELDS
if isinstance(stream_kwargs.get(f), (dict, list)) and _is_plain_json_data(stream_kwargs[f])}
if not moved:
return stream_kwargs
bypassed = {key: value for key, value in stream_kwargs.items() if key not in moved}
extra_body = bypassed.get("extra_body")
merged = dict(extra_body) if isinstance(extra_body, dict) else {}
for field, value in moved.items():
# An explicit caller-provided extra_body entry keeps precedence (SDK post-transform merge).
merged.setdefault(field, value)
bypassed["extra_body"] = merged
# An explicit caller-provided extra_body entry keeps precedence (SDK post-transform merge).
bypassed["extra_body"] = {**merged, **{f: v for f, v in moved.items() if f not in merged}}
return bypassed
@@ -882,8 +869,7 @@ def run_codex_stream(agent, api_kwargs: dict, client: Any = None, on_first_delta
from agent import relay_llm
transport_errors = (_httpx.RemoteProtocolError, _httpx.ReadTimeout, _httpx.ConnectError, ConnectionError)
active_client = client or agent._ensure_primary_openai_client(reason="codex_stream_direct")
max_stream_retries = 1
model = api_kwargs.get("model")
max_stream_retries, model = 1, api_kwargs.get("model")
# Accumulate streamed text so callers / compat shims can read it.
agent._codex_streamed_text_parts: list = []
# Retirement token for THIS request (installed by ``interruptible_api_call``). A
@@ -961,10 +947,9 @@ def run_codex_stream(agent, api_kwargs: dict, client: Any = None, on_first_delta
def _close_event_stream(event_stream: Any) -> None:
close_fn = getattr(event_stream, "close", None) # None while connect never succeeded
if not callable(close_fn):
return
try:
close_fn()
if callable(close_fn):
close_fn()
except Exception:
# A failed close can leave this connection checked out of the httpx pool while
# the caller reuse-caches the client; poison the slot so close really closes the
@@ -979,21 +964,16 @@ def run_codex_stream(agent, api_kwargs: dict, client: Any = None, on_first_delta
if agent._interrupt_requested:
raise InterruptedError("Agent interrupted before Codex stream retry")
intercepted_events: list = []
writer_token["value"] = None
event_stream = None
writer_token["value"] = event_stream = None
try:
try:
event_stream = relay_llm.stream(
dict(api_kwargs),
_open_codex_stream,
dict(api_kwargs), _open_codex_stream,
session_id=str(getattr(agent, "session_id", "") or ""),
name=str(getattr(agent, "provider", "") or "codex"),
model_name=str(model or ""),
name=str(getattr(agent, "provider", "") or "codex"), model_name=str(model or ""),
finalizer=lambda: _consume_codex_event_stream(list(intercepted_events), model=model),
on_stream_created=_codex_stream_created,
on_chunk=intercepted_events.append,
chunk_adapter=lambda chunk: chunk,
accept_chunk=_accept_codex_chunk,
on_stream_created=_codex_stream_created, on_chunk=intercepted_events.append,
chunk_adapter=lambda chunk: chunk, accept_chunk=_accept_codex_chunk,
completed_response_predicate=lambda r: bool(hasattr(r, "output") and not hasattr(r, "__iter__")),
metadata={
"api_mode": "codex_responses", "call_role": call_role, "retry_count": attempt,
@@ -1033,8 +1013,7 @@ def run_codex_stream(agent, api_kwargs: dict, client: Any = None, on_first_delta
"Codex Responses stream terminal status=%s "
"(incomplete_details=%s, error=%s, streamed_chars=%d). %s",
final.status, final.incomplete_details, final.error,
sum(len(p) for p in agent._codex_streamed_text_parts),
agent._client_log_context(),
sum(len(p) for p in agent._codex_streamed_text_parts), agent._client_log_context(),
)
return final
finally: