From 2c8fd38f2e169cf07496780c048b1f465cd75333 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Sun, 19 Jul 2026 09:00:27 -0400 Subject: [PATCH] fix(observability): defer logical LLM completion until validation Signed-off-by: Alex Fournier --- agent/chat_completion_helpers.py | 3 ++ agent/codex_runtime.py | 1 + agent/conversation_loop.py | 8 ++++ agent/relay_llm.py | 35 ++++++++++++-- tests/agent/test_relay_llm.py | 79 +++++++++++++++++++++++++++++++ tests/run_agent/test_run_agent.py | 44 +++++++++++++++++ 6 files changed, 165 insertions(+), 5 deletions(-) diff --git a/agent/chat_completion_helpers.py b/agent/chat_completion_helpers.py index da1ae4c70ea3..2dded16d8ddd 100644 --- a/agent/chat_completion_helpers.py +++ b/agent/chat_completion_helpers.py @@ -2418,6 +2418,7 @@ def interruptible_streaming_api_call(agent, api_kwargs: dict, *, on_first_delta= else "primary" ), }, + defer_logical_completion=True, ) streamed_response = stream_converse_with_callbacks( {"stream": stream}, @@ -2805,6 +2806,7 @@ def interruptible_streaming_api_call(agent, api_kwargs: dict, *, on_first_delta= else "primary" ), }, + defer_logical_completion=True, ) for chunk in stream: last_chunk_time["t"] = time.time() @@ -3220,6 +3222,7 @@ def interruptible_streaming_api_call(agent, api_kwargs: dict, *, on_first_delta= else "primary" ), }, + defer_logical_completion=True, ) try: for event in stream: diff --git a/agent/codex_runtime.py b/agent/codex_runtime.py index e0ea0680808d..65e80f1d8680 100644 --- a/agent/codex_runtime.py +++ b/agent/codex_runtime.py @@ -1270,6 +1270,7 @@ def run_codex_stream(agent, api_kwargs: dict, client: Any = None, on_first_delta ), "retry_count": attempt, }, + defer_logical_completion=True, ) def _interrupt_or_superseded() -> bool: diff --git a/agent/conversation_loop.py b/agent/conversation_loop.py index 3484c40eec8a..aee7b7529f29 100644 --- a/agent/conversation_loop.py +++ b/agent/conversation_loop.py @@ -1465,6 +1465,7 @@ def run_conversation( ), "retry_count": retry_count, }, + defer_logical_completion=True, ) from hermes_cli.middleware import run_llm_execution_middleware @@ -1745,6 +1746,13 @@ def run_conversation( ) continue # Retry the API call + from agent import relay_llm + + relay_llm.complete_logical_call( + api_request_id, + outcome="success", + ) + # Check finish_reason before proceeding if agent.api_mode == "codex_responses": status = getattr(response, "status", None) diff --git a/agent/relay_llm.py b/agent/relay_llm.py index 9d9195fb83c0..65e900c1d9fb 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -28,6 +28,7 @@ def execute( name: str, model_name: str, metadata: dict[str, Any] | None = None, + defer_logical_completion: bool = False, ) -> Any: """Run one non-streaming physical provider attempt through Relay.""" runtime, session, parent = relay_runtime.resolve_execution_context(session_id) @@ -81,10 +82,10 @@ def execute( raise callback_error raise - if "value" in raw_response and _json_equal(managed, raw_response["json"]): + if not defer_logical_completion: _complete_logical(logical, outcome="success") + if "value" in raw_response and _json_equal(managed, raw_response["json"]): return raw_response["value"] - _complete_logical(logical, outcome="success") return managed @@ -96,6 +97,7 @@ async def execute_async( name: str, model_name: str, metadata: dict[str, Any] | None = None, + defer_logical_completion: bool = False, ) -> Any: """Run one asynchronous physical provider attempt through Relay.""" runtime, session, parent = relay_runtime.resolve_execution_context(session_id) @@ -147,7 +149,8 @@ async def execute_async( raise callback_error raise - _complete_logical(logical, outcome="success") + if not defer_logical_completion: + _complete_logical(logical, outcome="success") if "value" in raw_response and _json_equal(managed, raw_response["json"]): return raw_response["value"] return managed @@ -160,6 +163,7 @@ def execute_current( name: str, model_name: str, metadata: dict[str, Any] | None = None, + defer_logical_completion: bool = False, ) -> Any: """Run a provider attempt under the inherited Hermes turn when present.""" turn = relay_runtime.active_turn() @@ -172,6 +176,7 @@ def execute_current( name=name, model_name=model_name, metadata=metadata, + defer_logical_completion=defer_logical_completion, ) @@ -182,6 +187,7 @@ async def execute_current_async( name: str, model_name: str, metadata: dict[str, Any] | None = None, + defer_logical_completion: bool = False, ) -> Any: """Run an async provider attempt under the inherited turn when present.""" turn = relay_runtime.active_turn() @@ -194,6 +200,7 @@ async def execute_current_async( name=name, model_name=model_name, metadata=metadata, + defer_logical_completion=defer_logical_completion, ) @@ -205,6 +212,7 @@ def stream_current( model_name: str, finalizer: Callable[[], Any], metadata: dict[str, Any] | None = None, + defer_logical_completion: bool = False, ) -> Any: """Run a provider stream under the inherited Hermes turn when present.""" turn = relay_runtime.active_turn() @@ -218,6 +226,7 @@ def stream_current( model_name=model_name, finalizer=finalizer, metadata=metadata, + defer_logical_completion=defer_logical_completion, ) @@ -235,6 +244,7 @@ def stream( accept_chunk: Callable[[Any], bool] | None = None, completed_response_predicate: Callable[[Any], bool] | None = None, metadata: dict[str, Any] | None = None, + defer_logical_completion: bool = False, ) -> "ManagedLlmStream": """Return a synchronous view of one Relay-managed provider stream.""" return ManagedLlmStream( @@ -250,6 +260,7 @@ def stream( accept_chunk=accept_chunk, completed_response_predicate=completed_response_predicate, metadata=metadata, + defer_logical_completion=defer_logical_completion, ) @@ -271,6 +282,7 @@ class ManagedLlmStream(Iterator[Any]): accept_chunk: Callable[[Any], bool] | None, completed_response_predicate: Callable[[Any], bool] | None, metadata: dict[str, Any] | None, + defer_logical_completion: bool, ) -> None: self.final_response: Any = None self._loop: asyncio.AbstractEventLoop | None = None @@ -278,6 +290,7 @@ class ManagedLlmStream(Iterator[Any]): self._closed = False self._callback_error: BaseException | None = None self._logical: tuple[relay_runtime.RelayTurnContext, Any, str] | None = None + self._defer_logical_completion = defer_logical_completion self._on_chunk = on_chunk self._chunk_adapter = chunk_adapter or _namespace self._accept_chunk = accept_chunk @@ -392,8 +405,9 @@ class ManagedLlmStream(Iterator[Any]): try: chunk = self._loop.run_until_complete(next_chunk()) except StopAsyncIteration: - _complete_logical(self._logical, outcome="success") - self._logical = None + if not self._defer_logical_completion: + _complete_logical(self._logical, outcome="success") + self._logical = None self.close() raise StopIteration from None except BaseException as exc: @@ -619,6 +633,17 @@ def _complete_logical( turn.logical_llm_calls.pop(request_id, None) +def complete_logical_call(api_request_id: str, *, outcome: str) -> None: + """Complete the active turn's logical LLM call after caller validation.""" + turn = relay_runtime.active_turn() + if turn is None or not api_request_id: + return + with turn.logical_llm_lock: + handle = turn.logical_llm_calls.get(api_request_id) + if handle is not None: + _complete_logical((turn, handle, api_request_id), outcome=outcome) + + def _provider_request( original: dict[str, Any], request: Any, diff --git a/tests/agent/test_relay_llm.py b/tests/agent/test_relay_llm.py index 71eb26c81b93..a62d192ff39b 100644 --- a/tests/agent/test_relay_llm.py +++ b/tests/agent/test_relay_llm.py @@ -179,6 +179,40 @@ def test_non_stream_preserves_raw_provider_response_identity(relay_turn): assert result is raw_response +def test_non_stream_defers_logical_success_and_reuses_scope_for_retry(relay_turn): + _relay, turn = relay_turn + metadata = {"api_mode": "custom", "api_request_id": "request-retry"} + + first = relay_llm.execute( + {"model": "test-model", "messages": []}, + lambda _request: {"content": "invalid"}, + session_id="session-1", + name="test-provider", + model_name="test-model", + metadata=metadata, + defer_logical_completion=True, + ) + first_handle = turn.logical_llm_calls["request-retry"] + + second = relay_llm.execute( + {"model": "test-model", "messages": []}, + lambda _request: {"content": "valid"}, + session_id="session-1", + name="test-provider", + model_name="test-model", + metadata=metadata, + defer_logical_completion=True, + ) + + assert first == {"content": "invalid"} + assert second == {"content": "valid"} + assert turn.logical_llm_calls == {"request-retry": first_handle} + + relay_llm.complete_logical_call("request-retry", outcome="success") + + assert turn.logical_llm_calls == {} + + def test_non_stream_result_survives_logical_scope_close_failure( relay_turn, monkeypatch ): @@ -230,6 +264,51 @@ async def test_async_non_stream_preserves_raw_provider_response_identity(relay_t assert result is raw_response +@pytest.mark.asyncio +async def test_async_non_stream_defers_logical_success_for_validation(relay_turn): + _relay, turn = relay_turn + + async def provider(_request): + return {"content": "pending-validation"} + + await relay_llm.execute_current_async( + {"model": "test-model", "messages": []}, + provider, + name="test-provider", + model_name="test-model", + metadata={"api_mode": "custom", "api_request_id": "request-async-defer"}, + defer_logical_completion=True, + ) + + assert "request-async-defer" in turn.logical_llm_calls + + relay_llm.complete_logical_call("request-async-defer", outcome="success") + + assert turn.logical_llm_calls == {} + + +def test_stream_defers_logical_success_for_response_validation(relay_turn): + _relay, turn = relay_turn + + stream = relay_llm.stream( + {"model": "test-model", "messages": []}, + lambda _request: iter([{"delta": "pending-validation"}]), + session_id="session-1", + name="test-provider", + model_name="test-model", + finalizer=lambda: {"content": "pending-validation"}, + metadata={"api_mode": "custom", "api_request_id": "request-stream-defer"}, + defer_logical_completion=True, + ) + + assert list(stream) == [{"delta": "pending-validation"}] + assert "request-stream-defer" in turn.logical_llm_calls + + relay_llm.complete_logical_call("request-stream-defer", outcome="success") + + assert turn.logical_llm_calls == {} + + def test_current_attempt_bypasses_relay_without_an_active_turn(monkeypatch): monkeypatch.setattr(relay_runtime, "current_turn", lambda: None) request = {"model": "test-model", "messages": []} diff --git a/tests/run_agent/test_run_agent.py b/tests/run_agent/test_run_agent.py index a34a22c1515c..9d8762b37c9d 100644 --- a/tests/run_agent/test_run_agent.py +++ b/tests/run_agent/test_run_agent.py @@ -6080,6 +6080,50 @@ class TestRetryExhaustion: assert "Invalid API response" in result["error"] assert result.get("final_response") == result["error"] + def test_invalid_response_retry_completes_one_logical_call(self, agent): + self._setup_agent(agent) + agent.client.chat.completions.create.side_effect = [ + SimpleNamespace(choices=[], model="test/model", usage=None), + _mock_response(content="recovered"), + ] + relay_attempts = [] + logical_completions = [] + + def execute(request, callback, **kwargs): + relay_attempts.append(kwargs) + return callback(request) + + from agent import conversation_loop as _conv_loop + + with ( + patch.object(agent, "_persist_session"), + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + patch("run_agent.time", self._make_fast_time_mock()), + patch.object(_conv_loop, "time", self._make_fast_time_mock()), + patch.object(_conv_loop, "jittered_backoff", lambda *a, **k: 0.0), + patch("agent.relay_llm.execute", side_effect=execute), + patch( + "agent.relay_llm.complete_logical_call", + side_effect=lambda request_id, *, outcome: logical_completions.append( + (request_id, outcome) + ), + ), + ): + result = agent.run_conversation("hello") + + assert result["completed"] is True + assert len(relay_attempts) == 2 + assert all( + attempt["defer_logical_completion"] is True + for attempt in relay_attempts + ) + request_ids = { + attempt["metadata"]["api_request_id"] for attempt in relay_attempts + } + assert len(request_ids) == 1 + assert logical_completions == [(request_ids.pop(), "success")] + def test_content_filter_refusal_surfaced_not_retried(self, agent): """A model refusal must be surfaced immediately, NOT laundered into the empty-response retry loop and reported as "rate limited" / "no