From 63abae70f147cb8f85b24707f3c0ab3040a4339f Mon Sep 17 00:00:00 2001 From: Mohammed Alshyakh Date: Tue, 8 Sep 2026 21:57:01 +0200 Subject: [PATCH] fix(streaming): preserve reasoning tokens and usage in chunk builder --- .../streaming_chunk_builder_utils.py | 49 ++++---- .../litellm_core_utils/streaming_handler.py | 13 +- .../test_streaming_chunk_builder_utils.py | 111 +++++++++++++++++- 3 files changed, 145 insertions(+), 28 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index 90698296142..37286fbd6d9 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -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 diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index db23929e0c3..4ee160fbae6 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -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( diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py index 626b8a63b20..dc603a25caa 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py @@ -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 +