This commit is contained in:
Zuraiz Anjum 2026-10-05 08:59:53 +08:00 • committed by GitHub
commit aaaf593a19
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 79 additions and 10 deletions

View file

@ -279,7 +279,7 @@ class CustomStreamWrapper:
self._repeated_messages_count = 1
self.is_function_call = self.check_is_function_call(logging_obj=logging_obj)
self.created: int | None = None
self._last_returned_hidden_params: dict | None = None
self._last_returned_hidden_params: dict[str, object] | None = None
_cached_logging_provider: Final = self.logging_obj.model_call_details.get("custom_llm_provider", None)
self._cached_logging_llm_provider: str | None = _cached_logging_provider
@ -1837,8 +1837,9 @@ class CustomStreamWrapper:
continue
# add usage as hidden param
if self.sent_last_chunk is True and self.stream_options is None:
usage = calculate_total_usage(chunks=self.chunks)
response._hidden_params["usage"] = usage
usage = _reported_total_usage(chunks=self.chunks)
if usage is not None:
response._hidden_params["usage"] = usage
self._last_returned_hidden_params = response._hidden_params
# Add MCP metadata to final chunk if present
response = self._add_mcp_metadata_to_final_chunk(response)
@ -1935,8 +1936,10 @@ class CustomStreamWrapper:
if self.received_finish_reason is not None or self.intermittent_finish_reason is not None:
self.chunks.append(processed_chunk)
if self.stream_options is None: # add usage as hidden param
usage = calculate_total_usage(chunks=self.chunks)
processed_chunk._hidden_params["usage"] = usage
self._last_returned_hidden_params = processed_chunk._hidden_params
usage = _reported_total_usage(chunks=self.chunks)
if usage is not None:
self._last_returned_hidden_params["usage"] = usage
## LOGGING
executor.submit(
self.run_success_logging_and_cache_storage,
@ -2047,8 +2050,9 @@ class CustomStreamWrapper:
# add usage as hidden param
if self.sent_last_chunk is True and self.stream_options is None:
usage = calculate_total_usage(chunks=self.chunks)
processed_chunk._hidden_params["usage"] = usage
usage = _reported_total_usage(chunks=self.chunks)
if usage is not None:
processed_chunk._hidden_params["usage"] = usage
self._last_returned_hidden_params = processed_chunk._hidden_params
# Call post-call streaming deployment hook for final chunk
@ -2203,8 +2207,10 @@ class CustomStreamWrapper:
if self.received_finish_reason is not None or self.intermittent_finish_reason is not None:
self.chunks.append(processed_chunk)
if self.stream_options is None:
usage: Final = calculate_total_usage(chunks=self.chunks)
processed_chunk._hidden_params["usage"] = usage # pyright: ignore[reportPrivateUsage] # sync parity
self._last_returned_hidden_params = processed_chunk._hidden_params # pyright: ignore[reportPrivateUsage] # sync parity
usage: Final = _reported_total_usage(chunks=self.chunks)
if usage is not None:
self._last_returned_hidden_params["usage"] = usage
# see sync __next__'s sibling branch: deliberately do NOT restore
# here - this chunk is still this call's own data, and restoring
# before returning it would corrupt the caller's own log
@ -2412,7 +2418,13 @@ def _coerce_token_details(
return details_type(**(raw if isinstance(raw, dict) else raw.model_dump()))
def calculate_total_usage(chunks: list[ModelResponse]) -> Usage:
def _reported_total_usage(chunks: Sequence[ModelResponse]) -> Usage | None:
if not any("usage" in chunk and chunk["usage"] is not None for chunk in chunks):
return None
return calculate_total_usage(chunks=chunks)
def calculate_total_usage(chunks: Sequence[ModelResponse]) -> Usage:
"""Assume most recent usage chunk has total usage uptil then."""
from litellm.litellm_core_utils.streaming_chunk_builder_utils import (
attach_cache_creation_token_details,

View file

@ -1,3 +1,4 @@
import datetime
import json
import time
from unittest.mock import AsyncMock, MagicMock, Mock, patch
@ -2346,6 +2347,62 @@ def test_usage_chunk_after_finish_reason_updates_hidden_params(logging_obj):
), f"Expected completion_tokens=135 from provider, got {hidden_usage.completion_tokens}"
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.parametrize("ends_with_finish_chunk", [True, False])
@pytest.mark.asyncio
async def test_stream_without_provider_usage_falls_back_to_the_token_estimate(
sync_mode: bool, ends_with_finish_chunk: bool
) -> None:
model: Final = "gpt-4o"
messages: Final = [{"role": "user", "content": "Write a long sentence about a fox. " * 30}]
text: Final = "The quick brown fox jumps over the lazy dog. " * 20
content_chunk: Final = ModelResponseStream(
id="chatcmpl-no-usage",
created=1741037890,
model=model,
choices=[StreamingChoices(index=0, delta=Delta(role="assistant", content=text))],
)
finish_chunk: Final = ModelResponseStream(
id="chatcmpl-no-usage",
created=1741037890,
model=model,
choices=[StreamingChoices(index=0, delta=Delta(content=""), finish_reason="stop")],
)
wrapper: Final = CustomStreamWrapper(
completion_stream=ModelResponseListIterator(
model_responses=[content_chunk, finish_chunk] if ends_with_finish_chunk else [content_chunk]
),
model=model,
custom_llm_provider="openai",
logging_obj=Logging(
model=model,
messages=messages,
stream=True,
call_type="completion",
start_time=datetime.datetime(2025, 3, 3, 21, 38, 10),
litellm_call_id="no-usage-call",
function_id="no-usage-fn",
),
stream_options=None,
)
collected: Final = [chunk for chunk in wrapper] if sync_mode else [chunk async for chunk in wrapper]
expected: Final = (
litellm.token_counter(model=model, messages=messages),
litellm.token_counter(model=model, text=text, count_response_tokens=True),
)
assembled: Final = litellm.stream_chunk_builder(chunks=collected, messages=messages)
assert assembled is not None
assert (assembled.usage.prompt_tokens, assembled.usage.completion_tokens) == expected
hidden_usage: Final = collected[-1]._hidden_params["usage"]
assert (hidden_usage.prompt_tokens, hidden_usage.completion_tokens) == expected
rates: Final = litellm.model_cost[model]
spend: Final = expected[0] * rates["input_cost_per_token"] + expected[1] * rates["output_cost_per_token"]
assert litellm.completion_cost(completion_response=assembled) == pytest.approx(spend)
@pytest.mark.asyncio
async def test_custom_stream_wrapper_aclose():
"""Test that aclose() delegates to the underlying completion_stream's aclose()"""