fix(streaming): preserve reasoning tokens and usage in chunk builder

This commit is contained in:
Mohammed Alshyakh 2026-09-08 21:57:01 +02:00
parent e5da59336d
commit 63abae70f1
3 changed files with 145 additions and 28 deletions

View file

@ -759,26 +759,27 @@ class ChunkProcessor:
prompt_tokens_details: PromptTokensDetailsWrapper | None = None
cost: float | None = None
if "prompt_tokens" in usage_chunk:
prompt_tokens = usage_chunk.get("prompt_tokens", 0) or 0
if "completion_tokens" in usage_chunk:
completion_tokens = usage_chunk.get("completion_tokens", 0) or 0
if "cache_creation_input_tokens" in usage_chunk:
cache_creation_input_tokens = usage_chunk.get("cache_creation_input_tokens")
if "cache_read_input_tokens" in usage_chunk:
cache_read_input_tokens = usage_chunk.get("cache_read_input_tokens")
if "cost" in usage_chunk:
cost = usage_chunk.get("cost")
if hasattr(usage_chunk, "completion_tokens_details"):
if isinstance(usage_chunk.completion_tokens_details, dict):
completion_tokens_details = CompletionTokensDetails(**usage_chunk.completion_tokens_details)
elif isinstance(usage_chunk.completion_tokens_details, CompletionTokensDetails):
completion_tokens_details = usage_chunk.completion_tokens_details
if hasattr(usage_chunk, "prompt_tokens_details"):
if isinstance(usage_chunk.prompt_tokens_details, dict):
prompt_tokens_details = PromptTokensDetailsWrapper(**usage_chunk.prompt_tokens_details)
elif isinstance(usage_chunk.prompt_tokens_details, PromptTokensDetailsWrapper):
prompt_tokens_details = usage_chunk.prompt_tokens_details
_get = (
usage_chunk.get
if isinstance(usage_chunk, dict)
else lambda k, default=None: getattr(usage_chunk, k, default)
)
prompt_tokens = _get("prompt_tokens", 0) or 0
completion_tokens = _get("completion_tokens", 0) or 0
cache_creation_input_tokens = _get("cache_creation_input_tokens")
cache_read_input_tokens = _get("cache_read_input_tokens")
cost = _get("cost")
raw_completion_details = _get("completion_tokens_details")
if isinstance(raw_completion_details, dict):
completion_tokens_details = CompletionTokensDetails(**raw_completion_details)
elif isinstance(raw_completion_details, CompletionTokensDetails):
completion_tokens_details = raw_completion_details
raw_prompt_details = _get("prompt_tokens_details")
if isinstance(raw_prompt_details, dict):
prompt_tokens_details = PromptTokensDetailsWrapper(**raw_prompt_details)
elif isinstance(raw_prompt_details, PromptTokensDetailsWrapper):
prompt_tokens_details = raw_prompt_details
return {
"prompt_tokens": prompt_tokens,
@ -1080,6 +1081,14 @@ class ChunkProcessor:
returned_usage.completion_tokens_details.text_tokens = (
returned_usage.completion_tokens - capped_reasoning_tokens
)
effective_reasoning_tokens: Final[int | None] = (
getattr(returned_usage.completion_tokens_details, "reasoning_tokens", None)
if returned_usage.completion_tokens_details is not None
else None
)
if effective_reasoning_tokens is not None and returned_usage.completion_tokens < effective_reasoning_tokens:
returned_usage.completion_tokens = max(returned_usage.completion_tokens, effective_reasoning_tokens)
returned_usage.total_tokens = returned_usage.prompt_tokens + returned_usage.completion_tokens
if prompt_tokens_details is not None:
returned_usage.prompt_tokens_details = prompt_tokens_details

View file

@ -1527,11 +1527,7 @@ class CustomStreamWrapper:
setattr(
model_response,
"usage",
litellm.Usage(
prompt_tokens=response_obj["usage"].get("prompt_tokens", None) or None,
completion_tokens=response_obj["usage"].get("completion_tokens", None) or None,
total_tokens=response_obj["usage"].get("total_tokens", None) or None,
),
litellm.Usage(**response_obj["usage"]),
)
elif isinstance(response_obj["usage"], Usage):
setattr(
@ -1656,7 +1652,12 @@ class CustomStreamWrapper:
self.tool_call = True
if hasattr(chunk, "usage") and chunk.usage is not None:
model_response.usage = chunk.usage
if isinstance(chunk.usage, dict):
model_response.usage = litellm.Usage(**chunk.usage)
elif isinstance(chunk.usage, BaseModel) and not isinstance(chunk.usage, Usage):
model_response.usage = litellm.Usage(**chunk.usage.model_dump())
else:
model_response.usage = chunk.usage
## RETURN ARG
result: Final = self.return_processed_chunk_logic(

View file

@ -4,15 +4,14 @@ from typing import Final
import pytest
from litellm import ChatCompletionUsageBlock, stream_chunk_builder
from litellm.types.utils import GenericStreamingChunk
from litellm.litellm_core_utils.streaming_chunk_builder_utils import ChunkProcessor
from litellm.types.utils import (
ChatCompletionDeltaToolCall,
ChatCompletionMessageToolCall,
Delta,
Function,
GenericStreamingChunk,
ModelResponseStream,
PromptTokensDetails,
ServerToolUse,
@ -1554,3 +1553,111 @@ def test_stream_chunk_builder_reads_role_from_first_frame_with_choices() -> None
assert response is not None
assert response.choices[0].message.role == "user"
assert response.choices[0].message.content == "Hi"
def test_calculate_usage_preserves_reasoning_tokens_and_counts_from_completion_usage():
from unittest.mock import MagicMock
from openai.types.completion_usage import CompletionTokensDetails, CompletionUsage
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
chunk_sdk = ModelResponseStream(
id="chatcmpl-reasoning-sdk",
model="gpt-4o",
choices=[StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="42", role="assistant"))],
usage=CompletionUsage(
prompt_tokens=25,
completion_tokens=326,
total_tokens=351,
completion_tokens_details=CompletionTokensDetails(reasoning_tokens=326),
),
)
processor_sdk = ChunkProcessor(chunks=[chunk_sdk])
usage_sdk = processor_sdk.calculate_usage(
chunks=[chunk_sdk],
model="gpt-4o",
completion_output="42",
)
assert usage_sdk.prompt_tokens == 25
assert usage_sdk.completion_tokens == 326
assert usage_sdk.total_tokens == 351
assert usage_sdk.completion_tokens_details is not None
assert usage_sdk.completion_tokens_details.reasoning_tokens == 326
chunk_dict = {
"id": "chatcmpl-reasoning-dict",
"model": "gpt-4o",
"choices": [{"finish_reason": "stop", "index": 0, "delta": {"content": "42"}}],
"usage": {
"prompt_tokens": 25,
"completion_tokens": 326,
"total_tokens": 351,
"completion_tokens_details": {"reasoning_tokens": 326},
},
}
processor_dict = ChunkProcessor(chunks=[chunk_dict])
usage_dict = processor_dict.calculate_usage(
chunks=[chunk_dict],
model="gpt-4o",
completion_output="42",
)
assert usage_dict.prompt_tokens == 25
assert usage_dict.completion_tokens == 326
assert usage_dict.total_tokens == 351
assert usage_dict.completion_tokens_details is not None
assert usage_dict.completion_tokens_details.reasoning_tokens == 326
chunk_under = ModelResponseStream(
id="chatcmpl-reasoning-under",
model="gpt-4o",
choices=[StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="hi"))],
usage=Usage(
prompt_tokens=10,
completion_tokens=5,
total_tokens=15,
completion_tokens_details={"reasoning_tokens": 50},
),
)
usage_under = ChunkProcessor(chunks=[chunk_under]).calculate_usage(
chunks=[chunk_under],
model="gpt-4o",
completion_output="hi",
)
assert usage_under.completion_tokens >= 50
assert usage_under.total_tokens == usage_under.prompt_tokens + usage_under.completion_tokens
wrapper = CustomStreamWrapper(
completion_stream=iter([]),
model="gpt-4o",
custom_llm_provider="openai",
logging_obj=MagicMock(),
)
mock_chunk_dict = ModelResponseStream(
id="chatcmpl-dict-usage",
choices=[StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="hi"))],
)
mock_chunk_dict.usage = {
"prompt_tokens": 25,
"completion_tokens": 326,
"total_tokens": 351,
"completion_tokens_details": {"reasoning_tokens": 326},
}
parsed_response_dict = wrapper.chunk_creator(mock_chunk_dict)
assert parsed_response_dict.usage is not None
assert parsed_response_dict.usage.completion_tokens_details is not None
assert parsed_response_dict.usage.completion_tokens_details.reasoning_tokens == 326
mock_chunk_sdk = ModelResponseStream(
id="chatcmpl-sdk-usage",
choices=[StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="hi"))],
)
mock_chunk_sdk.usage = CompletionUsage(
prompt_tokens=25,
completion_tokens=326,
total_tokens=351,
completion_tokens_details=CompletionTokensDetails(reasoning_tokens=326),
)
parsed_response_sdk = wrapper.chunk_creator(mock_chunk_sdk)
assert parsed_response_sdk.usage is not None
assert parsed_response_sdk.usage.completion_tokens_details is not None
assert parsed_response_sdk.usage.completion_tokens_details.reasoning_tokens == 326