fix(relay): preserve buffered provider stream chunks

Signed-off-by: Alex Fournier <afournier@nvidia.com>
This commit is contained in:
Alex Fournier 2026-07-27 08:16:16 -07:00
parent bc597571fe
commit ec6c881453
2 changed files with 116 additions and 8 deletions

View file

@ -467,12 +467,7 @@ class ManagedLlmStream(Iterator[Any]):
"preserving the provider result",
exc_info=True,
)
if not self._defer_logical_completion:
_complete_logical(self._logical, outcome="success")
self._logical = None
loop.close()
self._loop = None
self._stream = iter(())
self._preserve_pending_provider_chunks()
return
if not self._defer_logical_completion:
_complete_logical(
@ -532,8 +527,8 @@ class ManagedLlmStream(Iterator[Any]):
"preserving the provider result",
exc_info=True,
)
self._close(logical_outcome="success")
raise StopIteration from None
self._preserve_pending_provider_chunks()
return next(self)
self._close(
logical_outcome="cancelled" if _is_cancellation(exc) else "failed"
)
@ -553,6 +548,35 @@ class ManagedLlmStream(Iterator[Any]):
"""Close an explicitly abandoned stream and cancel its logical call."""
self._close(logical_outcome="cancelled")
def _preserve_pending_provider_chunks(self) -> None:
"""Switch a failed Relay stream to its undelivered provider chunks."""
pending = [raw for _encoded, raw in self._raw_chunks]
self._raw_chunks.clear()
loop = self._loop
relay_stream = self._stream
self._loop = None
self._stream = iter(pending)
self._raw_stream_resource = None
self._accept_chunk = None
if loop is not None:
close = getattr(relay_stream, "aclose", None)
if callable(close):
async def close_stream() -> None:
await close()
try:
loop.run_until_complete(close_stream())
except Exception:
logger.debug(
"Relay stream cleanup failed during provider fallback",
exc_info=True,
)
loop.close()
if not self._defer_logical_completion:
_complete_logical(self._logical, outcome="success")
self._logical = None
def _close(self, *, logical_outcome: str) -> None:
if self._closed:
return

View file

@ -690,6 +690,90 @@ def test_stream_finishes_after_relay_post_processing_failure(
assert "preserving the provider result" in caplog.text
def test_stream_flushes_buffered_provider_chunks_after_relay_failure(
relay_turn, monkeypatch
):
relay, turn = relay_turn
raw_chunks = [{"delta": "first"}, {"delta": "second"}]
async def fail_with_buffered_chunk(
_name,
request,
callback,
observe_chunk,
finalizer,
**_kwargs,
):
async def generate():
upstream = callback(request)
first = await anext(upstream)
observe_chunk(first)
yield first
second = await anext(upstream)
observe_chunk(second)
with pytest.raises(StopAsyncIteration):
await anext(upstream)
finalizer()
raise RuntimeError("simulated buffered Relay failure")
return generate()
monkeypatch.setattr(relay.llm, "stream_execute", fail_with_buffered_chunk)
stream = relay_llm.stream(
{"model": "test-model", "messages": []},
lambda _request: iter(raw_chunks),
session_id="session-1",
name="test-provider",
model_name="test-model",
finalizer=lambda: {"content": "complete"},
metadata={
"api_mode": "custom",
"api_request_id": "request-buffered-failure",
},
)
assert list(stream) == raw_chunks
assert turn.logical_llm_calls == {}
def test_stream_constructor_flushes_provider_chunks_after_relay_failure(
relay_turn, monkeypatch
):
relay, turn = relay_turn
raw_chunks = [{"delta": "first"}, {"delta": "second"}]
async def fail_during_stream_setup(
_name,
request,
callback,
observe_chunk,
finalizer,
**_kwargs,
):
upstream = callback(request)
async for chunk in upstream:
observe_chunk(chunk)
finalizer()
raise RuntimeError("simulated Relay setup failure")
monkeypatch.setattr(relay.llm, "stream_execute", fail_during_stream_setup)
stream = relay_llm.stream(
{"model": "test-model", "messages": []},
lambda _request: iter(raw_chunks),
session_id="session-1",
name="test-provider",
model_name="test-model",
finalizer=lambda: {"content": "complete"},
metadata={
"api_mode": "custom",
"api_request_id": "request-setup-failure",
},
)
assert list(stream) == raw_chunks
assert turn.logical_llm_calls == {}
def test_stream_does_not_swallow_interrupt_after_provider_success(
relay_turn, monkeypatch
):