mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-20 00:11:50 +00:00
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:
commit
f93b31679b
2 changed files with 174 additions and 49 deletions
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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=[])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue