fix: continue recovered streams in a new segment
maintainer edit: streamed timeout recovery was returning the retried response internally while the channel still treated the final outbound as already streamed. End the current stream segment before retry/fallback recovery so subsequent deltas are delivered in a new segment.
This commit is contained in:
@@ -492,6 +492,61 @@ class TestToolEventProgress:
|
||||
assert turn_end_msgs[0].content == ""
|
||||
provider.chat_with_retry.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_timeout_recovery_continues_in_new_segment(
|
||||
self,
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""Recovered streaming output should use a new stream segment."""
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.supports_progress_deltas = True
|
||||
provider.get_default_model.return_value = "openai-codex/gpt-5.5"
|
||||
|
||||
async def chat_stream_with_retry(*, on_content_delta, on_stream_recover, **kwargs):
|
||||
await on_content_delta("partial")
|
||||
await on_stream_recover()
|
||||
await on_content_delta("full retry response")
|
||||
return LLMResponse(content="full retry response", tool_calls=[])
|
||||
|
||||
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||
provider.chat_with_retry = AsyncMock()
|
||||
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="openai-codex/gpt-5.5")
|
||||
_attach_webui_runtime_events(loop, bus)
|
||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
|
||||
|
||||
await loop._dispatch(InboundMessage(
|
||||
channel="websocket",
|
||||
sender_id="u1",
|
||||
chat_id="chat1",
|
||||
content="say hello",
|
||||
metadata={"_wants_stream": True},
|
||||
))
|
||||
|
||||
outbound = []
|
||||
while bus.outbound_size > 0:
|
||||
outbound.append(await bus.consume_outbound())
|
||||
|
||||
deltas = [m for m in outbound if m.metadata.get("_stream_delta")]
|
||||
stream_end = [m for m in outbound if m.metadata.get("_stream_end")]
|
||||
final = [
|
||||
m for m in outbound
|
||||
if not m.metadata.get("_stream_delta")
|
||||
and not m.metadata.get("_stream_end")
|
||||
and not m.metadata.get("_turn_end")
|
||||
and not m.metadata.get("_goal_status")
|
||||
]
|
||||
|
||||
assert [m.content for m in deltas] == ["partial", "full retry response"]
|
||||
assert [m.metadata.get("_resuming") for m in stream_end] == [True, False]
|
||||
assert deltas[0].metadata.get("_stream_id") == stream_end[0].metadata.get("_stream_id")
|
||||
assert deltas[1].metadata.get("_stream_id") == stream_end[1].metadata.get("_stream_id")
|
||||
assert deltas[0].metadata.get("_stream_id") != deltas[1].metadata.get("_stream_id")
|
||||
assert final[-1].content == "full retry response"
|
||||
assert final[-1].metadata.get("_streamed") is True
|
||||
provider.chat_with_retry.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streamed_progress_is_not_repeated_before_tool_execution(
|
||||
self,
|
||||
|
||||
@@ -323,18 +323,24 @@ class TestFallbackOnStreamStalledAfterContent:
|
||||
)
|
||||
|
||||
streamed: list[str] = []
|
||||
recoveries: list[str] = []
|
||||
|
||||
async def _delta(text: str) -> None:
|
||||
streamed.append(text)
|
||||
|
||||
async def _recover() -> None:
|
||||
recoveries.append("recover")
|
||||
|
||||
result = await fb.chat_stream(
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
on_content_delta=_delta,
|
||||
on_stream_recover=_recover,
|
||||
)
|
||||
assert result.finish_reason == "stop"
|
||||
assert result.content == "fallback ok"
|
||||
factory.assert_called_once_with(_fallback("fallback-a"))
|
||||
assert "stream stalled" in streamed
|
||||
assert streamed == ["stream stalled", "fallback ok"]
|
||||
assert recoveries == ["recover"]
|
||||
|
||||
|
||||
class TestFailoverOnTransientError:
|
||||
|
||||
@@ -199,6 +199,49 @@ async def test_chat_stream_with_retry_retries_timeout_after_emitting_content(mon
|
||||
assert provider.last_kwargs.get("on_content_delta") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_stream_with_retry_retries_timeout_in_new_stream_segment(
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
first = LLMResponse(
|
||||
content="Error calling LLM: stream stalled for more than 30 seconds",
|
||||
finish_reason="error",
|
||||
error_kind="timeout",
|
||||
)
|
||||
first._test_stream_delta = "partial" # type: ignore[attr-defined]
|
||||
second = LLMResponse(content="full retry response")
|
||||
second._test_stream_delta = "full retry response" # type: ignore[attr-defined]
|
||||
provider = ScriptedProvider([first, second])
|
||||
deltas: list[str] = []
|
||||
recoveries: list[str] = []
|
||||
delays: list[int] = []
|
||||
|
||||
async def _fake_sleep(delay: int) -> None:
|
||||
delays.append(delay)
|
||||
|
||||
async def _on_delta(delta: str) -> None:
|
||||
deltas.append(delta)
|
||||
|
||||
async def _on_stream_recover() -> None:
|
||||
recoveries.append("recover")
|
||||
|
||||
monkeypatch.setattr("nanobot.providers.base.asyncio.sleep", _fake_sleep)
|
||||
|
||||
response = await provider.chat_stream_with_retry(
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
on_content_delta=_on_delta,
|
||||
on_stream_recover=_on_stream_recover,
|
||||
)
|
||||
|
||||
assert response.content == "full retry response"
|
||||
assert response.finish_reason == "stop"
|
||||
assert provider.calls == 2
|
||||
assert deltas == ["partial", "full retry response"]
|
||||
assert recoveries == ["recover"]
|
||||
assert delays == [1]
|
||||
assert provider.last_kwargs.get("on_content_delta") is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_with_retry_uses_provider_generation_defaults() -> None:
|
||||
"""When callers omit generation params, provider.generation defaults are used."""
|
||||
|
||||
Reference in New Issue
Block a user