diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 9d33f86d841..0b3d66be6ae 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -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, diff --git a/tests/unit/litellm_core_utils/test_streaming_handler.py b/tests/unit/litellm_core_utils/test_streaming_handler.py index f88e082d577..f28ab873c80 100644 --- a/tests/unit/litellm_core_utils/test_streaming_handler.py +++ b/tests/unit/litellm_core_utils/test_streaming_handler.py @@ -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()"""