fix(observability): defer logical LLM completion until validation

Signed-off-by: Alex Fournier <afournier@nvidia.com>
This commit is contained in:
Alex Fournier 2026-07-19 09:00:27 -04:00
parent 55b7903cf7
commit 2c8fd38f2e
6 changed files with 165 additions and 5 deletions

View file

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

View file

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

View file

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

View file

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

View file

@ -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": []}

View file

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