diff --git a/agent/relay_llm.py b/agent/relay_llm.py index 6d018d462fa..1dc80b0ba11 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -387,6 +387,9 @@ class ManagedLlmStream(Iterator[Any]): ) ) except BaseException: + if not self._defer_logical_completion: + _complete_logical(self._logical, outcome="failed") + self._logical = None loop.close() self._loop = None raise @@ -419,10 +422,7 @@ class ManagedLlmStream(Iterator[Any]): raise StopIteration from None except BaseException as exc: callback_error = self._callback_error - self.close() - if not self._defer_logical_completion: - _complete_logical(self._logical, outcome="failed") - self._logical = None + self._close(logical_outcome="failed") if ( callback_error is not None and relay_runtime._is_relay_wrapped_callback_error(exc, callback_error) @@ -441,6 +441,10 @@ class ManagedLlmStream(Iterator[Any]): return self._chunk_adapter(chunk) def close(self) -> None: + """Close an explicitly abandoned stream and cancel its logical call.""" + self._close(logical_outcome="cancelled") + + def _close(self, *, logical_outcome: str) -> None: if self._closed: return self._closed = True @@ -450,6 +454,9 @@ class ManagedLlmStream(Iterator[Any]): close = getattr(self._stream, "close", None) if callable(close): close() + if not self._defer_logical_completion: + _complete_logical(self._logical, outcome=logical_outcome) + self._logical = None return close = getattr(self._stream, "aclose", None) if callable(close): @@ -461,6 +468,9 @@ class ManagedLlmStream(Iterator[Any]): loop.run_until_complete(close_stream()) except Exception: pass + if not self._defer_logical_completion: + _complete_logical(self._logical, outcome=logical_outcome) + self._logical = None loop.close() def __del__(self) -> None: diff --git a/tests/agent/test_relay_llm.py b/tests/agent/test_relay_llm.py index 40f9cb8372e..b3b62cb7ab8 100644 --- a/tests/agent/test_relay_llm.py +++ b/tests/agent/test_relay_llm.py @@ -170,6 +170,41 @@ def test_deferred_stream_preserves_provider_error_and_logical_scope_for_retry( assert "request-2" in turn.logical_llm_calls +def test_non_deferred_partial_stream_close_cancels_logical_call( + relay_turn, + monkeypatch, +): + relay, turn = relay_turn + original_pop = relay.scope.pop + terminal_outputs = [] + + def record_pop(handle, *args, **kwargs): + terminal_outputs.append(kwargs.get("output")) + return original_pop(handle, *args, **kwargs) + + monkeypatch.setattr(relay.scope, "pop", record_pop) + stream = relay_llm.stream( + {"model": "test-model", "messages": []}, + lambda _request: iter([{"delta": "partial"}, {"delta": "unused"}]), + session_id="session-1", + name="test-provider", + model_name="test-model", + finalizer=lambda: {"content": "partial"}, + metadata={ + "api_mode": "custom", + "api_request_id": "request-partial-close", + }, + ) + + assert next(stream) == {"delta": "partial"} + assert "request-partial-close" in turn.logical_llm_calls + + stream.close() + + assert "request-partial-close" not in turn.logical_llm_calls + assert {"outcome": "cancelled"} in terminal_outputs + + def test_non_stream_preserves_raw_provider_response_identity(relay_turn): _relay, _turn = relay_turn raw_response = SimpleNamespace(model="test-model", content="raw")