fix(gemini): count Anthropic messages, system and tools on /v1/messages/count_tokens and fall back on provider errors

This commit is contained in:
shrey kharbanda 2026-09-24 15:55:37 +00:00
parent cfa2830bde
commit 629a068e10
9 changed files with 481 additions and 29 deletions

View file

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

View file

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

View file

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

View 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)

View file

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

View file

@ -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():
"""

View 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

View file

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