feat(stream_chunk_builder_utils.py): correctly return web_search_requests on stream chunk builder

This commit is contained in:
Krrish Dholakia 2025-07-03 10:30:45 -07:00 • committed by Ishaan Jaff
parent c3909d6f50
commit 82a0a443c6
5 changed files with 78 additions and 5 deletions

View file

@ -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

View file

@ -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(

View file

@ -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(

View file

@ -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

View file

@ -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