diff --git a/agent/relay_llm.py b/agent/relay_llm.py index ba9cb6ad248..e139c23baca 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -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 diff --git a/tests/agent/test_relay_llm.py b/tests/agent/test_relay_llm.py index 411a1220f4f..d286f8af1a3 100644 --- a/tests/agent/test_relay_llm.py +++ b/tests/agent/test_relay_llm.py @@ -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 ):