From 82a0a443c642aacce3ed69c2afa39597ae46028a Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 3 Jul 2025 10:30:45 -0700 Subject: [PATCH] feat(stream_chunk_builder_utils.py): correctly return web_search_requests on stream chunk builder --- .../streaming_chunk_builder_utils.py | 39 +++++++++++++++++-- .../litellm_core_utils/streaming_handler.py | 5 ++- litellm/types/utils.py | 4 ++ tests/llm_translation/test_gemini.py | 1 + tests/test_litellm/types/test_types_utils.py | 34 ++++++++++++++++ 5 files changed, 78 insertions(+), 5 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index 0517d27e299..0cfdd138e69 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -17,6 +17,7 @@ from litellm.types.utils import ( ModelResponse, ModelResponseStream, PromptTokensDetails, + PromptTokensDetailsWrapper, Usage, ) from litellm.utils import print_verbose, token_counter @@ -256,7 +257,7 @@ class ChunkProcessor: cache_creation_input_tokens: Optional[int] = None cache_read_input_tokens: Optional[int] = None completion_tokens_details: Optional[CompletionTokensDetails] = None - prompt_tokens_details: Optional[PromptTokensDetails] = None + prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None if "prompt_tokens" in usage_chunk: prompt_tokens = usage_chunk.get("prompt_tokens", 0) or 0 @@ -277,10 +278,12 @@ class ChunkProcessor: 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 = PromptTokensDetails( + prompt_tokens_details = PromptTokensDetailsWrapper( **usage_chunk.prompt_tokens_details ) - elif isinstance(usage_chunk.prompt_tokens_details, PromptTokensDetails): + elif isinstance( + usage_chunk.prompt_tokens_details, PromptTokensDetailsWrapper + ): prompt_tokens_details = usage_chunk.prompt_tokens_details return { @@ -324,8 +327,10 @@ class ChunkProcessor: ## anthropic prompt caching information ## cache_creation_input_tokens: Optional[int] = None cache_read_input_tokens: Optional[int] = None + + web_search_requests: Optional[int] = None completion_tokens_details: Optional[CompletionTokensDetails] = None - prompt_tokens_details: Optional[PromptTokensDetails] = None + prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None for chunk in chunks: usage_chunk: Optional[Usage] = None if "usage" in chunk: @@ -366,6 +371,20 @@ class ChunkProcessor: completion_tokens_details = usage_chunk_dict[ "completion_tokens_details" ] + if ( + usage_chunk_dict["prompt_tokens_details"] is not None + and getattr( + usage_chunk_dict["prompt_tokens_details"], + "web_search_requests", + None, + ) + is not None + ): + web_search_requests = getattr( + usage_chunk_dict["prompt_tokens_details"], + "web_search_requests", + ) + prompt_tokens_details = usage_chunk_dict["prompt_tokens_details"] try: returned_usage.prompt_tokens = prompt_tokens or token_counter( @@ -415,8 +434,20 @@ class ChunkProcessor: if prompt_tokens_details is not None: returned_usage.prompt_tokens_details = prompt_tokens_details + if web_search_requests is not None: + if returned_usage.prompt_tokens_details is None: + returned_usage.prompt_tokens_details = PromptTokensDetailsWrapper( + web_search_requests=web_search_requests + ) + else: + returned_usage.prompt_tokens_details.web_search_requests = ( + web_search_requests + ) + # Return a new usage object with the new values + returned_usage = Usage(**returned_usage.model_dump()) + return returned_usage diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 5a592e6092f..07de7d647d6 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1204,7 +1204,9 @@ class CustomStreamWrapper: if response_obj is None: return completion_obj["content"] = response_obj["text"] - self.intermittent_finish_reason = response_obj.get("finish_reason", None) + self.intermittent_finish_reason = response_obj.get( + "finish_reason", None + ) if response_obj["is_finished"]: if response_obj["finish_reason"] == "error": raise Exception( @@ -1563,6 +1565,7 @@ class CustomStreamWrapper: complete_streaming_response = litellm.stream_chunk_builder( chunks=self.chunks, messages=self.messages ) + response = self.model_response_creator() if complete_streaming_response is not None: setattr( diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 7b75784326c..54cc158ef2d 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -915,6 +915,9 @@ class Usage(CompletionUsage): server_tool_use: Optional[ServerToolUse] = None + prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None + """Breakdown of tokens used in the prompt.""" + def __init__( self, prompt_tokens: Optional[int] = None, @@ -949,6 +952,7 @@ class Usage(CompletionUsage): # handle prompt_tokens_details _prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None + if prompt_tokens_details: if isinstance(prompt_tokens_details, dict): _prompt_tokens_details = PromptTokensDetailsWrapper( diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py index 0010a35773a..e0e1223f1cd 100644 --- a/tests/llm_translation/test_gemini.py +++ b/tests/llm_translation/test_gemini.py @@ -252,6 +252,7 @@ def test_gemini_with_grounding(): ) chunks = [] for chunk in response: + print(f"received chunk: {chunk}") chunks.append(chunk) print(f"chunks before stream_chunk_builder: {chunks}") assert len(chunks) > 0 diff --git a/tests/test_litellm/types/test_types_utils.py b/tests/test_litellm/types/test_types_utils.py index 6bb998b8d1b..1bf005db503 100644 --- a/tests/test_litellm/types/test_types_utils.py +++ b/tests/test_litellm/types/test_types_utils.py @@ -39,3 +39,37 @@ def test_empty_choices(): from litellm.types.utils import Choices Choices() + + +def test_usage_dump(): + from litellm.types.utils import ( + CompletionTokensDetailsWrapper, + PromptTokensDetailsWrapper, + Usage, + ) + + current_usage = Usage( + completion_tokens=37, + prompt_tokens=7, + total_tokens=44, + completion_tokens_details=CompletionTokensDetailsWrapper( + accepted_prediction_tokens=None, + audio_tokens=None, + reasoning_tokens=0, + rejected_prediction_tokens=None, + text_tokens=None, + ), + prompt_tokens_details=PromptTokensDetailsWrapper( + audio_tokens=None, + cached_tokens=None, + text_tokens=7, + image_tokens=None, + web_search_requests=1, + ), + web_search_requests=None, + ) + + assert current_usage.prompt_tokens_details.web_search_requests == 1 + + new_usage = Usage(**current_usage.model_dump()) + assert new_usage.prompt_tokens_details.web_search_requests == 1