fix(runtime): close abandoned Relay streams

Signed-off-by: Alex Fournier <afournier@nvidia.com>
This commit is contained in:
Alex Fournier 2026-07-22 10:32:49 -07:00
parent 761fce7f79
commit 537425ebe4
2 changed files with 49 additions and 4 deletions

View file

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

View file

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