mirror of
https://github.com/NousResearch/hermes-agent.git
synced 2026-07-29 18:46:59 +00:00
fix(observability): defer logical LLM completion until validation
Signed-off-by: Alex Fournier <afournier@nvidia.com>
This commit is contained in:
parent
55b7903cf7
commit
2c8fd38f2e
6 changed files with 165 additions and 5 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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": []}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue