mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
43514a7ffe
commit
884e96958c
2 changed files with 68 additions and 33 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue