mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(gemini): count Anthropic messages, system and tools on /v1/messages/count_tokens and fall back on provider errors
This commit is contained in:
parent
cfa2830bde
commit
629a068e10
9 changed files with 481 additions and 29 deletions
|
|
@ -27,7 +27,13 @@ class BaseTokenCounter(ABC):
|
|||
tools: list[dict[str, Any]] | None = None,
|
||||
system: Any | None = None,
|
||||
) -> TokenCountResponse | None:
|
||||
pass
|
||||
"""Count tokens with the provider's API.
|
||||
|
||||
Exactly one of `messages` (Anthropic or OpenAI chat shape, translated by the counter) or `contents`
|
||||
(provider-native shape, forwarded as is) is set. Provider failures are returned as
|
||||
`TokenCountResponse(error=True, status_code=...)`, never raised, so the proxy can decide between the
|
||||
local fallback and surfacing the provider status. `None` means the counter has nothing to count
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def should_use_token_counting_api(
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import litellm
|
|||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
|
@ -474,6 +475,9 @@ def get_api_key_from_env() -> str | None:
|
|||
class GoogleAIStudioTokenCounter(BaseTokenCounter):
|
||||
"""Token counter implementation for Google AI Studio provider."""
|
||||
|
||||
def __init__(self, client: AsyncHTTPHandler | None = None) -> None:
|
||||
self.client: Final = client
|
||||
|
||||
def should_use_token_counting_api(
|
||||
self,
|
||||
custom_llm_provider: str | None = None,
|
||||
|
|
@ -495,25 +499,57 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
|
|||
import copy
|
||||
|
||||
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
|
||||
from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload
|
||||
|
||||
if contents is None and not messages:
|
||||
return None
|
||||
|
||||
payload: Final = (
|
||||
build_count_tokens_payload(model=model_to_use, messages=messages or [], system=system, tools=tools)
|
||||
if contents is None
|
||||
else None
|
||||
)
|
||||
deployment = deployment or {}
|
||||
count_tokens_params_request: Final = copy.deepcopy(deployment.get("litellm_params", {}))
|
||||
count_tokens_params: Final = {
|
||||
"model": model_to_use,
|
||||
"contents": contents,
|
||||
"contents": payload.contents if payload is not None else contents,
|
||||
**({"system_instruction": payload.system_instruction} if payload and payload.system_instruction else {}),
|
||||
**({"tools": payload.tools} if payload and payload.tools else {}),
|
||||
}
|
||||
count_tokens_params_request.update(count_tokens_params)
|
||||
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
**count_tokens_params_request,
|
||||
)
|
||||
|
||||
if result is not None:
|
||||
try:
|
||||
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
**count_tokens_params_request,
|
||||
client=self.client,
|
||||
)
|
||||
except (litellm.APIError, litellm.APIConnectionError) as e:
|
||||
return TokenCountResponse(
|
||||
total_tokens=result.get("totalTokens", 0),
|
||||
total_tokens=0,
|
||||
request_model=request_model,
|
||||
model_used=model_to_use,
|
||||
tokenizer_type=result.get("tokenizer_used", ""),
|
||||
original_response=result,
|
||||
tokenizer_type="gemini_api",
|
||||
error=True,
|
||||
error_message=e.message,
|
||||
status_code=e.status_code,
|
||||
)
|
||||
|
||||
return None
|
||||
if "totalTokens" not in result:
|
||||
return TokenCountResponse(
|
||||
total_tokens=0,
|
||||
request_model=request_model,
|
||||
model_used=model_to_use,
|
||||
tokenizer_type="gemini_api",
|
||||
original_response=result,
|
||||
error=True,
|
||||
error_message="Google Gen AI Studio countTokens response has no totalTokens",
|
||||
status_code=502,
|
||||
)
|
||||
|
||||
return TokenCountResponse(
|
||||
total_tokens=result["totalTokens"],
|
||||
request_model=request_model,
|
||||
model_used=model_to_use,
|
||||
tokenizer_type=result.get("tokenizer_used", ""),
|
||||
original_response=result,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,8 @@ from typing import TYPE_CHECKING, Any, Final
|
|||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client
|
||||
from litellm.types.llms.vertex_ai import SystemInstructions, Tools
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -84,6 +85,9 @@ class GoogleAIStudioTokenCounter:
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
system_instruction: SystemInstructions | None = None,
|
||||
tools: list[Tools] | None = None,
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
**kwargs: object,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
|
|
@ -96,6 +100,9 @@ class GoogleAIStudioTokenCounter:
|
|||
api_key: Optional Google API key (will fall back to environment)
|
||||
api_base: Optional API base URL (defaults to Google Gen AI Studio)
|
||||
timeout: Optional timeout for the request
|
||||
system_instruction: Optional Gemini systemInstruction, counted alongside contents
|
||||
tools: Optional Gemini tool declarations, counted alongside contents
|
||||
client: Optional HTTP client, defaults to the shared Gemini client
|
||||
**kwargs: Additional parameters
|
||||
|
||||
Returns:
|
||||
|
|
@ -114,9 +121,8 @@ class GoogleAIStudioTokenCounter:
|
|||
|
||||
Raises:
|
||||
ValueError: If API key is missing
|
||||
litellm.APIError: If the API call fails
|
||||
litellm.APIError: If the API call fails or returns a non-JSON body
|
||||
litellm.APIConnectionError: If the connection fails
|
||||
Exception: For any other unexpected errors
|
||||
"""
|
||||
|
||||
# Prepare headers
|
||||
|
|
@ -130,22 +136,25 @@ class GoogleAIStudioTokenCounter:
|
|||
|
||||
# Prepare request body - clean up contents to remove unsupported fields
|
||||
cleaned_contents: Final = self._clean_contents_for_gemini_api(contents)
|
||||
request_body: Final = {"contents": cleaned_contents}
|
||||
request_body: Final = (
|
||||
{"contents": cleaned_contents}
|
||||
if system_instruction is None and tools is None
|
||||
else {
|
||||
"generateContentRequest": {
|
||||
"model": f"models/{model}",
|
||||
"contents": cleaned_contents,
|
||||
**({"systemInstruction": system_instruction} if system_instruction is not None else {}),
|
||||
**({"tools": tools} if tools is not None else {}),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
async_httpx_client: Final = get_async_httpx_client(
|
||||
async_httpx_client: Final = client or get_async_httpx_client(
|
||||
llm_provider=LlmProviders.GEMINI,
|
||||
)
|
||||
|
||||
try:
|
||||
response: Final = await async_httpx_client.post(url=url, headers=headers, json=request_body)
|
||||
|
||||
# Check for HTTP errors
|
||||
response.raise_for_status()
|
||||
|
||||
# Parse response
|
||||
result: Final = response.json()
|
||||
return result
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_msg = f"Google Gen AI Studio API error: {e.response.status_code} - {e.response.text}"
|
||||
raise litellm.APIError(
|
||||
|
|
@ -157,6 +166,14 @@ class GoogleAIStudioTokenCounter:
|
|||
except httpx.RequestError as e:
|
||||
error_msg = f"Request to Google Gen AI Studio failed: {e}"
|
||||
raise litellm.APIConnectionError(message=error_msg, llm_provider="gemini", model=model) from e
|
||||
except Exception as e:
|
||||
error_msg = f"Unexpected error during token counting: {e}"
|
||||
raise Exception(error_msg) from e
|
||||
|
||||
try:
|
||||
result: Final = response.json()
|
||||
except ValueError as e:
|
||||
raise litellm.APIError(
|
||||
message=f"Google Gen AI Studio API returned a non-JSON body: {response.text}",
|
||||
llm_provider="gemini",
|
||||
model=model,
|
||||
status_code=response.status_code,
|
||||
) from e
|
||||
return result
|
||||
|
|
|
|||
55
litellm/llms/gemini/count_tokens/transformation.py
Normal file
55
litellm/llms/gemini/count_tokens/transformation.py
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
from dataclasses import dataclass
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
|
||||
LiteLLMAnthropicMessagesAdapter,
|
||||
)
|
||||
from litellm.llms.vertex_ai.gemini.transformation import (
|
||||
_gemini_convert_messages_with_history, # pyright: ignore[reportPrivateUsage] # shared helper already used by gemini/chat, context_caching, and vertex_and_google_ai_studio_gemini
|
||||
_transform_system_message, # pyright: ignore[reportPrivateUsage] # same helper the Gemini chat transformation uses to split system prompts
|
||||
)
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
||||
from litellm.types.llms.anthropic import AnthropicMessagesRequest
|
||||
from litellm.types.llms.vertex_ai import ContentType, SystemInstructions, Tools
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GeminiCountTokensPayload:
|
||||
contents: list[ContentType]
|
||||
system_instruction: SystemInstructions | None
|
||||
tools: list[Tools] | None
|
||||
|
||||
|
||||
def build_count_tokens_payload(
|
||||
model: str,
|
||||
messages: list[dict[str, Any]],
|
||||
system: object | None,
|
||||
tools: list[dict[str, Any]] | None,
|
||||
) -> GeminiCountTokensPayload:
|
||||
"""Translate an Anthropic Messages token-count request into the Gemini countTokens shape."""
|
||||
anthropic_request: Final = AnthropicMessagesRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
**({"system": system} if isinstance(system, (str, list)) else {}),
|
||||
**({"tools": tools} if tools else {}),
|
||||
)
|
||||
openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
|
||||
anthropic_request, custom_llm_provider="gemini"
|
||||
)
|
||||
system_instruction, remaining_messages = _transform_system_message(
|
||||
supports_system_message=True,
|
||||
messages=list(openai_request["messages"]),
|
||||
)
|
||||
contents: Final = _gemini_convert_messages_with_history(
|
||||
messages=remaining_messages, model=model, custom_llm_provider="gemini"
|
||||
)
|
||||
openai_tools: Final = openai_request.get("tools")
|
||||
gemini_tools: Final = (
|
||||
VertexGeminiConfig()._map_function( # pyright: ignore[reportPrivateUsage] # same tool mapper the Gemini chat path uses
|
||||
value=[dict(tool) for tool in openai_tools],
|
||||
optional_params={},
|
||||
)
|
||||
if openai_tools
|
||||
else None
|
||||
)
|
||||
return GeminiCountTokensPayload(contents=contents, system_instruction=system_instruction, tools=gemini_tools)
|
||||
|
|
@ -13467,6 +13467,20 @@ async def _try_provider_token_count(
|
|||
param="model",
|
||||
code=status_code,
|
||||
)
|
||||
except (litellm.APIError, litellm.APIConnectionError) as e:
|
||||
if litellm.disable_token_counter is True:
|
||||
raise ProxyException(
|
||||
message=e.message,
|
||||
type="token_counting_error",
|
||||
param="model",
|
||||
code=e.status_code,
|
||||
)
|
||||
verbose_proxy_logger.warning(
|
||||
"Provider token counting raised (%s): %s. Falling back to local tokenizer.",
|
||||
e.status_code,
|
||||
e.message,
|
||||
)
|
||||
return None
|
||||
if result is not None and result.error is True:
|
||||
if litellm.disable_token_counter is True:
|
||||
raise ProxyException(
|
||||
|
|
|
|||
|
|
@ -29,7 +29,8 @@ from litellm.proxy._types import ProxyException, TokenCountRequest
|
|||
from litellm.proxy.anthropic_endpoints.endpoints import (
|
||||
count_tokens as anthropic_count_tokens,
|
||||
)
|
||||
from litellm.proxy.proxy_server import token_counter
|
||||
from litellm.llms.base_llm.base_utils import BaseTokenCounter
|
||||
from litellm.proxy.proxy_server import _try_provider_token_count, token_counter
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
verbose_proxy_logger.setLevel(level=logging.DEBUG)
|
||||
|
|
@ -1067,6 +1068,108 @@ async def test_token_counter_httpx_status_error_raises_proxy_exception():
|
|||
litellm.proxy.proxy_server.llm_router = original_router
|
||||
|
||||
|
||||
class _RaisingCounter(BaseTokenCounter):
|
||||
def __init__(self, error: Exception) -> None:
|
||||
self.error = error
|
||||
|
||||
def should_use_token_counting_api(self, custom_llm_provider: str | None = None) -> bool:
|
||||
return True
|
||||
|
||||
async def count_tokens(self, **kwargs) -> TokenCountResponse | None:
|
||||
raise self.error
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"provider_error, expected_code",
|
||||
[
|
||||
(
|
||||
litellm.APIError(
|
||||
status_code=400, message="contents is not specified", llm_provider="gemini", model="gemini-2.5-flash"
|
||||
),
|
||||
"400",
|
||||
),
|
||||
(litellm.APIConnectionError(message="connection refused", llm_provider="gemini", model="gemini-2.5-flash"), "500"),
|
||||
],
|
||||
)
|
||||
async def test_provider_counter_raising_litellm_error_falls_back_or_surfaces_provider_status(
|
||||
provider_error, expected_code, monkeypatch
|
||||
):
|
||||
counter = _RaisingCounter(provider_error)
|
||||
call = dict(
|
||||
provider_counter=counter,
|
||||
custom_llm_provider="gemini",
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
contents=None,
|
||||
deployment={"litellm_params": {"model": "gemini/gemini-2.5-flash"}},
|
||||
request_model="gemini-flash",
|
||||
)
|
||||
|
||||
monkeypatch.setattr(litellm, "disable_token_counter", False)
|
||||
assert await _try_provider_token_count(**call) is None
|
||||
|
||||
monkeypatch.setattr(litellm, "disable_token_counter", True)
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _try_provider_token_count(**call)
|
||||
assert exc_info.value.code == expected_code
|
||||
assert provider_error.message in exc_info.value.message
|
||||
assert exc_info.value.type == "token_counting_error"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gemini_deployment_counts_anthropic_messages_through_provider(monkeypatch):
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.llms.gemini.common_utils import GoogleAIStudioTokenCounter
|
||||
|
||||
seen_bodies: list[dict] = []
|
||||
|
||||
def google_ai_studio(request: httpx.Request) -> httpx.Response:
|
||||
body = json.loads(request.content)
|
||||
seen_bodies.append(body)
|
||||
contents = body.get("contents") or body.get("generateContentRequest", {}).get("contents")
|
||||
if not contents:
|
||||
return httpx.Response(400, json={"error": {"message": "contents is not specified", "status": "INVALID_ARGUMENT"}})
|
||||
return httpx.Response(200, json={"totalTokens": 17})
|
||||
|
||||
counter = GoogleAIStudioTokenCounter(client=AsyncHTTPHandler(transport=httpx.MockTransport(google_ai_studio)))
|
||||
deployment = {"litellm_params": {"model": "gemini/gemini-2.5-flash", "api_key": "test-key"}, "model_info": {}}
|
||||
router = MagicMock()
|
||||
router.async_get_available_deployment = AsyncMock(return_value=deployment)
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router)
|
||||
monkeypatch.setattr(
|
||||
litellm.proxy.proxy_server,
|
||||
"_get_provider_token_counter",
|
||||
lambda deployment, model_to_use: (counter, "gemini-2.5-flash", "gemini"),
|
||||
)
|
||||
|
||||
response = await token_counter(
|
||||
request=TokenCountRequest(
|
||||
model="gemini-flash",
|
||||
messages=[{"role": "user", "content": "What is the weather in Paris?"}],
|
||||
system="Be terse",
|
||||
tools=[{"name": "get_weather", "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}}}],
|
||||
),
|
||||
call_endpoint=True,
|
||||
)
|
||||
|
||||
assert seen_bodies == [
|
||||
{
|
||||
"generateContentRequest": {
|
||||
"model": "models/gemini-2.5-flash",
|
||||
"contents": [{"role": "user", "parts": [{"text": "What is the weather in Paris?"}]}],
|
||||
"systemInstruction": {"parts": [{"text": "Be terse"}]},
|
||||
"tools": [
|
||||
{"function_declarations": [{"name": "get_weather", "parameters": {"type": "object", "properties": {"city": {"type": "string"}}}}]}
|
||||
],
|
||||
}
|
||||
}
|
||||
]
|
||||
assert response.total_tokens == 17
|
||||
assert response.error is False
|
||||
assert response.model_used == "gemini-2.5-flash"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_token_counter_error_raises_exception_when_disabled():
|
||||
"""
|
||||
|
|
|
|||
0
tests/test_litellm/llms/gemini/count_tokens/__init__.py
Normal file
0
tests/test_litellm/llms/gemini/count_tokens/__init__.py
Normal file
|
|
@ -0,0 +1,98 @@
|
|||
from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload
|
||||
|
||||
MODEL = "gemini-2.5-flash"
|
||||
|
||||
|
||||
def test_anthropic_tool_turns_become_gemini_function_call_and_response_parts():
|
||||
payload = build_count_tokens_payload(
|
||||
model=MODEL,
|
||||
messages=[
|
||||
{"role": "user", "content": "What is the weather in Paris?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "tool_use", "id": "toolu_1", "name": "get_weather", "input": {"city": "Paris"}}],
|
||||
},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "Sunny"}]},
|
||||
],
|
||||
system=None,
|
||||
tools=None,
|
||||
)
|
||||
|
||||
assert payload.contents == [
|
||||
{"role": "user", "parts": [{"text": "What is the weather in Paris?"}]},
|
||||
{"role": "model", "parts": [{"function_call": {"name": "get_weather", "args": {"city": "Paris"}}}]},
|
||||
{
|
||||
"role": "user",
|
||||
"parts": [{"function_response": {"name": "get_weather", "response": {"content": "Sunny"}}}],
|
||||
},
|
||||
], payload
|
||||
assert payload.system_instruction is None
|
||||
assert payload.tools is None
|
||||
|
||||
|
||||
def test_system_prompt_is_lifted_out_of_contents_into_system_instruction():
|
||||
payload = build_count_tokens_payload(
|
||||
model=MODEL,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
system=[{"type": "text", "text": "Be terse"}, {"type": "text", "text": "Answer in French"}],
|
||||
tools=None,
|
||||
)
|
||||
|
||||
assert payload.contents == [{"role": "user", "parts": [{"text": "hi"}]}], payload
|
||||
assert payload.system_instruction is not None
|
||||
system_text = "".join(part["text"] for part in payload.system_instruction["parts"])
|
||||
assert "Be terse" in system_text and "Answer in French" in system_text, payload.system_instruction
|
||||
|
||||
|
||||
def test_string_system_prompt_becomes_system_instruction():
|
||||
payload = build_count_tokens_payload(
|
||||
model=MODEL, messages=[{"role": "user", "content": "hi"}], system="You are terse", tools=None
|
||||
)
|
||||
|
||||
assert payload.system_instruction == {"parts": [{"text": "You are terse"}]}
|
||||
assert payload.contents == [{"role": "user", "parts": [{"text": "hi"}]}], payload
|
||||
|
||||
|
||||
def test_anthropic_tools_become_gemini_function_declarations():
|
||||
payload = build_count_tokens_payload(
|
||||
model=MODEL,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
system=None,
|
||||
tools=[
|
||||
{
|
||||
"name": "get_weather",
|
||||
"description": "Weather lookup",
|
||||
"input_schema": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
"required": ["city"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
assert payload.tools == [
|
||||
{
|
||||
"function_declarations": [
|
||||
{
|
||||
"name": "get_weather",
|
||||
"description": "Weather lookup",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
"required": ["city"],
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
], payload.tools
|
||||
|
||||
|
||||
def test_unrecognised_system_value_is_dropped_not_sent():
|
||||
payload = build_count_tokens_payload(
|
||||
model=MODEL, messages=[{"role": "user", "content": "hi"}], system={"unexpected": "shape"}, tools=None
|
||||
)
|
||||
|
||||
assert payload.system_instruction is None
|
||||
assert payload.contents == [{"role": "user", "parts": [{"text": "hi"}]}], payload
|
||||
|
|
@ -1,8 +1,12 @@
|
|||
import json
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.llms.gemini.common_utils import GeminiModelInfo, GoogleAIStudioTokenCounter
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
|
||||
class TestGeminiModelInfo:
|
||||
|
|
@ -158,9 +162,128 @@ class TestGoogleAIStudioTokenCounter:
|
|||
|
||||
# Verify the mock was called correctly
|
||||
mock_acount_tokens.assert_called_once_with(
|
||||
model=model_to_use, contents=contents
|
||||
model=model_to_use, contents=contents, client=None
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _counter_with_upstream(upstream_response: httpx.Response) -> tuple[GoogleAIStudioTokenCounter, list[httpx.Request]]:
|
||||
seen_requests: list[httpx.Request] = []
|
||||
|
||||
def upstream(request: httpx.Request) -> httpx.Response:
|
||||
seen_requests.append(request)
|
||||
return upstream_response
|
||||
|
||||
client = AsyncHTTPHandler(transport=httpx.MockTransport(upstream))
|
||||
return GoogleAIStudioTokenCounter(client=client), seen_requests
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_messages_are_sent_as_gemini_contents_with_system_and_tools(self):
|
||||
counter, seen = self._counter_with_upstream(httpx.Response(200, json={"totalTokens": 42}))
|
||||
|
||||
result = await counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "What is the weather in Paris?"}],
|
||||
contents=None,
|
||||
deployment={"litellm_params": {"model": "gemini/gemini-2.5-flash", "api_key": "test-key"}},
|
||||
request_model="gemini-flash",
|
||||
tools=[{"name": "get_weather", "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}}}],
|
||||
system="Be terse",
|
||||
)
|
||||
|
||||
assert len(seen) == 1, seen
|
||||
assert seen[0].url.path == "/v1beta/models/gemini-2.5-flash:countTokens"
|
||||
assert seen[0].headers["x-goog-api-key"] == "test-key"
|
||||
assert json.loads(seen[0].content) == {
|
||||
"generateContentRequest": {
|
||||
"model": "models/gemini-2.5-flash",
|
||||
"contents": [{"role": "user", "parts": [{"text": "What is the weather in Paris?"}]}],
|
||||
"systemInstruction": {"parts": [{"text": "Be terse"}]},
|
||||
"tools": [
|
||||
{"function_declarations": [{"name": "get_weather", "parameters": {"type": "object", "properties": {"city": {"type": "string"}}}}]}
|
||||
],
|
||||
}
|
||||
}
|
||||
assert result == TokenCountResponse(
|
||||
total_tokens=42,
|
||||
request_model="gemini-flash",
|
||||
model_used="gemini-2.5-flash",
|
||||
tokenizer_type="",
|
||||
original_response={"totalTokens": 42},
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_contents_are_sent_unchanged(self):
|
||||
counter, seen = self._counter_with_upstream(httpx.Response(200, json={"totalTokens": 3}))
|
||||
contents = [{"role": "user", "parts": [{"text": "Hello world"}]}]
|
||||
|
||||
result = await counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=None,
|
||||
contents=contents,
|
||||
deployment={"litellm_params": {"api_key": "test-key"}},
|
||||
request_model="gemini-flash",
|
||||
)
|
||||
|
||||
assert json.loads(seen[0].content) == {"contents": contents}
|
||||
assert result is not None and result.total_tokens == 3
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_rejection_is_returned_as_error_value_not_raised(self):
|
||||
counter, _ = self._counter_with_upstream(
|
||||
httpx.Response(400, json={"error": {"message": "contents is not specified"}})
|
||||
)
|
||||
|
||||
result = await counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
contents=None,
|
||||
deployment={"litellm_params": {"api_key": "test-key"}},
|
||||
request_model="gemini-flash",
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.error is True
|
||||
assert result.status_code == 400
|
||||
assert result.error_message is not None and "contents is not specified" in result.error_message
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_success_without_total_tokens_is_an_error_value_not_zero(self):
|
||||
counter, _ = self._counter_with_upstream(httpx.Response(200, json={"promptTokensDetails": []}))
|
||||
|
||||
result = await counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
contents=None,
|
||||
deployment={"litellm_params": {"api_key": "test-key"}},
|
||||
request_model="gemini-flash",
|
||||
)
|
||||
|
||||
assert result == TokenCountResponse(
|
||||
total_tokens=0,
|
||||
request_model="gemini-flash",
|
||||
model_used="gemini-2.5-flash",
|
||||
tokenizer_type="gemini_api",
|
||||
original_response={"promptTokensDetails": []},
|
||||
error=True,
|
||||
error_message="Google Gen AI Studio countTokens response has no totalTokens",
|
||||
status_code=502,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_nothing_to_count_skips_the_provider_call(self):
|
||||
counter, seen = self._counter_with_upstream(httpx.Response(200, json={"totalTokens": 0}))
|
||||
|
||||
result = await counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=[],
|
||||
contents=None,
|
||||
deployment={"litellm_params": {"api_key": "test-key"}},
|
||||
request_model="gemini-flash",
|
||||
)
|
||||
|
||||
assert result is None
|
||||
assert seen == []
|
||||
|
||||
def test_clean_contents_for_gemini_api_removes_id_field(self):
|
||||
"""Test that _clean_contents_for_gemini_api removes unsupported 'id' field from function responses"""
|
||||
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue