mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
2fdab5684e
commit
923ff2e327
9 changed files with 189 additions and 36 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
@ -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 ##
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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