mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-15 16:49:24 +03:00
fix(anthropic): treat stream idle timeout as inactivity only, not total time
This commit is contained in:
@@ -782,67 +782,68 @@ class AnthropicProvider(LLMProvider):
|
|||||||
idle_timeout_s = resolve_stream_idle_timeout_s()
|
idle_timeout_s = resolve_stream_idle_timeout_s()
|
||||||
try:
|
try:
|
||||||
async with self._client.messages.stream(**kwargs) as stream:
|
async with self._client.messages.stream(**kwargs) as stream:
|
||||||
if on_content_delta or on_thinking_delta or on_tool_call_delta:
|
# Idle timeout must track *any* SSE chunk (thinking_delta,
|
||||||
# Idle timeout must track *any* SSE chunk (thinking_delta,
|
# tool JSON deltas, etc.), not only text_stream tokens.
|
||||||
# tool JSON deltas, etc.), not only text_stream tokens.
|
# Otherwise extended thinking can stall text_stream for minutes
|
||||||
# Otherwise extended thinking can stall text_stream for minutes
|
# while the connection is healthy (e.g. MiniMax Anthropic).
|
||||||
# while the connection is healthy (e.g. MiniMax Anthropic).
|
# Drain the whole stream with per-chunk idle waits so the
|
||||||
tool_blocks: dict[int, dict[str, str]] = {}
|
# timeout measures inactivity, not total generation time: a
|
||||||
while True:
|
# long but continuously-active stream must never be killed.
|
||||||
try:
|
# The SDK accumulates the final message snapshot during
|
||||||
chunk = await asyncio.wait_for(
|
# iteration, so get_final_message() below returns instantly.
|
||||||
stream.__anext__(),
|
tool_blocks: dict[int, dict[str, str]] = {}
|
||||||
timeout=idle_timeout_s,
|
while True:
|
||||||
)
|
try:
|
||||||
except StopAsyncIteration:
|
chunk = await asyncio.wait_for(
|
||||||
break
|
stream.__anext__(),
|
||||||
if chunk.type == "content_block_start":
|
timeout=idle_timeout_s,
|
||||||
block = getattr(chunk, "content_block", None)
|
)
|
||||||
if getattr(block, "type", None) == "tool_use":
|
except StopAsyncIteration:
|
||||||
index = int(getattr(chunk, "index", 0) or 0)
|
break
|
||||||
state = {
|
if chunk.type == "content_block_start":
|
||||||
"call_id": str(getattr(block, "id", "") or ""),
|
block = getattr(chunk, "content_block", None)
|
||||||
"name": str(getattr(block, "name", "") or ""),
|
if getattr(block, "type", None) == "tool_use":
|
||||||
}
|
index = int(getattr(chunk, "index", 0) or 0)
|
||||||
tool_blocks[index] = state
|
state = {
|
||||||
if on_tool_call_delta:
|
"call_id": str(getattr(block, "id", "") or ""),
|
||||||
await on_tool_call_delta({
|
"name": str(getattr(block, "name", "") or ""),
|
||||||
"index": index,
|
}
|
||||||
**state,
|
tool_blocks[index] = state
|
||||||
"arguments_delta": "",
|
if on_tool_call_delta:
|
||||||
})
|
|
||||||
elif (
|
|
||||||
chunk.type == "content_block_delta"
|
|
||||||
and getattr(chunk.delta, "type", None) == "thinking_delta"
|
|
||||||
):
|
|
||||||
piece = getattr(chunk.delta, "thinking", None) or ""
|
|
||||||
if piece and on_thinking_delta:
|
|
||||||
await on_thinking_delta(piece)
|
|
||||||
elif (
|
|
||||||
chunk.type == "content_block_delta"
|
|
||||||
and getattr(chunk.delta, "type", None) == "text_delta"
|
|
||||||
):
|
|
||||||
text = getattr(chunk.delta, "text", None) or ""
|
|
||||||
if text and on_content_delta:
|
|
||||||
await on_content_delta(text)
|
|
||||||
elif (
|
|
||||||
chunk.type == "content_block_delta"
|
|
||||||
and getattr(chunk.delta, "type", None) == "input_json_delta"
|
|
||||||
):
|
|
||||||
partial = getattr(chunk.delta, "partial_json", None) or ""
|
|
||||||
if partial and on_tool_call_delta:
|
|
||||||
index = int(getattr(chunk, "index", 0) or 0)
|
|
||||||
state = tool_blocks.get(index, {})
|
|
||||||
await on_tool_call_delta({
|
await on_tool_call_delta({
|
||||||
"index": index,
|
"index": index,
|
||||||
"call_id": state.get("call_id", ""),
|
**state,
|
||||||
"name": state.get("name", ""),
|
"arguments_delta": "",
|
||||||
"arguments_delta": partial,
|
|
||||||
})
|
})
|
||||||
response = await asyncio.wait_for(
|
elif (
|
||||||
stream.get_final_message(),
|
chunk.type == "content_block_delta"
|
||||||
timeout=idle_timeout_s,
|
and getattr(chunk.delta, "type", None) == "thinking_delta"
|
||||||
)
|
):
|
||||||
|
piece = getattr(chunk.delta, "thinking", None) or ""
|
||||||
|
if piece and on_thinking_delta:
|
||||||
|
await on_thinking_delta(piece)
|
||||||
|
elif (
|
||||||
|
chunk.type == "content_block_delta"
|
||||||
|
and getattr(chunk.delta, "type", None) == "text_delta"
|
||||||
|
):
|
||||||
|
text = getattr(chunk.delta, "text", None) or ""
|
||||||
|
if text and on_content_delta:
|
||||||
|
await on_content_delta(text)
|
||||||
|
elif (
|
||||||
|
chunk.type == "content_block_delta"
|
||||||
|
and getattr(chunk.delta, "type", None) == "input_json_delta"
|
||||||
|
):
|
||||||
|
partial = getattr(chunk.delta, "partial_json", None) or ""
|
||||||
|
if partial and on_tool_call_delta:
|
||||||
|
index = int(getattr(chunk, "index", 0) or 0)
|
||||||
|
state = tool_blocks.get(index, {})
|
||||||
|
await on_tool_call_delta({
|
||||||
|
"index": index,
|
||||||
|
"call_id": state.get("call_id", ""),
|
||||||
|
"name": state.get("name", ""),
|
||||||
|
"arguments_delta": partial,
|
||||||
|
})
|
||||||
|
response = await stream.get_final_message()
|
||||||
return self._parse_response(response)
|
return self._parse_response(response)
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
return LLMResponse(
|
return LLMResponse(
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
@@ -48,6 +49,107 @@ class _FakeAsyncStream:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class _ConsumingFakeAsyncStream:
|
||||||
|
"""Mimics the real AsyncMessageStream: ``__anext__`` yields chunks after a
|
||||||
|
per-chunk network delay, and ``get_final_message()`` consumes the remaining
|
||||||
|
chunks (like the SDK's ``until_done()``) before returning."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
chunks: list[SimpleNamespace],
|
||||||
|
per_chunk_delay: float,
|
||||||
|
) -> None:
|
||||||
|
self._chunks = chunks
|
||||||
|
self._idx = 0
|
||||||
|
self._delay = per_chunk_delay
|
||||||
|
|
||||||
|
async def __anext__(self) -> SimpleNamespace:
|
||||||
|
if self._idx >= len(self._chunks):
|
||||||
|
raise StopAsyncIteration
|
||||||
|
c = self._chunks[self._idx]
|
||||||
|
self._idx += 1
|
||||||
|
await asyncio.sleep(self._delay)
|
||||||
|
return c
|
||||||
|
|
||||||
|
def __aiter__(self) -> _ConsumingFakeAsyncStream:
|
||||||
|
return self
|
||||||
|
|
||||||
|
async def get_final_message(self) -> SimpleNamespace:
|
||||||
|
async for _ in self:
|
||||||
|
pass
|
||||||
|
return _final_message_stub("ok")
|
||||||
|
|
||||||
|
async def __aenter__(self) -> _ConsumingFakeAsyncStream:
|
||||||
|
return self
|
||||||
|
|
||||||
|
async def __aexit__(self, *_exc: object) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_chat_stream_without_callback_survives_long_active_stream(monkeypatch) -> None:
|
||||||
|
"""Regression: the idle timeout must not double as a total timeout.
|
||||||
|
|
||||||
|
A stream that keeps producing chunks (5 x 0.06s = 0.30s) well past the
|
||||||
|
idle timeout (0.15s) must complete. Currently the no-callback path wraps
|
||||||
|
``stream.get_final_message()`` in ``wait_for(timeout=idle_timeout_s)``,
|
||||||
|
which measures total wall-clock time and kills the stream even though it
|
||||||
|
is continuously active.
|
||||||
|
"""
|
||||||
|
monkeypatch.setenv("NANOBOT_STREAM_IDLE_TIMEOUT_S", "0.15")
|
||||||
|
provider = AnthropicProvider(api_key="sk-test")
|
||||||
|
provider._client = MagicMock()
|
||||||
|
|
||||||
|
chunks = [
|
||||||
|
SimpleNamespace(
|
||||||
|
type="content_block_delta",
|
||||||
|
delta=SimpleNamespace(type="text_delta", text="a"),
|
||||||
|
)
|
||||||
|
for _ in range(5)
|
||||||
|
]
|
||||||
|
fake = _ConsumingFakeAsyncStream(chunks, per_chunk_delay=0.06)
|
||||||
|
stream_cm = MagicMock()
|
||||||
|
stream_cm.__aenter__ = AsyncMock(return_value=fake)
|
||||||
|
stream_cm.__aexit__ = AsyncMock(return_value=None)
|
||||||
|
provider._client.messages.stream = MagicMock(return_value=stream_cm)
|
||||||
|
|
||||||
|
res = await provider.chat_stream(
|
||||||
|
messages=[{"role": "user", "content": "hello"}],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert res.finish_reason != "error", (
|
||||||
|
f"active stream was killed by total-timeout misuse: {res.content}"
|
||||||
|
)
|
||||||
|
assert res.content == "ok"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_chat_stream_without_callback_still_enforces_idle_timeout(monkeypatch) -> None:
|
||||||
|
"""A genuinely stalled stream must still be cut off by the idle timeout."""
|
||||||
|
monkeypatch.setenv("NANOBOT_STREAM_IDLE_TIMEOUT_S", "0.05")
|
||||||
|
provider = AnthropicProvider(api_key="sk-test")
|
||||||
|
provider._client = MagicMock()
|
||||||
|
|
||||||
|
class _StalledStream(_FakeAsyncStream):
|
||||||
|
async def __anext__(self) -> SimpleNamespace:
|
||||||
|
await asyncio.sleep(3600)
|
||||||
|
raise StopAsyncIteration
|
||||||
|
|
||||||
|
fake = _StalledStream([])
|
||||||
|
stream_cm = MagicMock()
|
||||||
|
stream_cm.__aenter__ = AsyncMock(return_value=fake)
|
||||||
|
stream_cm.__aexit__ = AsyncMock(return_value=None)
|
||||||
|
provider._client.messages.stream = MagicMock(return_value=stream_cm)
|
||||||
|
|
||||||
|
res = await provider.chat_stream(
|
||||||
|
messages=[{"role": "user", "content": "hello"}],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert res.finish_reason == "error"
|
||||||
|
assert res.error_kind == "timeout"
|
||||||
|
assert "stalled" in (res.content or "")
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_chat_stream_calls_on_content_delta_only_for_text_delta() -> None:
|
async def test_chat_stream_calls_on_content_delta_only_for_text_delta() -> None:
|
||||||
"""Thinking deltas must be consumed without invoking on_content_delta."""
|
"""Thinking deltas must be consumed without invoking on_content_delta."""
|
||||||
|
|||||||
Reference in New Issue
Block a user