fix(gemini): validate Anthropic token-count input with TypeAdapter and type the countTokens request body

This commit is contained in:
shrey kharbanda 2026-09-24 16:36:19 +00:00
parent 629a068e10
commit 78a8a4afe9
7 changed files with 271 additions and 118 deletions

View file

@ -6,6 +6,7 @@ from collections.abc import Mapping, Sequence
from typing import Any, Final
import httpx
from pydantic import TypeAdapter
import litellm
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
@ -13,6 +14,7 @@ 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.gemini import GeminiCountTokensDeploymentParams
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import TokenCountResponse
@ -472,6 +474,10 @@ def get_api_key_from_env() -> str | None:
return get_secret_str("GOOGLE_API_KEY") or get_secret_str("GEMINI_API_KEY")
_COUNT_TOKENS_DEPLOYMENT_PARAMS: Final = TypeAdapter(GeminiCountTokensDeploymentParams)
_NO_DEPLOYMENT_PARAMS: Final[GeminiCountTokensDeploymentParams] = {}
class GoogleAIStudioTokenCounter(BaseTokenCounter):
"""Token counter implementation for Google AI Studio provider."""
@ -496,31 +502,45 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
tools: list[dict[str, object]] | None = None,
system: object | None = None,
) -> TokenCountResponse | None:
import copy
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload
from litellm.llms.gemini.count_tokens.transformation import (
AnthropicCountTokensInput,
InvalidAnthropicRequest,
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 = {
anthropic_input: Final[AnthropicCountTokensInput] = {
"model": model_to_use,
"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 {}),
"messages": messages or (),
"system": system,
"tools": tools,
}
count_tokens_params_request.update(count_tokens_params)
payload: Final = build_count_tokens_payload(anthropic_input) if contents is None else None
if isinstance(payload, InvalidAnthropicRequest):
return TokenCountResponse(
total_tokens=0,
request_model=request_model,
model_used=model_to_use,
tokenizer_type="gemini_api",
error=True,
error_message=payload.message,
status_code=400,
)
deployment_params: Final = (
_COUNT_TOKENS_DEPLOYMENT_PARAMS.validate_python(deployment["litellm_params"])
if deployment and "litellm_params" in deployment
else _NO_DEPLOYMENT_PARAMS
)
try:
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
**count_tokens_params_request,
model=model_to_use,
api_key=deployment_params.get("api_key") or deployment_params.get("gemini_api_key"),
api_base=deployment_params.get("api_base"),
contents=contents if payload is None else payload.contents,
system_instruction=None if payload is None else payload.system_instruction,
tools=None if payload is None else payload.tools,
client=self.client,
)
except (litellm.APIError, litellm.APIConnectionError) as e:

View file

@ -1,16 +1,48 @@
from typing import TYPE_CHECKING, Any, Final
from collections.abc import Sequence
from typing import Any, Final
import httpx
import litellm
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.llms.gemini import GeminiCountTokensRequest
from litellm.types.llms.vertex_ai import ContentType, SystemInstructions, Tools
from litellm.types.utils import LlmProviders
if TYPE_CHECKING:
from litellm.types.google_genai.main import GenerateContentContentListUnionDict
else:
GenerateContentContentListUnionDict = Any
def build_count_tokens_request(
model: str,
contents: Sequence[ContentType],
system_instruction: SystemInstructions | None,
tools: Sequence[Tools] | None,
) -> GeminiCountTokensRequest:
model_name: Final = f"models/{model}"
if tools is None:
if system_instruction is None:
bare: Final[GeminiCountTokensRequest] = {"contents": contents}
return bare
with_system: Final[GeminiCountTokensRequest] = {
"generateContentRequest": {
"model": model_name,
"contents": contents,
"systemInstruction": system_instruction,
}
}
return with_system
if system_instruction is None:
with_tools: Final[GeminiCountTokensRequest] = {
"generateContentRequest": {"model": model_name, "contents": contents, "tools": tools}
}
return with_tools
with_both: Final[GeminiCountTokensRequest] = {
"generateContentRequest": {
"model": model_name,
"contents": contents,
"systemInstruction": system_instruction,
"tools": tools,
}
}
return with_both
class GoogleAIStudioTokenCounter:
@ -86,7 +118,7 @@ class GoogleAIStudioTokenCounter:
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
system_instruction: SystemInstructions | None = None,
tools: list[Tools] | None = None,
tools: Sequence[Tools] | None = None,
client: AsyncHTTPHandler | None = None,
**kwargs: object,
) -> dict[str, Any]:
@ -135,18 +167,11 @@ 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}
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 {}),
}
}
request_body: Final = build_count_tokens_request(
model=model,
contents=self._clean_contents_for_gemini_api(contents),
system_instruction=system_instruction,
tools=tools,
)
async_httpx_client: Final = client or get_async_httpx_client(
@ -154,7 +179,11 @@ class GoogleAIStudioTokenCounter:
)
try:
response: Final = await async_httpx_client.post(url=url, headers=headers, json=request_body)
response: Final = await async_httpx_client.post(
url=url,
headers=headers,
json=request_body, # pyright: ignore[reportArgumentType] # post() takes a bare dict; a TypedDict is one at runtime
)
except httpx.HTTPStatusError as e:
error_msg = f"Google Gen AI Studio API error: {e.response.status_code} - {e.response.text}"
raise litellm.APIError(

View file

@ -1,5 +1,9 @@
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Any, Final
from typing import Final, cast
from pydantic import TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
LiteLLMAnthropicMessagesAdapter,
@ -13,41 +17,57 @@ from litellm.types.llms.anthropic import AnthropicMessagesRequest
from litellm.types.llms.vertex_ai import ContentType, SystemInstructions, Tools
class AnthropicCountTokensInput(TypedDict):
"""The Anthropic-shaped fields of a token-count request before validation."""
model: ReadOnly[str]
messages: ReadOnly[Sequence[Mapping[str, object]]]
system: ReadOnly[object]
tools: ReadOnly[Sequence[Mapping[str, object]] | None]
@dataclass(frozen=True, slots=True)
class GeminiCountTokensPayload:
contents: list[ContentType]
contents: Sequence[ContentType]
system_instruction: SystemInstructions | None
tools: list[Tools] | None
tools: Sequence[Tools] | None
def build_count_tokens_payload(
model: str,
messages: list[dict[str, Any]],
system: object | None,
tools: list[dict[str, Any]] | None,
) -> GeminiCountTokensPayload:
@dataclass(frozen=True, slots=True)
class InvalidAnthropicRequest:
message: str
_ANTHROPIC_REQUEST: Final = TypeAdapter(AnthropicMessagesRequest)
def _validated_request(raw: AnthropicCountTokensInput) -> AnthropicMessagesRequest | InvalidAnthropicRequest:
try:
_ANTHROPIC_REQUEST.validate_python(raw)
except ValidationError as e:
return InvalidAnthropicRequest(message=str(e))
return cast(AnthropicMessagesRequest, raw) # cast-ok: validated above; pydantic returns lazy Iterable validators
def build_count_tokens_payload(raw: AnthropicCountTokensInput) -> GeminiCountTokensPayload | InvalidAnthropicRequest:
"""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 {}),
)
request: Final = _validated_request(raw)
if isinstance(request, InvalidAnthropicRequest):
return request
openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
anthropic_request, custom_llm_provider="gemini"
request, custom_llm_provider="gemini"
)
system_instruction, remaining_messages = _transform_system_message(
supports_system_message=True,
messages=list(openai_request["messages"]),
supports_system_message=True, messages=openai_request["messages"]
)
contents: Final = _gemini_convert_messages_with_history(
messages=remaining_messages, model=model, custom_llm_provider="gemini"
messages=remaining_messages, model=raw["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={},
value=[dict(tool) for tool in openai_tools], # mutable-ok: _map_function only accepts list[dict]
optional_params={}, # mutable-ok: _map_function writes retrieval config into the dict it is given
)
if openai_tools
else None

View file

@ -1,17 +1,42 @@
from collections.abc import Sequence
from enum import Enum
from typing import Any, Literal
from typing_extensions import Required, TypedDict
from typing_extensions import NotRequired, ReadOnly, Required, TypedDict
from .vertex_ai import (
ContentType,
GenerationConfig,
HttpxBlobType,
HttpxContentType,
SystemInstructions,
Tools,
UsageMetadata,
)
class GeminiGenerateContentRequest(TypedDict):
model: ReadOnly[str]
contents: ReadOnly[Sequence[ContentType]]
systemInstruction: ReadOnly[NotRequired[SystemInstructions]]
tools: ReadOnly[NotRequired[Sequence[Tools]]]
class GeminiCountTokensDeploymentParams(TypedDict, total=False):
"""The deployment litellm_params the countTokens handler reads: everything else is ignored."""
api_key: ReadOnly[str | None]
gemini_api_key: ReadOnly[str | None]
api_base: ReadOnly[str | None]
class GeminiCountTokensRequest(TypedDict, total=False):
"""Body of models/{model}:countTokens: bare contents, or a generateContentRequest when system or tools are set."""
contents: ReadOnly[Sequence[ContentType]]
generateContentRequest: ReadOnly[GeminiGenerateContentRequest]
class GeminiFilesState(Enum):
STATE_UNSPECIFIED = "STATE_UNSPECIFIED"
PROCESSING = "PROCESSING"

View file

@ -22,6 +22,7 @@ from fastapi import HTTPException, Request
import litellm
from litellm import Router
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.base_utils import BaseTokenCounter
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter
from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler
@ -29,7 +30,6 @@ from litellm.proxy._types import ProxyException, TokenCountRequest
from litellm.proxy.anthropic_endpoints.endpoints import (
count_tokens as anthropic_count_tokens,
)
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

View file

@ -1,21 +1,28 @@
from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload
from litellm.llms.gemini.count_tokens.transformation import (
GeminiCountTokensPayload,
InvalidAnthropicRequest,
build_count_tokens_payload,
)
MODEL = "gemini-2.5-flash"
def _payload(messages, system=None, tools=None) -> GeminiCountTokensPayload:
payload = build_count_tokens_payload({"model": MODEL, "messages": messages, "system": system, "tools": tools})
assert isinstance(payload, GeminiCountTokensPayload), payload
return payload
def test_anthropic_tool_turns_become_gemini_function_call_and_response_parts():
payload = build_count_tokens_payload(
model=MODEL,
messages=[
payload = _payload(
[
{"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 == [
@ -30,12 +37,37 @@ def test_anthropic_tool_turns_become_gemini_function_call_and_response_parts():
assert payload.tools is None
def test_tool_result_with_block_list_content_keeps_its_text():
payload = _payload(
[
{"role": "user", "content": "Weather?"},
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "toolu_1", "name": "get_weather", "input": {}}],
},
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "toolu_1",
"content": [{"type": "text", "text": "Sunny"}, {"type": "text", "text": "21C"}],
}
],
},
]
)
assert payload.contents[2] == {
"role": "user",
"parts": [{"function_response": {"name": "get_weather", "response": {"content": "Sunny21C"}}}],
}, payload.contents
def test_system_prompt_is_lifted_out_of_contents_into_system_instruction():
payload = build_count_tokens_payload(
model=MODEL,
messages=[{"role": "user", "content": "hi"}],
payload = _payload(
[{"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
@ -45,19 +77,15 @@ def test_system_prompt_is_lifted_out_of_contents_into_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
)
payload = _payload([{"role": "user", "content": "hi"}], system="You are terse")
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,
payload = _payload(
[{"role": "user", "content": "hi"}],
tools=[
{
"name": "get_weather",
@ -89,10 +117,24 @@ def test_anthropic_tools_become_gemini_function_declarations():
], 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
def test_system_of_an_unrecognised_shape_is_rejected_as_invalid_not_sent():
result = 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
assert isinstance(result, InvalidAnthropicRequest), result
assert "system" in result.message, result.message
def test_tools_that_are_not_a_list_are_rejected_as_invalid():
result = build_count_tokens_payload(
{"model": MODEL, "messages": [{"role": "user", "content": "hi"}], "system": None, "tools": {"name": "one"}}
)
assert isinstance(result, InvalidAnthropicRequest), result
assert "tools" in result.message, result.message

View file

@ -40,18 +40,10 @@ class TestGeminiModelInfo:
# Test edge cases where model names end with characters from "models/"
# These would be incorrectly processed if using strip("models/") instead of replace("models/", "")
models = [
{
"name": "models/gemini-1.5-pro"
}, # ends with 'o' - would become "gemini-1.5-pr" with strip()
{
"name": "models/test-model"
}, # ends with 'l' - would become "gemini/test-mode" with strip()
{
"name": "models/custom-models"
}, # ends with 's' - would become "gemini/custom-model" with strip()
{
"name": "models/demo"
}, # ends with 'o' - would become "gemini/dem" with strip()
{"name": "models/gemini-1.5-pro"}, # ends with 'o' - would become "gemini-1.5-pr" with strip()
{"name": "models/test-model"}, # ends with 'l' - would become "gemini/test-mode" with strip()
{"name": "models/custom-models"}, # ends with 's' - would become "gemini/custom-model" with strip()
{"name": "models/demo"}, # ends with 'o' - would become "gemini/dem" with strip()
]
result = gemini_model_info.process_model_name(models)
@ -102,16 +94,10 @@ class TestGoogleAIStudioTokenCounter:
token_counter = GoogleAIStudioTokenCounter()
# Test with gemini provider - should return True
assert (
token_counter.should_use_token_counting_api(LlmProviders.GEMINI.value)
is True
)
assert token_counter.should_use_token_counting_api(LlmProviders.GEMINI.value) is True
# Test with other providers - should return False
assert (
token_counter.should_use_token_counting_api(LlmProviders.OPENAI.value)
is False
)
assert token_counter.should_use_token_counting_api(LlmProviders.OPENAI.value) is False
assert token_counter.should_use_token_counting_api("anthropic") is False
assert token_counter.should_use_token_counting_api("vertex_ai") is False
@ -162,11 +148,19 @@ class TestGoogleAIStudioTokenCounter:
# Verify the mock was called correctly
mock_acount_tokens.assert_called_once_with(
model=model_to_use, contents=contents, client=None
model=model_to_use,
api_key=None,
api_base=None,
contents=contents,
system_instruction=None,
tools=None,
client=None,
)
@staticmethod
def _counter_with_upstream(upstream_response: httpx.Response) -> tuple[GoogleAIStudioTokenCounter, list[httpx.Request]]:
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:
@ -186,7 +180,9 @@ class TestGoogleAIStudioTokenCounter:
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"}}}}],
tools=[
{"name": "get_weather", "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}}}
],
system="Be terse",
)
@ -199,7 +195,14 @@ class TestGoogleAIStudioTokenCounter:
"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"}}}}]}
{
"function_declarations": [
{
"name": "get_weather",
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}},
}
]
}
],
}
}
@ -211,6 +214,27 @@ class TestGoogleAIStudioTokenCounter:
original_response={"totalTokens": 42},
)
@pytest.mark.asyncio
async def test_system_without_tools_still_wraps_in_generate_content_request(self):
counter, seen = self._counter_with_upstream(httpx.Response(200, json={"totalTokens": 9}))
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",
system=[{"type": "text", "text": "Be terse"}],
)
assert json.loads(seen[0].content) == {
"generateContentRequest": {
"model": "models/gemini-2.5-flash",
"contents": [{"role": "user", "parts": [{"text": "hi"}]}],
"systemInstruction": {"parts": [{"text": "Be terse"}]},
}
}
@pytest.mark.asyncio
async def test_native_contents_are_sent_unchanged(self):
counter, seen = self._counter_with_upstream(httpx.Response(200, json={"totalTokens": 3}))
@ -299,9 +323,7 @@ class TestGoogleAIStudioTokenCounter:
"functionResponse": {
"id": "read_many_files-1757526647518-730a691aac11c", # This should be removed
"name": "read_many_files",
"response": {
"output": "No files matching the criteria were found or all were skipped."
},
"response": {"output": "No files matching the criteria were found or all were skipped."},
}
}
],
@ -310,9 +332,7 @@ class TestGoogleAIStudioTokenCounter:
]
# Clean the contents
cleaned_contents = token_counter._clean_contents_for_gemini_api(
contents_with_id
)
cleaned_contents = token_counter._clean_contents_for_gemini_api(contents_with_id)
# Verify the 'id' field was removed
function_response = cleaned_contents[1]["parts"][0]["functionResponse"]
@ -321,8 +341,7 @@ class TestGoogleAIStudioTokenCounter:
assert "response" in function_response
assert function_response["name"] == "read_many_files"
assert (
function_response["response"]["output"]
== "No files matching the criteria were found or all were skipped."
function_response["response"]["output"] == "No files matching the criteria were found or all were skipped."
)
def test_clean_contents_for_gemini_api_preserves_other_fields(self):
@ -338,9 +357,7 @@ class TestGoogleAIStudioTokenCounter:
]
# Clean the contents
cleaned_contents = token_counter._clean_contents_for_gemini_api(
contents_without_function_response
)
cleaned_contents = token_counter._clean_contents_for_gemini_api(contents_without_function_response)
# Verify the contents are unchanged
assert cleaned_contents == contents_without_function_response