Merge pull request #41503 from BerriAI/litellm_fix_interrupted_anthropic_reasoning_usage

fix(streaming): estimate interrupted Anthropic stream usage from reasoning_content
This commit is contained in:
kerry-berri 2026-09-16 17:04:59 -07:00 committed by GitHub
commit f93b31679b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 174 additions and 49 deletions

View file

@ -148,6 +148,7 @@ class _ToolCallChunk(TypedDict):
class _UsageBearingChunk(TypedDict, total=False):
usage: Usage | None
_hidden_params: Mapping[str, str]
choices: ReadOnly[Sequence[StreamingChoices | Mapping[str, object]]]
class _UsageSummary(TypedDict):
@ -921,21 +922,22 @@ class ChunkProcessor:
prompt_tokens_details = attach_cache_creation_token_details(prompt_tokens_details, cache_creation_token_details)
completion_tokens = self._reset_anthropic_cursor_completion_tokens(
recovered_completion_tokens: Final = self._reset_anthropic_cursor_completion_tokens(
chunks=chunks,
completion_tokens=completion_tokens,
completion_usage_updates=completion_usage_updates,
)
cursor_was_reset: Final = recovered_completion_tokens != completion_tokens
return UsagePerChunk(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
completion_tokens=recovered_completion_tokens,
cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens,
server_tool_use=server_tool_use,
web_search_requests=web_search_requests,
google_maps_grounding_requests=google_maps_grounding_requests,
completion_tokens_details=completion_tokens_details,
completion_tokens_details=None if cursor_was_reset else completion_tokens_details,
prompt_tokens_details=prompt_tokens_details,
cost=cost,
inference_geo=self._last_provider_pricing_field(chunks, "inference_geo"),
@ -960,6 +962,30 @@ class ChunkProcessor:
]
return values[-1] if values else None
@staticmethod
def _finish_reason_of_choice(choice: object) -> str | None:
match choice:
case StreamingChoices(finish_reason=reason) | Choices(finish_reason=reason):
return reason
case {"finish_reason": str() as reason}:
return reason
case _:
return None
@staticmethod
def _chunk_choices(chunk: "_UsageBearingChunk | ModelResponse | ModelResponseStream") -> Sequence[object]:
if isinstance(chunk, dict):
return chunk.get("choices", ())
return getattr(chunk, "choices", ())
@staticmethod
def _saw_finish_reason(chunks: Sequence["_UsageBearingChunk | ModelResponse"]) -> bool:
return any(
ChunkProcessor._finish_reason_of_choice(choice) is not None
for chunk in chunks
for choice in ChunkProcessor._chunk_choices(chunk)
)
@staticmethod
def _reset_anthropic_cursor_completion_tokens(
chunks: Sequence["_UsageBearingChunk | ModelResponse"],
@ -970,18 +996,18 @@ class ChunkProcessor:
See the ``completion_usage_updates`` comment in
``_calculate_usage_per_chunk``. The accumulated value is NOT a stale
cursor when either it is > 1 (definitely not a placeholder) or we saw
>= 2 completion-bearing usage events (positive evidence ``message_delta``
arrived). Otherwise the only completion update we ever saw was the
Anthropic ``message_start`` cursor (=1) reset to 0 so
``calculate_usage()``'s ``or token_counter(text=...)`` fallback estimates
from the actually-received completion text instead of trusting the
placeholder. Gated on ``custom_llm_provider == "anthropic"`` so the
heuristic (which encodes Anthropic's specific message_start SSE shape)
does not silently affect other providers that may legitimately report
``completion_tokens=1`` from a single usage event.
cursor when we saw >= 2 completion-bearing usage events or any chunk
carried a ``finish_reason`` (positive evidence ``message_delta``
arrived). Otherwise the only completion update we ever saw was the
Anthropic ``message_start`` cursor, a small placeholder whose magnitude
varies per request (1 and 8 both observed live), so reset to 0 and let
``calculate_usage()``'s ``or token_counter(...)`` fallback estimate from
the actually-received text and reasoning instead. Gated on
``custom_llm_provider == "anthropic"`` so the heuristic (which encodes
Anthropic's specific message_start SSE shape) does not silently affect
other providers that legitimately report usage from a single event.
"""
saw_non_cursor_completion: Final = completion_tokens > 1 or completion_usage_updates >= 2
saw_non_cursor_completion: Final = completion_usage_updates >= 2 or ChunkProcessor._saw_finish_reason(chunks)
if saw_non_cursor_completion:
return completion_tokens
@ -995,7 +1021,7 @@ class ChunkProcessor:
if isinstance(hp, dict):
custom_llm_provider = hp.get("custom_llm_provider")
if custom_llm_provider == "anthropic" and completion_tokens == 1:
if custom_llm_provider == "anthropic":
return 0
return completion_tokens
@ -1039,10 +1065,13 @@ class ChunkProcessor:
returned_usage.prompt_tokens = 0
returned_usage.completion_tokens = (
completion_tokens
or token_counter(
model=model,
text=completion_output,
count_response_tokens=True, # count_response_tokens is a Flag to tell token counter this is a response, No need to add extra tokens we do for input messages
or (
token_counter(
model=model,
text=completion_output,
count_response_tokens=True, # count_response_tokens is a Flag to tell token counter this is a response, No need to add extra tokens we do for input messages
)
+ (reasoning_tokens or 0)
)
)
returned_usage.total_tokens = returned_usage.prompt_tokens + returned_usage.completion_tokens
@ -1066,15 +1095,16 @@ class ChunkProcessor:
returned_usage.completion_tokens_details = completion_tokens_details
if reasoning_tokens is not None:
capped_reasoning_tokens: Final = min(max(0, reasoning_tokens), returned_usage.completion_tokens)
if returned_usage.completion_tokens_details is None:
returned_usage.completion_tokens_details = CompletionTokensDetailsWrapper(
reasoning_tokens=reasoning_tokens
reasoning_tokens=capped_reasoning_tokens,
text_tokens=returned_usage.completion_tokens - capped_reasoning_tokens,
)
elif (
returned_usage.completion_tokens_details is not None
and returned_usage.completion_tokens_details.reasoning_tokens is None
):
capped_reasoning_tokens: Final = min(max(0, reasoning_tokens), returned_usage.completion_tokens)
returned_usage.completion_tokens_details.reasoning_tokens = capped_reasoning_tokens
if returned_usage.completion_tokens_details.text_tokens is None:
returned_usage.completion_tokens_details.text_tokens = (

View file

@ -19,12 +19,12 @@ to 0 when the only update we saw was the cursor, allowing the
text-based fallback to estimate from the real completion text.
"""
import pytest
import litellm
from litellm.litellm_core_utils.streaming_chunk_builder_utils import ChunkProcessor
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
Delta,
ModelResponseStream,
StreamingChoices,
@ -35,6 +35,7 @@ from litellm.types.utils import (
def _make_chunk(
*,
content: str = "",
reasoning_content: str | None = None,
usage: Usage = None,
finish_reason: str = None,
custom_llm_provider: str = "anthropic",
@ -48,7 +49,7 @@ def _make_chunk(
StreamingChoices(
finish_reason=finish_reason,
index=0,
delta=Delta(content=content, role="assistant"),
delta=Delta(content=content, role="assistant", reasoning_content=reasoning_content),
)
],
usage=usage,
@ -69,9 +70,7 @@ class TestAnthropicCursorBug:
token_counter fallback can estimate from completion text.
"""
# Anthropic message_start: input_tokens accurate, output_tokens=1 cursor
message_start = _make_chunk(
usage=Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025)
)
message_start = _make_chunk(usage=Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025))
# Several content_block_delta chunks (no usage attached)
text_chunks = [
_make_chunk(content="Hello"),
@ -97,9 +96,7 @@ class TestAnthropicCursorBug:
Normal complete stream: message_start cursor=1, then message_delta=3847.
Last-wins must give 3847 (the real value).
"""
message_start = _make_chunk(
usage=Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025)
)
message_start = _make_chunk(usage=Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025))
text_chunks = [_make_chunk(content=t) for t in ["Hello", " world", "!"]]
# message_delta with the real cumulative output_tokens
message_delta = _make_chunk(
@ -119,19 +116,14 @@ class TestAnthropicCursorBug:
End-to-end via calculate_usage(): cursor-only stream + real completion
text should produce a token-counter estimate, NOT 1.
"""
message_start = _make_chunk(
usage=Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025)
)
message_start = _make_chunk(usage=Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025))
# ~50 visible chars ≈ ~12 tokens (anthropic-style tokenizer ballpark)
text_chunks = [
_make_chunk(content="Based on your question, I think the answer is "),
_make_chunk(content="forty-two. Here is my reasoning: "),
]
chunks = [message_start, *text_chunks]
completion_output = (
"Based on your question, I think the answer is forty-two. "
"Here is my reasoning: "
)
completion_output = "Based on your question, I think the answer is forty-two. Here is my reasoning: "
processor = ChunkProcessor(chunks=chunks, messages=[])
usage = processor.calculate_usage(
@ -149,9 +141,7 @@ class TestAnthropicCursorBug:
def test_cache_fields_preserved_from_message_start(self):
"""cache_read / cache_creation come from message_start and must survive."""
message_start_usage = Usage(
prompt_tokens=1024, completion_tokens=1, total_tokens=1025
)
message_start_usage = Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025)
# Anthropic puts these in message_start
message_start_usage.cache_read_input_tokens = 512
message_start_usage.cache_creation_input_tokens = 128
@ -193,9 +183,7 @@ class TestAnthropicCursorBug:
on a 1-token string also gives ~1, so billing is still approximately
correct. This test pins that the result is sane (1 or 0).
"""
message_start = _make_chunk(
usage=Usage(prompt_tokens=20, completion_tokens=1, total_tokens=21)
)
message_start = _make_chunk(usage=Usage(prompt_tokens=20, completion_tokens=1, total_tokens=21))
text_chunk = _make_chunk(content="Yes.")
# Anthropic's message_delta also gives output_tokens=1 in this case
message_delta = _make_chunk(
@ -231,9 +219,7 @@ class TestAnthropicCursorBug:
must fire so token_counter estimates from completion text instead of
billing the placeholder.
"""
message_start_usage = Usage(
prompt_tokens=1024, completion_tokens=1, total_tokens=1025
)
message_start_usage = Usage(prompt_tokens=1024, completion_tokens=1, total_tokens=1025)
message_start_usage.cache_read_input_tokens = 4096
message_start = _make_chunk(usage=message_start_usage)
# Subsequent chunks with cache fields but no completion_tokens
@ -253,6 +239,114 @@ class TestAnthropicCursorBug:
"Reset to 0 forces token_counter fallback."
)
@pytest.mark.parametrize("placeholder", [1, 3, 8])
def test_interrupted_reasoning_only_stream_estimates_from_reasoning(self, placeholder: int):
message_start = _make_chunk(
usage=Usage(
prompt_tokens=100,
completion_tokens=placeholder,
total_tokens=100 + placeholder,
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=0, text_tokens=placeholder),
)
)
reasoning_text = "Let me work through the scheduling constraints step by step. " * 40
reasoning_chunks = [
_make_chunk(reasoning_content=reasoning_text[i : i + 50]) for i in range(0, len(reasoning_text), 50)
]
response = litellm.stream_chunk_builder(
chunks=[message_start, *reasoning_chunks],
messages=[{"role": "user", "content": "Plan the schedule."}],
)
assert response.choices[0].message.reasoning_content == reasoning_text
reasoning_tokens = response.usage.completion_tokens_details.reasoning_tokens
assert reasoning_tokens > placeholder
assert response.usage.completion_tokens == reasoning_tokens, (
f"Expected completion_tokens to be the reasoning estimate, got "
f"completion_tokens={response.usage.completion_tokens} reasoning_tokens={reasoning_tokens}"
)
assert response.usage.total_tokens == response.usage.prompt_tokens + reasoning_tokens
details = response.usage.completion_tokens_details
assert details.text_tokens + details.reasoning_tokens == response.usage.completion_tokens
def test_fallback_counts_reasoning_and_text_together(self):
reasoning = "First I should check whether the input is sorted. " * 10
text = "The list is already sorted, so no work is needed."
chunks = [_make_chunk(reasoning_content=reasoning), _make_chunk(content=text)]
response = litellm.stream_chunk_builder(chunks=chunks, messages=[{"role": "user", "content": "Sort it."}])
text_only = litellm.token_counter(model="claude-sonnet-4-6", text=text, count_response_tokens=True)
details = response.usage.completion_tokens_details
assert details.reasoning_tokens > 0
assert response.usage.completion_tokens == text_only + details.reasoning_tokens
assert details.text_tokens == text_only
def test_lone_usage_event_with_finish_reason_is_trusted(self):
chunks = [
_make_chunk(content="Yes, "),
_make_chunk(content="that works."),
_make_chunk(
usage=Usage(prompt_tokens=20, completion_tokens=5, total_tokens=25),
finish_reason="stop",
),
]
processor = ChunkProcessor(chunks=chunks, messages=[])
result = processor._calculate_usage_per_chunk(chunks=chunks)
assert result["completion_tokens"] == 5
def test_dict_chunks_with_finish_reason_are_trusted(self):
chunks = [
{
"_hidden_params": {"custom_llm_provider": "anthropic"},
"choices": [{"delta": {"content": "Yes, "}, "finish_reason": None}],
},
{
"_hidden_params": {"custom_llm_provider": "anthropic"},
"choices": [{"delta": {"content": "that works."}, "finish_reason": "stop"}],
"usage": Usage(prompt_tokens=20, completion_tokens=5, total_tokens=25),
},
]
processor = ChunkProcessor(chunks=chunks, messages=[])
result = processor._calculate_usage_per_chunk(chunks=chunks)
assert result["completion_tokens"] == 5
def test_dict_chunks_without_finish_reason_reset_placeholder(self):
chunks = [
{
"_hidden_params": {"custom_llm_provider": "anthropic"},
"choices": [],
"usage": Usage(prompt_tokens=20, completion_tokens=1, total_tokens=21),
},
{
"_hidden_params": {"custom_llm_provider": "anthropic"},
"choices": [{"delta": {"content": "partial"}, "finish_reason": None}],
},
]
processor = ChunkProcessor(chunks=chunks, messages=[])
result = processor._calculate_usage_per_chunk(chunks=chunks)
assert result["completion_tokens"] == 0
assert result["completion_tokens_details"] is None
def test_estimated_reasoning_is_capped_to_trusted_completion_total(self):
chunks = [
_make_chunk(reasoning_content="Let me reason about this carefully and at length. " * 20),
_make_chunk(
finish_reason="stop",
usage=Usage(prompt_tokens=20, completion_tokens=5, total_tokens=25),
),
]
response = litellm.stream_chunk_builder(
chunks=chunks,
messages=[{"role": "user", "content": "Go."}],
)
details = response.usage.completion_tokens_details
assert response.usage.completion_tokens == 5
assert details.reasoning_tokens <= response.usage.completion_tokens
assert details.reasoning_tokens + details.text_tokens == response.usage.completion_tokens
assert details.text_tokens >= 0
class TestProviderGuard:
"""Class A: the cursor-reset heuristic must NOT silently affect non-Anthropic
@ -297,11 +391,12 @@ class TestNonAnthropicStreamingIntact:
"""Make sure providers without cursor pattern still work."""
def test_completion_tokens_above_one_never_resets(self):
"""Any chunk reporting completion_tokens > 1 sets saw_non_cursor
and prevents the reset."""
"""A non-Anthropic provider reporting completion_tokens > 1 from a
single usage event keeps that value."""
chunks = [
_make_chunk(
usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
custom_llm_provider="openai",
),
]
processor = ChunkProcessor(chunks=chunks, messages=[])