mirror of
https://github.com/NousResearch/hermes-agent.git
synced 2026-07-31 19:16:29 +00:00
fix(relay): preserve buffered provider stream chunks
Signed-off-by: Alex Fournier <afournier@nvidia.com>
This commit is contained in:
parent
bc597571fe
commit
ec6c881453
2 changed files with 116 additions and 8 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue