fix(streaming): restore the token estimate for streams that send no usage

A stream without stream_options got a made-up Usage(0, 0) written into its
final chunk's hidden params. Since #42323 the chunk builder treats that as a
provider-reported zero, so these streams were logged and billed at 0 tokens
and $0 instead of the tokenizer estimate. Only attach the summed usage when
some chunk actually carried usage.
This commit is contained in:
Zuraiz Anjum 2026-10-03 00:41:36 +05:00
parent f932e292c5
commit 9e51552d60
2 changed files with 74 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)
@ -1931,8 +1932,10 @@ class CustomStreamWrapper:
self.sent_last_chunk = True
processed_chunk: Final = self.finish_reason_handler()
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,
@ -2043,8 +2046,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
@ -2195,8 +2199,10 @@ class CustomStreamWrapper:
self.sent_last_chunk = True
processed_chunk: Final = self.finish_reason_handler()
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
@ -2404,7 +2410,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

@ -2346,6 +2346,58 @@ 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=time.time(),
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
@pytest.mark.asyncio
async def test_custom_stream_wrapper_aclose():
"""Test that aclose() delegates to the underlying completion_stream's aclose()"""