mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(stream_chunk_builder_utils.py): correctly return web_search_requests on stream chunk builder
This commit is contained in:
parent
c3909d6f50
commit
82a0a443c6
5 changed files with 78 additions and 5 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue