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

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-23 12:47:12 +00:00
parent 860bc7811d
commit d9b612e1e2
9 changed files with 416 additions and 6 deletions

View file

@ -495,17 +495,43 @@ 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
deployment = deployment or {}
count_tokens_params_request: Final = copy.deepcopy(deployment.get("litellm_params", {}))
payload: Final = (
build_count_tokens_payload(model=model_to_use, messages=messages, system=system, tools=tools)
if contents is None
else None
)
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 is not None and payload.system_instruction is not None
else {}
),
**({"tools": payload.tools} if payload is not None and payload.tools is not None else {}),
}
count_tokens_params_request.update(count_tokens_params)
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
**count_tokens_params_request,
)
try:
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
**count_tokens_params_request,
)
except (litellm.APIError, litellm.APIConnectionError) as e:
return TokenCountResponse(
total_tokens=0,
request_model=request_model,
model_used=model_to_use,
tokenizer_type="gemini_api",
error=True,
error_message=e.message,
status_code=e.status_code,
)
if result is not None:
return TokenCountResponse(

View file

@ -4,6 +4,8 @@ import httpx
import litellm
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.vertex_ai import SystemInstructions, Tools
from litellm.types.utils import LlmProviders
if TYPE_CHECKING:
@ -51,7 +53,7 @@ class GoogleAIStudioTokenCounter:
"""
Construct the URL for the Google Gen AI Studio countTokens endpoint.
"""
base_url: Final = api_base or "https://generativelanguage.googleapis.com"
base_url: Final = api_base or get_secret_str("GEMINI_API_BASE") or "https://generativelanguage.googleapis.com"
return f"{base_url}/v1beta/models/{model}:countTokens"
async def validate_environment(
@ -84,6 +86,8 @@ 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,
**kwargs: object,
) -> dict[str, Any]:
"""
@ -130,7 +134,18 @@ 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(
llm_provider=LlmProviders.GEMINI,

View file

@ -0,0 +1,61 @@
"""Translate an Anthropic /v1/messages/count_tokens request into a Gemini
countTokens payload (contents + systemInstruction + tools)."""
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Final, cast
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,
)
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: Sequence[Mapping[str, object]],
system: object | None,
tools: Sequence[Mapping[str, object]] | None,
) -> GeminiCountTokensPayload:
anthropic_request: Final[AnthropicMessagesRequest] = cast(
AnthropicMessagesRequest, # cast-ok: untrusted client payload, adapter reads the anthropic-shape keys only
{
"model": model,
"messages": list(messages),
**({"system": system} if system else {}),
**({"tools": list(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")
return GeminiCountTokensPayload(
contents=contents,
system_instruction=system_instruction,
tools=VertexGeminiConfig()._map_function(value=[dict(tool) for tool in openai_tools], optional_params={})
if openai_tools
else None,
)

View file

@ -0,0 +1,93 @@
"""Live e2e: POST /v1/messages/count_tokens against a gemini/ deployment.
Registers a fresh gemini/gemini-2.5-flash deployment via /model/new (deleted on
teardown) and drives the endpoint through the shared transport, since no
official provider SDK covers this route.
"""
import pytest
from e2e_config import unique_marker
from e2e_http import unwrap
from lifecycle import ResourceManager
from models import (
AnthropicCustomTool,
ChatMessage,
CountTokensBody,
JsonSchemaProperty,
LiteLLMParamsBody,
ToolInputSchema,
)
from proxy_client import ProxyClient
pytestmark = pytest.mark.e2e
BACKEND_MODEL = "gemini/gemini-2.5-flash"
GEMINI_API_KEY = "os.environ/GEMINI_API_KEY"
WEATHER_TOOL = AnthropicCustomTool(
name="get_weather",
description="Get the current weather for a city.",
input_schema=ToolInputSchema(
properties={"city": JsonSchemaProperty(type="string")},
required=["city"],
),
)
def _provision(proxy: ProxyClient, resources: ResourceManager) -> str:
model_name = f"e2e-count-tokens-gemini-{unique_marker()}"
model_id = proxy.create_model(
model_name,
LiteLLMParamsBody(model=BACKEND_MODEL, api_key=GEMINI_API_KEY),
)
resources.defer(lambda: proxy.delete_model(model_id))
return model_name
class TestMessagesCountTokens:
def test_count_tokens_gemini_returns_input_tokens(
self, proxy: ProxyClient, resources: ResourceManager, scoped_key: str
) -> None:
model = _provision(proxy, resources)
response = unwrap(
proxy.count_tokens(
scoped_key,
CountTokensBody(
model=model,
messages=[ChatMessage(role="user", content=f"hello world {unique_marker()}")],
),
)
)
assert response.input_tokens > 0, f"input_tokens not positive: {response.input_tokens}"
def test_count_tokens_gemini_with_system_and_tools(
self, proxy: ProxyClient, resources: ResourceManager, scoped_key: str
) -> None:
model = _provision(proxy, resources)
prompt = f"hello world {unique_marker()}"
plain = unwrap(
proxy.count_tokens(
scoped_key,
CountTokensBody(
model=model,
messages=[ChatMessage(role="user", content=prompt)],
),
)
)
with_system_and_tools = unwrap(
proxy.count_tokens(
scoped_key,
CountTokensBody(
model=model,
messages=[ChatMessage(role="user", content=prompt)],
system="You are a helpful assistant",
tools=[WEATHER_TOOL],
),
)
)
assert with_system_and_tools.input_tokens > plain.input_tokens, (
f"system + tools count {with_system_and_tools.input_tokens} did not exceed "
f"the plain message count {plain.input_tokens}; the extra prompt was not counted"
)

View file

@ -574,6 +574,8 @@ class CountTokensBody(BaseModel):
model: str
messages: list[ChatMessage]
system: str | None = None
tools: list[AnthropicTool] | None = None
class AnthropicMessagesResponse(BaseModel):

View file

@ -0,0 +1,64 @@
import json
import httpx
import pytest
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
COUNT_TOKENS_URL = "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:countTokens"
@pytest.mark.asyncio
async def test_acount_tokens_sends_generate_content_request_when_system_or_tools_present(monkeypatch):
recorded: list[httpx.Request] = []
def _handler(request: httpx.Request) -> httpx.Response:
recorded.append(request)
return httpx.Response(200, json={"totalTokens": 42})
monkeypatch.setattr(
"litellm.llms.gemini.count_tokens.handler.get_async_httpx_client",
lambda **kwargs: httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
)
result = await GoogleAIStudioTokenCounter().acount_tokens(
model="gemini-2.5-flash",
contents=[{"role": "user", "parts": [{"text": "hello world"}]}],
api_key="test-key",
system_instruction={"parts": [{"text": "You are a helpful assistant"}]},
tools=[{"function_declarations": [{"name": "get_weather"}]}],
)
assert result == {"totalTokens": 42}
request = recorded[-1]
assert request.url == COUNT_TOKENS_URL
body = json.loads(request.content)
assert "contents" not in body
generate_content_request = body["generateContentRequest"]
assert generate_content_request["model"] == "models/gemini-2.5-flash"
assert generate_content_request["contents"] == [{"role": "user", "parts": [{"text": "hello world"}]}]
assert generate_content_request["systemInstruction"] == {"parts": [{"text": "You are a helpful assistant"}]}
assert generate_content_request["tools"][0]["function_declarations"][0]["name"] == "get_weather"
@pytest.mark.asyncio
async def test_acount_tokens_keeps_contents_body_without_system_or_tools(monkeypatch):
recorded: list[httpx.Request] = []
def _handler(request: httpx.Request) -> httpx.Response:
recorded.append(request)
return httpx.Response(200, json={"totalTokens": 4})
monkeypatch.setattr(
"litellm.llms.gemini.count_tokens.handler.get_async_httpx_client",
lambda **kwargs: httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
)
await GoogleAIStudioTokenCounter().acount_tokens(
model="gemini-2.5-flash",
contents=[{"role": "user", "parts": [{"text": "hi"}]}],
api_key="test-key",
)
body = json.loads(recorded[-1].content)
assert body == {"contents": [{"role": "user", "parts": [{"text": "hi"}]}]}

View file

@ -0,0 +1,62 @@
from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload
def test_build_count_tokens_payload_translates_anthropic_request():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[{"role": "user", "content": "hello world"}],
system="You are a helpful assistant",
tools=[
{
"name": "get_weather",
"description": "Get the current weather for a city.",
"input_schema": {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
}
],
)
assert payload.contents
assert payload.contents[0]["parts"][0].get("text") == "hello world"
assert payload.system_instruction is not None
assert payload.system_instruction["parts"][0].get("text") == "You are a helpful assistant"
assert payload.tools is not None
function_declarations = payload.tools[0]["function_declarations"]
assert function_declarations[0]["name"] == "get_weather"
assert function_declarations[0].get("parameters", {}).get("required") == ["city"]
def test_build_count_tokens_payload_without_system_or_tools():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[{"role": "user", "content": "hi"}],
system=None,
tools=None,
)
assert payload.contents
assert payload.system_instruction is None
assert payload.tools is None
def test_build_count_tokens_payload_passes_openai_tools_through():
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[{"role": "user", "content": "hi"}],
system=None,
tools=[
{
"type": "function",
"function": {
"name": "get_weather",
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}},
},
}
],
)
assert payload.tools is not None
assert payload.tools[0]["function_declarations"][0]["name"] == "get_weather"

View file

@ -161,6 +161,93 @@ class TestGoogleAIStudioTokenCounter:
model=model_to_use, contents=contents
)
@pytest.mark.asyncio
async def test_count_tokens_translates_anthropic_messages_system_and_tools(self):
"""Anthropic-format messages are converted to gemini contents/system/tools
before hitting the countTokens endpoint."""
token_counter = GoogleAIStudioTokenCounter()
with patch(
"litellm.llms.gemini.count_tokens.handler.GoogleAIStudioTokenCounter.acount_tokens",
new_callable=AsyncMock,
) as mock_acount_tokens:
mock_acount_tokens.return_value = {"totalTokens": 12}
result = await token_counter.count_tokens(
model_to_use="gemini-2.5-flash",
messages=[{"role": "user", "content": "hello world"}],
contents=None,
deployment=None,
request_model="gemini/gemini-2.5-flash",
tools=[
{
"name": "get_weather",
"description": "Get the current weather for a city.",
"input_schema": {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
}
],
system="You are a helpful assistant",
)
assert result is not None
assert result.total_tokens == 12
kwargs = mock_acount_tokens.call_args.kwargs
assert kwargs["contents"]
assert kwargs["contents"][0]["parts"][0].get("text") == "hello world"
assert kwargs["system_instruction"]["parts"][0].get("text") == "You are a helpful assistant"
assert kwargs["tools"][0]["function_declarations"][0]["name"] == "get_weather"
@pytest.mark.asyncio
async def test_count_tokens_provider_error_returns_error_response(self):
"""A provider APIError must surface as an error TokenCountResponse so the
proxy falls back to the local tokenizer instead of 500ing."""
import litellm
token_counter = GoogleAIStudioTokenCounter()
with patch(
"litellm.llms.gemini.count_tokens.handler.GoogleAIStudioTokenCounter.acount_tokens",
new_callable=AsyncMock,
) as mock_acount_tokens:
mock_acount_tokens.side_effect = litellm.APIError(
status_code=400,
message="Google Gen AI Studio API error: 400",
llm_provider="gemini",
model="gemini-2.5-flash",
)
result = await token_counter.count_tokens(
model_to_use="gemini-2.5-flash",
messages=[{"role": "user", "content": "hello world"}],
contents=None,
deployment=None,
request_model="gemini/gemini-2.5-flash",
)
assert result is not None
assert result.error is True
assert result.status_code == 400
assert result.total_tokens == 0
assert result.error_message is not None
@pytest.mark.asyncio
async def test_count_tokens_returns_none_without_contents_or_messages(self):
token_counter = GoogleAIStudioTokenCounter()
result = await token_counter.count_tokens(
model_to_use="gemini-2.5-flash",
messages=None,
contents=None,
deployment=None,
request_model="gemini/gemini-2.5-flash",
)
assert result is None
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