mirror of
https://github.com/NousResearch/hermes-agent.git
synced 2026-07-29 18:46:59 +00:00
fix(relay): preserve managed cancellation signals
Signed-off-by: Alex Fournier <afournier@nvidia.com>
This commit is contained in:
parent
fb1e417367
commit
06ed41aa1f
4 changed files with 163 additions and 7 deletions
|
|
@ -87,6 +87,7 @@ def execute(
|
|||
raise callback_error
|
||||
if _recover_successful_callback(
|
||||
raw_response,
|
||||
relay_error=exc,
|
||||
callback_error=callback_error,
|
||||
logical=logical,
|
||||
defer_logical_completion=defer_logical_completion,
|
||||
|
|
@ -169,6 +170,7 @@ async def execute_async(
|
|||
raise callback_error
|
||||
if _recover_successful_callback(
|
||||
raw_response,
|
||||
relay_error=exc,
|
||||
callback_error=callback_error,
|
||||
logical=logical,
|
||||
defer_logical_completion=defer_logical_completion,
|
||||
|
|
@ -433,8 +435,12 @@ class ManagedLlmStream(Iterator[Any]):
|
|||
response_codec=_codec(runtime.relay, metadata),
|
||||
)
|
||||
)
|
||||
except BaseException:
|
||||
if self._provider_completed and self._callback_error is None:
|
||||
except BaseException as exc:
|
||||
if (
|
||||
isinstance(exc, Exception)
|
||||
and self._provider_completed
|
||||
and self._callback_error is None
|
||||
):
|
||||
logger.warning(
|
||||
"NeMo Relay stream post-processing failed after provider success; "
|
||||
"preserving the provider result",
|
||||
|
|
@ -448,7 +454,10 @@ class ManagedLlmStream(Iterator[Any]):
|
|||
self._stream = iter(())
|
||||
return
|
||||
if not self._defer_logical_completion:
|
||||
_complete_logical(self._logical, outcome="failed")
|
||||
_complete_logical(
|
||||
self._logical,
|
||||
outcome="cancelled" if _is_cancellation(exc) else "failed",
|
||||
)
|
||||
self._logical = None
|
||||
loop.close()
|
||||
self._loop = None
|
||||
|
|
@ -492,7 +501,11 @@ class ManagedLlmStream(Iterator[Any]):
|
|||
):
|
||||
self._close(logical_outcome="failed")
|
||||
raise callback_error
|
||||
if self._provider_completed and callback_error is None:
|
||||
if (
|
||||
isinstance(exc, Exception)
|
||||
and self._provider_completed
|
||||
and callback_error is None
|
||||
):
|
||||
logger.warning(
|
||||
"NeMo Relay stream post-processing failed after provider success; "
|
||||
"preserving the provider result",
|
||||
|
|
@ -500,7 +513,9 @@ class ManagedLlmStream(Iterator[Any]):
|
|||
)
|
||||
self._close(logical_outcome="success")
|
||||
raise StopIteration from None
|
||||
self._close(logical_outcome="failed")
|
||||
self._close(
|
||||
logical_outcome="cancelled" if _is_cancellation(exc) else "failed"
|
||||
)
|
||||
raise
|
||||
if not self._relay_observes_chunks and self._on_chunk is not None:
|
||||
self._on_chunk(chunk)
|
||||
|
|
@ -751,11 +766,16 @@ def _complete_logical(
|
|||
def _recover_successful_callback(
|
||||
raw_response: dict[str, Any],
|
||||
*,
|
||||
relay_error: BaseException,
|
||||
callback_error: BaseException | None,
|
||||
logical: tuple[relay_runtime.RelayTurnContext, Any, str] | None,
|
||||
defer_logical_completion: bool,
|
||||
) -> bool:
|
||||
if callback_error is not None or "value" not in raw_response:
|
||||
if (
|
||||
not isinstance(relay_error, Exception)
|
||||
or callback_error is not None
|
||||
or "value" not in raw_response
|
||||
):
|
||||
return False
|
||||
logger.warning(
|
||||
"NeMo Relay LLM post-processing failed after provider success; "
|
||||
|
|
@ -767,6 +787,13 @@ def _recover_successful_callback(
|
|||
return True
|
||||
|
||||
|
||||
def _is_cancellation(error: BaseException) -> bool:
|
||||
return isinstance(
|
||||
error,
|
||||
(asyncio.CancelledError, InterruptedError, KeyboardInterrupt),
|
||||
)
|
||||
|
||||
|
||||
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()
|
||||
|
|
|
|||
|
|
@ -63,7 +63,11 @@ def execute(
|
|||
and relay_runtime._is_relay_wrapped_callback_error(exc, callback_error)
|
||||
):
|
||||
raise callback_error
|
||||
if callback_error is None and "value" in raw_result:
|
||||
if (
|
||||
isinstance(exc, Exception)
|
||||
and callback_error is None
|
||||
and "value" in raw_result
|
||||
):
|
||||
logger.warning(
|
||||
"NeMo Relay tool post-processing failed after dispatch success; "
|
||||
"returning the Hermes tool result",
|
||||
|
|
|
|||
|
|
@ -508,6 +508,34 @@ def test_non_stream_returns_provider_response_after_relay_post_processing_failur
|
|||
assert "returning the provider response" in caplog.text
|
||||
|
||||
|
||||
def test_non_stream_does_not_swallow_interrupt_after_provider_success(
|
||||
relay_turn, monkeypatch
|
||||
):
|
||||
relay, turn = relay_turn
|
||||
|
||||
async def interrupt_after_callback(_name, request, callback, **_kwargs):
|
||||
callback(request)
|
||||
raise KeyboardInterrupt
|
||||
|
||||
monkeypatch.setattr(relay.llm, "execute", interrupt_after_callback)
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
relay_llm.execute(
|
||||
{"model": "test-model", "messages": []},
|
||||
lambda _request: {"content": "already returned"},
|
||||
session_id="session-1",
|
||||
name="test-provider",
|
||||
model_name="test-model",
|
||||
metadata={
|
||||
"api_mode": "custom",
|
||||
"api_request_id": "request-post-interrupt",
|
||||
},
|
||||
)
|
||||
|
||||
assert "request-post-interrupt" in turn.logical_llm_calls
|
||||
relay_llm.complete_logical_call("request-post-interrupt", outcome="cancelled")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_non_stream_preserves_raw_provider_response_identity(relay_turn):
|
||||
_relay, _turn = relay_turn
|
||||
|
|
@ -560,6 +588,40 @@ async def test_async_non_stream_returns_provider_response_after_relay_failure(
|
|||
assert "returning the provider response" in caplog.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_non_stream_does_not_swallow_cancellation_after_provider_success(
|
||||
relay_turn, monkeypatch
|
||||
):
|
||||
relay, turn = relay_turn
|
||||
|
||||
async def provider(_request):
|
||||
return {"content": "already returned"}
|
||||
|
||||
async def cancel_after_callback(_name, request, callback, **_kwargs):
|
||||
await callback(request)
|
||||
raise asyncio.CancelledError
|
||||
|
||||
monkeypatch.setattr(relay.llm, "execute", cancel_after_callback)
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
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-post-cancel",
|
||||
},
|
||||
)
|
||||
|
||||
assert "request-async-post-cancel" in turn.logical_llm_calls
|
||||
relay_llm.complete_logical_call(
|
||||
"request-async-post-cancel",
|
||||
outcome="cancelled",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_non_stream_defers_logical_success_for_validation(relay_turn):
|
||||
_relay, turn = relay_turn
|
||||
|
|
@ -628,6 +690,49 @@ def test_stream_finishes_after_relay_post_processing_failure(
|
|||
assert "preserving the provider result" in caplog.text
|
||||
|
||||
|
||||
def test_stream_does_not_swallow_interrupt_after_provider_success(
|
||||
relay_turn, monkeypatch
|
||||
):
|
||||
relay, turn = relay_turn
|
||||
|
||||
async def interrupt_after_stream(
|
||||
_name,
|
||||
request,
|
||||
callback,
|
||||
observe_chunk,
|
||||
finalizer,
|
||||
**_kwargs,
|
||||
):
|
||||
async def generate():
|
||||
upstream = callback(request)
|
||||
async for chunk in upstream:
|
||||
observe_chunk(chunk)
|
||||
yield chunk
|
||||
finalizer()
|
||||
raise KeyboardInterrupt
|
||||
|
||||
return generate()
|
||||
|
||||
monkeypatch.setattr(relay.llm, "stream_execute", interrupt_after_stream)
|
||||
stream = relay_llm.stream(
|
||||
{"model": "test-model", "messages": []},
|
||||
lambda _request: iter([{"delta": "complete"}]),
|
||||
session_id="session-1",
|
||||
name="test-provider",
|
||||
model_name="test-model",
|
||||
finalizer=lambda: {"content": "complete"},
|
||||
metadata={
|
||||
"api_mode": "custom",
|
||||
"api_request_id": "request-stream-post-interrupt",
|
||||
},
|
||||
)
|
||||
|
||||
assert next(stream) == {"delta": "complete"}
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
next(stream)
|
||||
assert turn.logical_llm_calls == {}
|
||||
|
||||
|
||||
def test_stream_does_not_swallow_hermes_finalizer_failure(relay_turn, monkeypatch):
|
||||
relay, _turn = relay_turn
|
||||
finalizer_error = RuntimeError("Hermes finalizer failed")
|
||||
|
|
|
|||
|
|
@ -240,3 +240,23 @@ def test_tool_adapter_returns_dispatch_result_after_relay_post_processing_failur
|
|||
assert result is raw_result
|
||||
assert observed_args == {"command": "pwd"}
|
||||
assert "returning the Hermes tool result" in caplog.text
|
||||
|
||||
|
||||
def test_tool_adapter_does_not_swallow_interrupt_after_dispatch_success(
|
||||
relay_turn, monkeypatch
|
||||
):
|
||||
relay = relay_turn
|
||||
|
||||
async def interrupt_after_callback(_name, args, callback, **_kwargs):
|
||||
callback(args)
|
||||
raise KeyboardInterrupt
|
||||
|
||||
monkeypatch.setattr(relay.tools, "execute", interrupt_after_callback)
|
||||
|
||||
with pytest.raises(KeyboardInterrupt):
|
||||
relay_tools.execute(
|
||||
"terminal",
|
||||
{"command": "pwd"},
|
||||
lambda _args: '{"ok":true}',
|
||||
session_id="session-1",
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue