fix(apodex): preserve reasoning block lifecycle

This commit is contained in:
zhanghanduo 2026-08-16 20:47:02 +08:00
parent ec3e293ac2
commit 428380fb1e
4 changed files with 177 additions and 58 deletions

View file

@ -3,7 +3,7 @@
import json
import traceback
from collections import deque
from collections.abc import AsyncIterator
from collections.abc import AsyncIterator, Mapping
from typing import Any, Final
from litellm import verbose_logger
@ -36,6 +36,8 @@ class AnthropicResponsesStreamWrapper:
self.model = model
self._message_id: str = f"msg_{uuid.uuid4()}"
self._current_block_index: int = -1
self._open_block_index: int | None = None
self._open_block_type: str | None = None
# Map item_id -> content_block_index so we can stop the right block later
self._item_id_to_block_index: dict[str, int] = {}
# Track open function_call items by item_id so we can emit tool_use start
@ -68,6 +70,48 @@ class AnthropicResponsesStreamWrapper:
self._current_block_index += 1
return self._current_block_index
def _close_open_block(self) -> None:
if self._open_block_index is None:
return
self._chunk_queue.append(
{
"type": "content_block_stop",
"index": self._open_block_index,
}
)
self._open_block_index = None
self._open_block_type = None
def _start_block(self, block_idx: int, block_type: str, content_block: Mapping[str, object]) -> None:
self._close_open_block()
self._chunk_queue.append(
{
"type": "content_block_start",
"index": block_idx,
"content_block": dict(content_block),
}
)
self._open_block_index = block_idx
self._open_block_type = block_type
def _get_or_start_block(
self,
item_id: str | None,
block_type: str,
content_block: Mapping[str, object],
) -> int:
mapped_index: Final = self._item_id_to_block_index.get(item_id) if item_id else None
if mapped_index is not None:
return mapped_index
if item_id is None and self._open_block_index is not None and self._open_block_type == block_type:
return self._open_block_index
block_idx: Final = self._next_block_index()
if item_id:
self._item_id_to_block_index[item_id] = block_idx
self._start_block(block_idx, block_type, content_block)
return block_idx
def _process_event(self, event: Any) -> None:
"""Convert one Responses API event into zero or more Anthropic chunks queued for emission."""
event_type = getattr(event, "type", None)
@ -96,13 +140,7 @@ class AnthropicResponsesStreamWrapper:
block_idx = self._next_block_index()
if item_id:
self._item_id_to_block_index[item_id] = block_idx
self._chunk_queue.append(
{
"type": "content_block_start",
"index": block_idx,
"content_block": {"type": "text", "text": ""},
}
)
self._start_block(block_idx, "text", {"type": "text", "text": ""})
elif item_type == "function_call":
call_id: Final = (
getattr(item, "call_id", None) or (item.get("call_id") if isinstance(item, dict) else None) or ""
@ -112,54 +150,36 @@ class AnthropicResponsesStreamWrapper:
if item_id:
self._item_id_to_block_index[item_id] = block_idx
self._pending_tool_ids[item_id] = call_id
self._chunk_queue.append(
self._start_block(
block_idx,
"tool_use",
{
"type": "content_block_start",
"index": block_idx,
"content_block": {
"type": "tool_use",
"id": call_id,
"name": name,
"input": {},
},
}
"type": "tool_use",
"id": call_id,
"name": name,
"input": {},
},
)
elif item_type == "reasoning":
block_idx = self._next_block_index()
if item_id:
self._item_id_to_block_index[item_id] = block_idx
self._chunk_queue.append(
{
"type": "content_block_start",
"index": block_idx,
"content_block": {"type": "thinking", "thinking": ""},
}
)
self._start_block(block_idx, "thinking", {"type": "thinking", "thinking": ""})
return
# ---- text delta ----
if event_type == "response.output_text.delta":
item_id = getattr(event, "item_id", None) or (event.get("item_id") if isinstance(event, dict) else None)
delta = getattr(event, "delta", "") or (event.get("delta", "") if isinstance(event, dict) else "")
block_idx = self._item_id_to_block_index.get(item_id, -1) if item_id else self._current_block_index
if block_idx < 0:
# Some providers (e.g. LMStudio) skip response.output_item.added,
# so no text block is open yet; synthesize content_block_start
# instead of emitting a delta with index -1
block_idx = self._next_block_index()
if item_id:
self._item_id_to_block_index[item_id] = block_idx
self._chunk_queue.append(
{
"type": "content_block_start",
"index": block_idx,
"content_block": {"type": "text", "text": ""},
}
)
text_block_idx: Final = self._get_or_start_block(
item_id=item_id,
block_type="text",
content_block={"type": "text", "text": ""},
)
self._chunk_queue.append(
{
"type": "content_block_delta",
"index": block_idx,
"index": text_block_idx,
"delta": {"type": "text_delta", "text": delta},
}
)
@ -169,15 +189,15 @@ class AnthropicResponsesStreamWrapper:
if event_type == "response.reasoning_summary_text.delta":
item_id = getattr(event, "item_id", None) or (event.get("item_id") if isinstance(event, dict) else None)
delta = getattr(event, "delta", "") or (event.get("delta", "") if isinstance(event, dict) else "")
block_idx = (
self._item_id_to_block_index.get(item_id, self._current_block_index)
if item_id
else self._current_block_index
thinking_block_idx: Final = self._get_or_start_block(
item_id=item_id,
block_type="thinking",
content_block={"type": "thinking", "thinking": ""},
)
self._chunk_queue.append(
{
"type": "content_block_delta",
"index": block_idx,
"index": thinking_block_idx,
"delta": {"type": "thinking_delta", "thinking": delta},
}
)
@ -212,12 +232,8 @@ class AnthropicResponsesStreamWrapper:
if item_id
else self._current_block_index
)
self._chunk_queue.append(
{
"type": "content_block_stop",
"index": block_idx,
}
)
if block_idx == self._open_block_index:
self._close_open_block()
return
# ---- response completed -> message_delta + message_stop ----
@ -226,6 +242,7 @@ class AnthropicResponsesStreamWrapper:
"response.failed",
"response.incomplete",
):
self._close_open_block()
response_obj: Final = getattr(event, "response", None) or (
event.get("response") if isinstance(event, dict) else None
)

View file

@ -165,8 +165,11 @@ class ApodexResponsesConfig(OpenAIResponsesAPIConfig):
if not isinstance(data, Mapping):
return None
channel: Final = _SWARM_CHANNELS.get(data.get("channel"))
channel_name: Final = data.get("channel")
delta: Final = data.get("delta")
if not isinstance(channel_name, str):
return None
channel: Final = _SWARM_CHANNELS.get(channel_name)
if channel is None or not isinstance(delta, str):
return None

View file

@ -113,9 +113,10 @@ class TestProcessEventTextDeltaWithoutOutputItemAdded:
{"type": "response.output_text.delta", "item_id": "m1", "delta": "Hi"},
]
)
assert chunks[1]["type"] == "content_block_start"
assert chunks[1]["content_block"] == {"type": "text", "text": ""}
assert [c["index"] for c in chunks[1:]] == [1, 1]
assert chunks[1] == {"type": "content_block_stop", "index": 0}
assert chunks[2]["type"] == "content_block_start"
assert chunks[2]["content_block"] == {"type": "text", "text": ""}
assert [c["index"] for c in chunks[1:]] == [0, 1, 1]
def test_process_event_registered_item_id_does_not_synthesize_start(self):
chunks = _process_all(
@ -133,6 +134,54 @@ class TestProcessEventTextDeltaWithoutOutputItemAdded:
]
class TestProcessEventReasoningDeltaWithoutOutputItemAdded:
def test_reasoning_opens_thinking_block_before_delta(self):
chunks = _process_all(
[
{"type": "response.reasoning_summary_text.delta", "item_id": "rs_1", "delta": "Think "},
{"type": "response.reasoning_summary_text.delta", "item_id": "rs_1", "delta": "carefully"},
]
)
assert [chunk["type"] for chunk in chunks] == [
"content_block_start",
"content_block_delta",
"content_block_delta",
]
assert chunks[0] == {
"type": "content_block_start",
"index": 0,
"content_block": {"type": "thinking", "thinking": ""},
}
assert chunks[1] == {
"type": "content_block_delta",
"index": 0,
"delta": {"type": "thinking_delta", "thinking": "Think "},
}
def test_reasoning_and_text_get_separate_closed_blocks(self):
response = SimpleNamespace(status="completed", output=[], usage=None)
chunks = _process_all(
[
{"type": "response.reasoning_summary_text.delta", "item_id": "rs_1", "delta": "Think"},
{"type": "response.output_text.delta", "item_id": "msg_1", "delta": "Answer"},
{"type": "response.completed", "response": response},
]
)
assert [(chunk["type"], chunk.get("index")) for chunk in chunks] == [
("content_block_start", 0),
("content_block_delta", 0),
("content_block_stop", 0),
("content_block_start", 1),
("content_block_delta", 1),
("content_block_stop", 1),
("message_delta", None),
("message_stop", None),
]
assert chunks[3]["content_block"] == {"type": "text", "text": ""}
class TestResponseCompletedUsage:
"""The Anthropic ``message_delta`` usage must report cache reads/writes and
exclude them from ``input_tokens``, so spend is not billed at the uncached

View file

@ -7,11 +7,15 @@ the model rather than applied provider-wide.
"""
import gzip
from types import SimpleNamespace
import httpx
import pytest
import litellm
from litellm.llms.anthropic.experimental_pass_through.responses_adapters.streaming_iterator import (
AnthropicResponsesStreamWrapper,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
@ -290,10 +294,56 @@ class TestDeepResearchKeepsState:
assert event.summary_index == 0
assert not hasattr(event, "content_index")
def test_reasoning_and_answer_form_valid_anthropic_blocks(self):
config = _responses_config("apodex-1-1-deep-research")
reasoning_event = config.transform_streaming_response(
model="apodex-1-1-deep-research",
parsed_chunk={
"type": "response.swarm.llm_delta",
"response_id": "w_c4b77c96",
"sequence_number": 7,
"swarm": {"data": {"channel": "reasoning", "delta": "The user wants"}},
},
logging_obj=None,
)
answer_event = config.transform_streaming_response(
model="apodex-1-1-deep-research",
parsed_chunk={
"type": "response.swarm.llm_delta",
"response_id": "w_c4b77c96",
"sequence_number": 8,
"swarm": {"data": {"channel": "output_text", "delta": "Hello there, friend!"}},
},
logging_obj=None,
)
wrapper = AnthropicResponsesStreamWrapper(responses_stream=None, model="apodex-1-1-deep-research")
for event in (
reasoning_event,
answer_event,
{"type": "response.completed", "response": SimpleNamespace(status="completed", output=[], usage=None)},
):
wrapper._process_event(event)
chunks = list(wrapper._chunk_queue)
assert [(chunk["type"], chunk.get("index")) for chunk in chunks] == [
("content_block_start", 0),
("content_block_delta", 0),
("content_block_stop", 0),
("content_block_start", 1),
("content_block_delta", 1),
("content_block_stop", 1),
("message_delta", None),
("message_stop", None),
]
assert chunks[0]["content_block"]["type"] == "thinking"
assert chunks[1]["delta"] == {"type": "thinking_delta", "thinking": "The user wants"}
assert chunks[3]["content_block"]["type"] == "text"
assert chunks[4]["delta"] == {"type": "text_delta", "text": "Hello there, friend!"}
@pytest.mark.parametrize(
"channel",
(None, "tool_output"),
ids=("no-channel", "unknown-channel"),
(None, "tool_output", []),
ids=("no-channel", "unknown-channel", "non-string-channel"),
)
def test_intermediate_agent_deltas_are_not_claimed(self, channel):
"""The worker agent streams a draft answer on a channel-less delta.