From 06ed41aa1f4fdd4e319e58e55897a8cf7c6c8a00 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Thu, 23 Jul 2026 14:21:37 -0700 Subject: [PATCH] fix(relay): preserve managed cancellation signals Signed-off-by: Alex Fournier --- agent/relay_llm.py | 39 ++++++++++-- agent/relay_tools.py | 6 +- tests/agent/test_relay_llm.py | 105 ++++++++++++++++++++++++++++++++ tests/agent/test_relay_tools.py | 20 ++++++ 4 files changed, 163 insertions(+), 7 deletions(-) diff --git a/agent/relay_llm.py b/agent/relay_llm.py index 7480af3da5f..b4f18e935e8 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -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() diff --git a/agent/relay_tools.py b/agent/relay_tools.py index 3972b97d93a..5023df1bcf9 100644 --- a/agent/relay_tools.py +++ b/agent/relay_tools.py @@ -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", diff --git a/tests/agent/test_relay_llm.py b/tests/agent/test_relay_llm.py index c65ba2292c8..524a91fb090 100644 --- a/tests/agent/test_relay_llm.py +++ b/tests/agent/test_relay_llm.py @@ -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") diff --git a/tests/agent/test_relay_tools.py b/tests/agent/test_relay_tools.py index e7d15581f61..8f605dfd4d9 100644 --- a/tests/agent/test_relay_tools.py +++ b/tests/agent/test_relay_tools.py @@ -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", + )