Litellm dev fix gemini web search tracking (#12288)

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

* fix(types/utils.py): handle prompttokendetails

* fix(stream_chunk_builder_utils.py): fix ruff check error

* test: try-except rate limit error

* fix: fix import
This commit is contained in:
Krish Dholakia 2025-07-03 12:25:40 -07:00 • committed by GitHub
parent 2fdab5684e
commit 923ff2e327
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 189 additions and 36 deletions

View file

@ -1,6 +1,6 @@
import base64
import time
from typing import Any, Dict, List, Optional, Union, cast
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast
from litellm.types.llms.openai import (
ChatCompletionAssistantContentValue,
@ -16,11 +16,16 @@ from litellm.types.utils import (
FunctionCall,
ModelResponse,
ModelResponseStream,
PromptTokensDetails,
PromptTokensDetailsWrapper,
Usage,
)
from litellm.utils import print_verbose, token_counter
if TYPE_CHECKING:
from litellm.types.litellm_core_utils.streaming_chunk_builder_utils import (
UsagePerChunk,
)
class ChunkProcessor:
def __init__(self, chunks: List, messages: Optional[list] = None):
@ -256,7 +261,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 +282,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 {
@ -306,26 +313,24 @@ class ChunkProcessor:
return reasoning_tokens
def calculate_usage(
def _calculate_usage_per_chunk(
self,
chunks: List[Union[Dict[str, Any], ModelResponse]],
model: str,
completion_output: str,
messages: Optional[List] = None,
reasoning_tokens: Optional[int] = None,
) -> Usage:
"""
Calculate usage for the given chunks.
"""
returned_usage = Usage()
) -> "UsagePerChunk":
from litellm.types.litellm_core_utils.streaming_chunk_builder_utils import (
UsagePerChunk,
)
# # Update usage information if needed
prompt_tokens = 0
completion_tokens = 0
## 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,7 +371,67 @@ 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"]
return UsagePerChunk(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens,
web_search_requests=web_search_requests,
completion_tokens_details=completion_tokens_details,
prompt_tokens_details=prompt_tokens_details,
)
def calculate_usage(
self,
chunks: List[Union[Dict[str, Any], ModelResponse]],
model: str,
completion_output: str,
messages: Optional[List] = None,
reasoning_tokens: Optional[int] = None,
) -> Usage:
"""
Calculate usage for the given chunks.
"""
returned_usage = Usage()
# # Update usage information if needed
calculated_usage_per_chunk = self._calculate_usage_per_chunk(chunks=chunks)
prompt_tokens = calculated_usage_per_chunk["prompt_tokens"]
completion_tokens = calculated_usage_per_chunk["completion_tokens"]
## anthropic prompt caching information ##
cache_creation_input_tokens: Optional[int] = calculated_usage_per_chunk[
"cache_creation_input_tokens"
]
cache_read_input_tokens: Optional[int] = calculated_usage_per_chunk[
"cache_read_input_tokens"
]
web_search_requests: Optional[int] = calculated_usage_per_chunk[
"web_search_requests"
]
completion_tokens_details: Optional[CompletionTokensDetails] = (
calculated_usage_per_chunk["completion_tokens_details"]
)
prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = (
calculated_usage_per_chunk["prompt_tokens_details"]
)
try:
returned_usage.prompt_tokens = prompt_tokens or token_counter(
model=model, messages=messages
@ -415,8 +480,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

@ -4937,7 +4937,10 @@ def transcription(
provider_config=provider_config,
litellm_params=litellm_params_dict,
)
elif custom_llm_provider in [LlmProviders.DEEPGRAM.value, LlmProviders.ELEVENLABS.value]:
elif custom_llm_provider in [
LlmProviders.DEEPGRAM.value,
LlmProviders.ELEVENLABS.value,
]:
response = base_llm_http_handler.audio_transcriptions(
model=model,
audio_file=file,

View file

@ -0,0 +1,13 @@
from typing import TYPE_CHECKING, Optional, TypedDict
from ..utils import CompletionTokensDetails, PromptTokensDetailsWrapper
class UsagePerChunk(TypedDict):
prompt_tokens: int
completion_tokens: int
cache_creation_input_tokens: Optional[int]
cache_read_input_tokens: Optional[int]
web_search_requests: Optional[int]
completion_tokens_details: Optional[CompletionTokensDetails]
prompt_tokens_details: Optional[PromptTokensDetailsWrapper]

View file

@ -910,13 +910,18 @@ 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,
completion_tokens: Optional[int] = None,
total_tokens: Optional[int] = None,
reasoning_tokens: Optional[int] = None,
prompt_tokens_details: Optional[Union[PromptTokensDetailsWrapper, dict]] = None,
prompt_tokens_details: Optional[
Union[PromptTokensDetailsWrapper, PromptTokensDetails, dict]
] = None,
completion_tokens_details: Optional[
Union[CompletionTokensDetailsWrapper, dict]
] = None,
@ -944,12 +949,17 @@ 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(
**prompt_tokens_details
)
elif isinstance(prompt_tokens_details, PromptTokensDetails):
_prompt_tokens_details = PromptTokensDetailsWrapper(
**prompt_tokens_details.model_dump()
)
elif isinstance(prompt_tokens_details, PromptTokensDetailsWrapper):
_prompt_tokens_details = prompt_tokens_details
## DEEPSEEK MAPPING ##

View file

@ -249,6 +249,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

@ -131,6 +131,7 @@ def test_null_role_response():
assert response.choices[0].message.role == "assistant"
@pytest.mark.skip(reason="Cohere having RBAC issues")
def test_completion_azure_command_r():
try:
@ -175,7 +176,6 @@ def test_completion_azure_ai_gpt_4o(api_base):
pytest.fail(f"Error occurred: {e}")
def predibase_mock_post(url, data=None, json=None, headers=None, timeout=None):
mock_response = MagicMock()
mock_response.status_code = 200
@ -940,6 +940,8 @@ def test_completion_mistral_api_mistral_large_function_call():
tool_choice="auto",
)
print(second_response)
except litellm.RateLimitError:
pass
except Exception as e:
pytest.fail(f"Error occurred: {e}")
@ -1475,7 +1477,6 @@ HF Tests we should pass
"""
@pytest.mark.parametrize(
"provider", ["openai", "hosted_vllm", "lm_studio", "llamafile"]
) # "vertex_ai",
@ -1560,6 +1561,7 @@ async def test_openai_compatible_custom_api_video(provider):
mock_call.assert_called_once()
def test_lm_studio_completion(monkeypatch):
monkeypatch.delenv("LM_STUDIO_API_KEY", raising=False)
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
@ -1578,6 +1580,7 @@ def test_lm_studio_completion(monkeypatch):
except litellm.APIError as e:
print(e)
# ################### Hugging Face Conversational models ########################
# def hf_test_completion_conv():
# try:
@ -1625,7 +1628,6 @@ def mock_post(url, **kwargs):
return mock_response
def test_ollama_image():
"""
Test that datauri prefixes are removed, JPEG/PNG images are passed
@ -4394,6 +4396,7 @@ def test_humanloop_completion(monkeypatch):
messages=[{"role": "user", "content": "Tell me a joke."}],
)
def test_completion_novita_ai():
litellm.set_verbose = True
messages = [
@ -4403,10 +4406,11 @@ def test_completion_novita_ai():
"content": "Hey",
},
]
from openai import OpenAI
openai_client = OpenAI(api_key="fake-key")
with patch.object(
openai_client.chat.completions, "create", new=MagicMock()
) as mock_call:
@ -4417,21 +4421,22 @@ def test_completion_novita_ai():
client=openai_client,
api_base="https://api.novita.ai/v3/openai",
)
mock_call.assert_called_once()
# Verify model is passed correctly
assert mock_call.call_args.kwargs["model"] == "meta-llama/llama-3.3-70b-instruct"
assert (
mock_call.call_args.kwargs["model"]
== "meta-llama/llama-3.3-70b-instruct"
)
# Verify messages are passed correctly
assert mock_call.call_args.kwargs["messages"] == messages
except Exception as e:
pytest.fail(f"Error occurred: {e}")
@pytest.mark.parametrize(
"api_key", ["my-bad-api-key"]
)
@pytest.mark.parametrize("api_key", ["my-bad-api-key"])
def test_completion_novita_ai_dynamic_params(api_key):
try:
litellm.set_verbose = True
@ -4442,12 +4447,15 @@ def test_completion_novita_ai_dynamic_params(api_key):
"content": "Hey",
},
]
from openai import OpenAI
openai_client = OpenAI(api_key="fake-key")
with patch.object(
openai_client.chat.completions, "create", side_effect=Exception("Invalid API key")
openai_client.chat.completions,
"create",
side_effect=Exception("Invalid API key"),
) as mock_call:
try:
completion(
@ -4461,11 +4469,12 @@ def test_completion_novita_ai_dynamic_params(api_key):
except Exception as e:
# This should fail with the mocked exception
assert "Invalid API key" in str(e)
mock_call.assert_called_once()
except Exception as e:
pytest.fail(f"Unexpected error: {e}")
def test_deepseek_reasoning_content_completion():
try:
litellm.set_verbose = True

View file

@ -975,6 +975,8 @@ def test_completion_mistral_api_mistral_large_function_call_with_streaming():
elif chunk.choices[0].finish_reason is not None: # last chunk
validate_final_streaming_function_calling_chunk(chunk=chunk)
idx += 1
except litellm.RateLimitError:
pass
except Exception as e:
pytest.fail(f"Error occurred: {e}")
@ -2892,6 +2894,7 @@ def test_azure_streaming_and_function_calling():
pytest.fail(f"Error occurred: {e}")
raise e
@pytest.mark.asyncio
async def test_azure_astreaming_and_function_calling():
import uuid
@ -4022,4 +4025,4 @@ def test_is_delta_empty():
tool_calls=None,
audio=None,
)
)
)

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