fix(streaming): cap estimated reasoning tokens to the provider total and cover dict chunks

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-09-16 23:27:48 +00:00
parent 43514a7ffe
commit 884e96958c
2 changed files with 68 additions and 33 deletions

View file

@ -5,7 +5,6 @@ from itertools import groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, TypeAlias, TypedDict, Union, cast
from pydantic import BaseModel
from typing_extensions import ReadOnly, Required
from litellm._logging import verbose_logger
@ -974,12 +973,10 @@ class ChunkProcessor:
return None
@staticmethod
def _chunk_choices(chunk: "_UsageBearingChunk | BaseModel") -> Sequence[object]:
def _chunk_choices(chunk: "_UsageBearingChunk | ModelResponse | ModelResponseStream") -> Sequence[object]:
if isinstance(chunk, dict):
return chunk.get("choices", ())
if isinstance(chunk, (ModelResponse, ModelResponseStream)):
return chunk.choices
return ()
return getattr(chunk, "choices", ())
@staticmethod
def _saw_finish_reason(chunks: Sequence["_UsageBearingChunk | ModelResponse"]) -> bool:
@ -1099,19 +1096,22 @@ class ChunkProcessor:
if reasoning_tokens is not None:
if returned_usage.completion_tokens_details is None:
capped_reasoning_tokens: Final = min(max(0, reasoning_tokens), returned_usage.completion_tokens)
returned_usage.completion_tokens_details = CompletionTokensDetailsWrapper(
reasoning_tokens=reasoning_tokens,
text_tokens=max(0, returned_usage.completion_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
existing_capped_reasoning_tokens: Final = min(
max(0, reasoning_tokens), returned_usage.completion_tokens
)
returned_usage.completion_tokens_details.reasoning_tokens = existing_capped_reasoning_tokens
if returned_usage.completion_tokens_details.text_tokens is None:
returned_usage.completion_tokens_details.text_tokens = (
returned_usage.completion_tokens - capped_reasoning_tokens
returned_usage.completion_tokens - existing_capped_reasoning_tokens
)
if prompt_tokens_details is not None:
returned_usage.prompt_tokens_details = prompt_tokens_details

View file

@ -19,7 +19,6 @@ 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
@ -71,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"),
@ -99,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(
@ -121,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(
@ -151,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
@ -195,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(
@ -233,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
@ -312,6 +296,57 @@ class TestAnthropicCursorBug:
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