diff --git a/agent/codex_runtime.py b/agent/codex_runtime.py index 2e6e9d56bb..a1b3e8fc12 100644 --- a/agent/codex_runtime.py +++ b/agent/codex_runtime.py @@ -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: