mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(streaming): preserve reasoning tokens and usage in chunk builder
This commit is contained in:
parent
e5da59336d
commit
63abae70f1
3 changed files with 145 additions and 28 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue