fix(gemini): typed countTokens request, 400 on malformed input, 502 on missing totalTokens, proxy fallback on provider exceptions

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-24 23:37:16 +00:00
parent 0a25c92f2b
commit e3713c767f
9 changed files with 466 additions and 215 deletions

View file

@ -3,7 +3,7 @@ import datetime
import json
import math
from collections.abc import Mapping, Sequence
from typing import Any, Final
from typing import Any, Final, cast
import httpx
@ -497,68 +497,83 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
from litellm.llms.gemini.count_tokens.transformation import (
InvalidCountTokensRequest,
build_count_tokens_payload,
normalize_count_tokens_tools,
)
from litellm.types.llms.vertex_ai import SystemInstructions
if contents is None and not messages:
return None
deployment = deployment or {}
count_tokens_params_request: Final = copy.deepcopy(deployment.get("litellm_params", {}))
try:
payload: Final = (
build_count_tokens_payload(model=model_to_use, messages=messages, system=system, tools=tools)
if contents is None
else None
)
system_instruction: Final = (
payload.system_instruction
if payload is not None
else (
{"parts": [{"text": system}]} # mutable-ok: SystemInstructions wire shape
if isinstance(system, str)
else system
)
)
gemini_tools: Final = payload.tools if payload is not None else normalize_count_tokens_tools(tools)
count_tokens_params: Final = { # mutable-ok: kwargs dict for acount_tokens
"model": model_to_use,
"contents": payload.contents if payload is not None else contents,
**(
{"system_instruction": system_instruction} # mutable-ok: kwargs dict for acount_tokens
if system_instruction is not None
else {} # mutable-ok: kwargs dict for acount_tokens
),
**(
{"tools": gemini_tools} # mutable-ok: kwargs dict for acount_tokens
if gemini_tools is not None
else {} # mutable-ok: kwargs dict for acount_tokens
),
}
count_tokens_params_request.update(count_tokens_params)
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
client=client,
**count_tokens_params_request,
)
if result is not None:
return TokenCountResponse(
total_tokens=result.get("totalTokens", 0),
request_model=request_model,
model_used=model_to_use,
tokenizer_type="gemini_api",
original_response=result,
)
return None
except Exception as e:
# provider counting is best-effort: translation, credential, and request
# failures all degrade to the proxy's local-tokenizer fallback
payload: Final = (
build_count_tokens_payload(model=model_to_use, messages=messages, system=system, tools=tools)
if contents is None
else None
)
if isinstance(payload, InvalidCountTokensRequest):
return TokenCountResponse(
total_tokens=0,
request_model=request_model,
model_used=model_to_use,
tokenizer_type="gemini_api",
error=True,
error_message=getattr(e, "message", None) or str(e),
status_code=getattr(e, "status_code", None) or 500,
error_message=payload.message,
status_code=400,
)
system_instruction: Final[SystemInstructions | None] = (
payload.system_instruction
if payload is not None
else (
{"parts": [{"text": system}]} # mutable-ok: SystemInstructions wire shape
if isinstance(system, str)
else cast( # cast-ok: contents-path callers pass a Gemini-shaped systemInstruction
"SystemInstructions | None",
system,
)
)
)
gemini_tools: Final = payload.tools if payload is not None else normalize_count_tokens_tools(tools)
count_tokens_params_request.update(
{ # mutable-ok: kwargs dict for acount_tokens
"model": model_to_use,
"contents": payload.contents if payload is not None else contents,
}
)
try:
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
system_instruction=system_instruction,
tools=gemini_tools,
client=client,
**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 "totalTokens" not in result:
return TokenCountResponse(
total_tokens=0,
request_model=request_model,
model_used=model_to_use,
tokenizer_type="gemini_api",
error=True,
error_message="Google Gen AI Studio countTokens response has no totalTokens",
status_code=502,
original_response=result,
)
return TokenCountResponse(
total_tokens=result["totalTokens"],
request_model=request_model,
model_used=model_to_use,
tokenizer_type="gemini_api",
original_response=result,
)

View file

@ -1,3 +1,4 @@
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any, Final
import httpx
@ -5,7 +6,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.llms.gemini import GeminiCountTokensRequest
from litellm.types.llms.vertex_ai import ContentType, SystemInstructions, Tools
from litellm.types.utils import LlmProviders
if TYPE_CHECKING:
@ -14,6 +16,41 @@ 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:
def _clean_contents_for_gemini_api(self, contents: Any) -> Any:
"""
@ -87,7 +124,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: httpx.AsyncClient | None = None,
**kwargs: object,
) -> dict[str, Any]:
@ -133,45 +170,40 @@ class GoogleAIStudioTokenCounter:
litellm_params=kwargs,
)
cleaned_contents: Final = self._clean_contents_for_gemini_api(contents)
request_body: Final = (
{"contents": cleaned_contents} # mutable-ok: httpx json body takes a plain dict
if system_instruction is None and tools is None
else { # mutable-ok: httpx json body takes a plain dict
"generateContentRequest": { # mutable-ok: httpx json body takes a plain dict
"model": f"models/{model}",
"contents": cleaned_contents,
**(
{ # mutable-ok: httpx json body takes a plain dict
"systemInstruction": system_instruction,
}
if system_instruction is not None
else {} # mutable-ok: httpx json body takes a plain dict
),
**(
{ # mutable-ok: httpx json body takes a plain dict
"tools": tools,
}
if tools is not None
else {} # mutable-ok: httpx json body takes a plain dict
),
}
}
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(
llm_provider=LlmProviders.GEMINI,
)
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
)
# Check for HTTP errors
response.raise_for_status()
# Parse response
result: Final = response.json()
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
except litellm.APIError:
raise
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

@ -13,6 +13,8 @@ from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Final, cast
from pydantic import TypeAdapter, ValidationError
import litellm
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
LiteLLMAnthropicMessagesAdapter,
@ -34,6 +36,11 @@ class GeminiCountTokensPayload:
tools: list[Tools] | None
@dataclass(frozen=True, slots=True)
class InvalidCountTokensRequest:
message: str
_ANTHROPIC_PART_TYPES: Final = frozenset(
{
"tool_use",
@ -345,28 +352,36 @@ def _payload_from_openai_parts(
)
_ANTHROPIC_REQUEST: Final = TypeAdapter(AnthropicMessagesRequest)
def _build_anthropic_payload(
model: str,
messages: Sequence[Mapping[str, object]],
system: object | None,
tools: Sequence[Mapping[str, object]] | None,
) -> GeminiCountTokensPayload:
) -> GeminiCountTokensPayload | InvalidCountTokensRequest:
hosted_tools: Final = tuple(tool for tool in tools or () if _hosted_tool_type(tool) is not None)
adapter_tools: Final = tuple(tool for tool in tools or () if _hosted_tool_type(tool) is None)
anthropic_request: Final[AnthropicMessagesRequest] = cast( # cast-ok: adapter reads only the keys supplied
raw_request: Final = { # mutable-ok: transient request dict for the anthropic adapter
"model": model,
"messages": list( # mutable-ok: adapter contract takes a list of messages
_textify_server_side_blocks(messages)
),
**({"system": system} if system else {}), # mutable-ok: transient request dict for the anthropic adapter
**(
{"tools": list(adapter_tools)}
if adapter_tools
else {} # mutable-ok: transient request dict for the anthropic adapter
),
}
try:
_ANTHROPIC_REQUEST.validate_python(raw_request)
except ValidationError as e:
return InvalidCountTokensRequest(message=str(e))
anthropic_request: Final = cast( # cast-ok: validated above; pydantic returns lazy Iterable validators
AnthropicMessagesRequest,
{ # mutable-ok: transient request dict for the anthropic adapter
"model": model,
"messages": list( # mutable-ok: adapter contract takes a list of messages
_textify_server_side_blocks(messages)
),
**({"system": system} if system else {}), # mutable-ok: transient request dict for the anthropic adapter
**(
{"tools": list(adapter_tools)}
if adapter_tools
else {} # mutable-ok: transient request dict for the anthropic adapter
),
},
raw_request,
)
openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
anthropic_request, custom_llm_provider="gemini"
@ -422,10 +437,13 @@ def build_count_tokens_payload(
messages: Sequence[Mapping[str, object]],
system: object | None,
tools: Sequence[Mapping[str, object]] | None,
) -> GeminiCountTokensPayload:
if _has_anthropic_shape(system=system, tools=tools, messages=messages):
return _build_anthropic_payload(model=model, messages=messages, system=system, tools=tools)
return _build_openai_payload(model=model, messages=messages, system=system, tools=tools)
) -> GeminiCountTokensPayload | InvalidCountTokensRequest:
try:
if _has_anthropic_shape(system=system, tools=tools, messages=messages):
return _build_anthropic_payload(model=model, messages=messages, system=system, tools=tools)
return _build_openai_payload(model=model, messages=messages, system=system, tools=tools)
except (KeyError, TypeError, ValueError) as e:
return InvalidCountTokensRequest(message=f"Invalid token count request: {e!r}")
# Matches real inlineData blobs; a short or non-base64 `data` field (tool args,

View file

@ -13436,6 +13436,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

@ -1,17 +1,34 @@
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 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,7 @@ 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.proxy.proxy_server import _try_provider_token_count, token_counter
from litellm.types.utils import TokenCountResponse
verbose_proxy_logger.setLevel(level=logging.DEBUG)
@ -140,9 +141,7 @@ async def test_vLLM_token_counting():
print("response: ", response)
assert (
response.tokenizer_type == "openai_tokenizer"
) # SHOULD use the default tokenizer
assert response.tokenizer_type == "openai_tokenizer" # SHOULD use the default tokenizer
assert response.model_used == "wolfram/miquliz-120b-v2.0"
@ -175,9 +174,7 @@ async def test_token_counting_model_not_in_model_list():
print("response: ", response)
assert (
response.tokenizer_type == "openai_tokenizer"
) # SHOULD use the OpenAI tokenizer
assert response.tokenizer_type == "openai_tokenizer" # SHOULD use the OpenAI tokenizer
assert response.model_used == "special-alias"
@ -210,9 +207,7 @@ async def test_gpt_token_counting():
print("response: ", response)
assert (
response.tokenizer_type == "openai_tokenizer"
) # SHOULD use the OpenAI tokenizer
assert response.tokenizer_type == "openai_tokenizer" # SHOULD use the OpenAI tokenizer
assert response.request_model == "gpt-4"
@ -249,9 +244,7 @@ async def test_anthropic_messages_count_tokens_endpoint():
# Mock the internal token_counter function to return a controlled response
async def mock_token_counter(request, call_endpoint=False):
assert (
call_endpoint == True
), "Should be called with call_endpoint=True for Anthropic endpoint"
assert call_endpoint == True, "Should be called with call_endpoint=True for Anthropic endpoint"
assert request.model == "claude-3-sonnet-20240229"
assert request.messages == [{"role": "user", "content": "Hello Claude!"}]
@ -321,9 +314,7 @@ async def test_anthropic_messages_count_tokens_with_non_anthropic_model():
# Mock the internal token_counter function to return a controlled response
async def mock_token_counter(request, call_endpoint=True):
assert (
call_endpoint == True
), "Should be called with call_endpoint=True for Anthropic endpoint"
assert call_endpoint == True, "Should be called with call_endpoint=True for Anthropic endpoint"
assert request.model == "gpt-4"
assert request.messages == [{"role": "user", "content": "Hello GPT!"}]
@ -480,9 +471,7 @@ async def test_factory_anthropic_endpoint_calls_anthropic_counter():
# Mock the global handler instance in token_counter module
mock_handler = MagicMock()
mock_handler.handle_count_tokens_request = AsyncMock(
return_value={"input_tokens": 42}
)
mock_handler.handle_count_tokens_request = AsyncMock(return_value={"input_tokens": 42})
with patch(
"litellm.llms.anthropic.count_tokens.token_counter.anthropic_count_tokens_handler",
@ -537,9 +526,7 @@ async def test_factory_gpt4_endpoint_does_not_call_anthropic_counter():
# Mock the global handler instance in token_counter module
mock_handler = MagicMock()
mock_handler.handle_count_tokens_request = AsyncMock(
return_value={"input_tokens": 42}
)
mock_handler.handle_count_tokens_request = AsyncMock(return_value={"input_tokens": 42})
with patch(
"litellm.llms.anthropic.count_tokens.token_counter.anthropic_count_tokens_handler",
@ -596,9 +583,7 @@ async def test_factory_normal_token_counter_endpoint_does_not_call_anthropic():
# Mock the global handler instance in token_counter module
mock_handler = MagicMock()
mock_handler.handle_count_tokens_request = AsyncMock(
return_value={"input_tokens": 42}
)
mock_handler.handle_count_tokens_request = AsyncMock(return_value={"input_tokens": 42})
with patch(
"litellm.llms.anthropic.count_tokens.token_counter.anthropic_count_tokens_handler",
@ -613,9 +598,7 @@ async def test_factory_normal_token_counter_endpoint_does_not_call_anthropic():
mock_router.model_list = [
{
"model_name": "claude-3-5-sonnet",
"litellm_params": {
"model": "anthropic/claude-3-5-sonnet-20241022"
},
"litellm_params": {"model": "anthropic/claude-3-5-sonnet-20241022"},
"model_info": {},
}
]
@ -624,9 +607,7 @@ async def test_factory_normal_token_counter_endpoint_does_not_call_anthropic():
mock_router.async_get_available_deployment = AsyncMock(
return_value={
"model_name": "claude-3-5-sonnet",
"litellm_params": {
"model": "anthropic/claude-3-5-sonnet-20241022"
},
"litellm_params": {"model": "anthropic/claude-3-5-sonnet-20241022"},
"model_info": {},
}
)
@ -661,9 +642,7 @@ async def test_factory_registration():
assert counter is not None
# Create test deployments
anthropic_deployment = {
"litellm_params": {"model": "anthropic/claude-3-5-sonnet-20241022"}
}
anthropic_deployment = {"litellm_params": {"model": "anthropic/claude-3-5-sonnet-20241022"}}
non_anthropic_deployment = {"litellm_params": {"model": "openai/gpt-4"}}
@ -678,9 +657,7 @@ async def test_factory_registration():
assert not counter.should_use_token_counting_api(custom_llm_provider=None)
@pytest.mark.skip(
reason="Requires Google/Vertex AI credentials (GEMINI_API_KEY or VERTEX_AI_PRIVATE_KEY)."
)
@pytest.mark.skip(reason="Requires Google/Vertex AI credentials (GEMINI_API_KEY or VERTEX_AI_PRIVATE_KEY).")
@pytest.mark.asyncio
@pytest.mark.parametrize("model_name", ["gemini-2.5-pro", "vertex-ai-gemini-2.5-pro"])
async def test_vertex_ai_gemini_token_counting_with_contents(model_name):
@ -711,9 +688,7 @@ async def test_vertex_ai_gemini_token_counting_with_contents(model_name):
response = await token_counter(
request=TokenCountRequest(
model=model_name,
contents=[
{"parts": [{"text": "Hello world, how are you doing today? i am ij"}]}
],
contents=[{"parts": [{"text": "Hello world, how are you doing today? i am ij"}]}],
),
call_endpoint=True,
)
@ -750,9 +725,7 @@ async def test_bedrock_count_tokens_endpoint():
model_list=[
{
"model_name": "claude-bedrock",
"litellm_params": {
"model": "bedrock/anthropic.claude-3-sonnet-20240229-v1:0"
},
"litellm_params": {"model": "bedrock/anthropic.claude-3-sonnet-20240229-v1:0"},
}
]
)
@ -766,9 +739,7 @@ async def test_bedrock_count_tokens_endpoint():
}
# Test the mock handler directly to verify correct parameter extraction
await mock_count_tokens_handler(
request_data, {}, "anthropic.claude-3-sonnet-20240229-v1:0"
)
await mock_count_tokens_handler(request_data, {}, "anthropic.claude-3-sonnet-20240229-v1:0")
@pytest.mark.asyncio
@ -830,10 +801,7 @@ async def test_vertex_ai_anthropic_token_counting():
assert call_args is not None
assert call_args.kwargs["model"] == "claude-3-5-sonnet-20241022"
assert "messages" in call_args.kwargs["request_data"]
assert (
call_args.kwargs["request_data"]["messages"][0]["content"]
== "Hello Claude on Vertex AI! How are you?"
)
assert call_args.kwargs["request_data"]["messages"][0]["content"] == "Hello Claude on Vertex AI! How are you?"
# Validate response structure
assert response.model_used == "claude-3-5-sonnet-20241022"
@ -866,9 +834,7 @@ def test_vertex_ai_partner_models_token_counting_endpoint(vertex_location):
if vertex_location == "global":
assert endpoint.startswith("https://aiplatform.googleapis.com")
else:
assert endpoint.startswith(
f"https://{vertex_location}-aiplatform.googleapis.com"
)
assert endpoint.startswith(f"https://{vertex_location}-aiplatform.googleapis.com")
@pytest.mark.asyncio
@ -880,13 +846,9 @@ async def test_bedrock_token_counter_error_propagation_bedrock_error():
counter = BedrockTokenCounter()
# Mock the handler to raise BedrockError with specific status code
with patch.object(
counter, "count_tokens", wraps=counter.count_tokens
) as mock_count:
with patch.object(counter, "count_tokens", wraps=counter.count_tokens) as mock_count:
# We need to patch at the handler level
with patch(
"litellm.llms.bedrock.count_tokens.bedrock_token_counter.BedrockCountTokensHandler"
) as MockHandler:
with patch("litellm.llms.bedrock.count_tokens.bedrock_token_counter.BedrockCountTokensHandler") as MockHandler:
mock_handler_instance = MockHandler.return_value
mock_handler_instance.handle_count_tokens_request = AsyncMock(
side_effect=BedrockError(status_code=429, message="Rate limit exceeded")
@ -915,13 +877,9 @@ async def test_bedrock_token_counter_error_propagation_generic_exception():
"""
counter = BedrockTokenCounter()
with patch(
"litellm.llms.bedrock.count_tokens.bedrock_token_counter.BedrockCountTokensHandler"
) as MockHandler:
with patch("litellm.llms.bedrock.count_tokens.bedrock_token_counter.BedrockCountTokensHandler") as MockHandler:
mock_handler_instance = MockHandler.return_value
mock_handler_instance.handle_count_tokens_request = AsyncMock(
side_effect=Exception("Unexpected error")
)
mock_handler_instance.handle_count_tokens_request = AsyncMock(side_effect=Exception("Unexpected error"))
result = await counter.count_tokens(
model_to_use="anthropic.claude-3-sonnet",
@ -958,20 +916,14 @@ async def test_bedrock_handler_httpx_error_status_code_propagation():
with patch.object(handler, "validate_count_tokens_request"):
with patch.object(handler, "_get_aws_region_name", return_value="us-west-2"):
with patch.object(
handler, "transform_anthropic_to_bedrock_count_tokens", return_value={}
):
with patch.object(handler, "transform_anthropic_to_bedrock_count_tokens", return_value={}):
with patch.object(
handler,
"get_bedrock_count_tokens_endpoint",
return_value="https://example.com",
):
with patch.object(
handler, "_sign_request", return_value=({}, "{}")
):
with patch(
"litellm.llms.bedrock.count_tokens.handler.get_async_httpx_client"
) as mock_client:
with patch.object(handler, "_sign_request", return_value=({}, "{}")):
with patch("litellm.llms.bedrock.count_tokens.handler.get_async_httpx_client") as mock_client:
mock_async_client = AsyncMock()
mock_async_client.post = AsyncMock(side_effect=http_error)
mock_client.return_value = mock_async_client
@ -980,9 +932,7 @@ async def test_bedrock_handler_httpx_error_status_code_propagation():
await handler.handle_count_tokens_request(
request_data={
"model": "test",
"messages": [
{"role": "user", "content": "hello"}
],
"messages": [{"role": "user", "content": "hello"}],
},
litellm_params={},
resolved_model="anthropic.claude-3-sonnet",
@ -990,10 +940,7 @@ async def test_bedrock_handler_httpx_error_status_code_propagation():
assert exc_info.value.status_code == 403
# Message should be the raw response text
assert (
exc_info.value.message
== "Forbidden - Invalid credentials"
)
assert exc_info.value.message == "Forbidden - Invalid credentials"
@pytest.mark.asyncio
@ -1021,9 +968,7 @@ async def test_token_counter_httpx_status_error_raises_proxy_exception():
mock_counter.count_tokens = AsyncMock(side_effect=http_error)
# Save originals
original_get_provider_token_counter = (
litellm.proxy.proxy_server._get_provider_token_counter
)
original_get_provider_token_counter = litellm.proxy.proxy_server._get_provider_token_counter
original_router = litellm.proxy.proxy_server.llm_router
try:
@ -1031,9 +976,7 @@ async def test_token_counter_httpx_status_error_raises_proxy_exception():
def mock_get_provider_token_counter(deployment, model_to_use):
return (mock_counter, "claude-4-6-sonnet", "vertex_ai")
litellm.proxy.proxy_server._get_provider_token_counter = (
mock_get_provider_token_counter
)
litellm.proxy.proxy_server._get_provider_token_counter = mock_get_provider_token_counter
mock_router = MagicMock()
mock_router.async_get_available_deployment = AsyncMock(
@ -1061,12 +1004,74 @@ async def test_token_counter_httpx_status_error_raises_proxy_exception():
assert exc_info.value.type == "token_counting_error"
assert exc_info.value.param == "model"
finally:
litellm.proxy.proxy_server._get_provider_token_counter = (
original_get_provider_token_counter
)
litellm.proxy.proxy_server._get_provider_token_counter = original_get_provider_token_counter
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,
model_to_use: str,
messages: list | None,
contents: list | None,
deployment: dict | None = None,
request_model: str = "",
tools: list | None = None,
system: object | None = None,
) -> 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_proxy_token_counter_error_raises_exception_when_disabled():
"""
@ -1099,9 +1104,7 @@ async def test_proxy_token_counter_error_raises_exception_when_disabled():
# Save original value and function
original_disable = litellm.disable_token_counter
original_get_provider_token_counter = (
litellm.proxy.proxy_server._get_provider_token_counter
)
original_get_provider_token_counter = litellm.proxy.proxy_server._get_provider_token_counter
try:
litellm.disable_token_counter = True
@ -1115,9 +1118,7 @@ async def test_proxy_token_counter_error_raises_exception_when_disabled():
def mock_get_provider_token_counter(deployment, model_to_use):
return (mock_counter, "anthropic.claude-3-sonnet", "bedrock")
litellm.proxy.proxy_server._get_provider_token_counter = (
mock_get_provider_token_counter
)
litellm.proxy.proxy_server._get_provider_token_counter = mock_get_provider_token_counter
with pytest.raises(ProxyException) as exc_info:
await token_counter(
@ -1132,9 +1133,7 @@ async def test_proxy_token_counter_error_raises_exception_when_disabled():
assert "Rate limit exceeded" in exc_info.value.message
finally:
litellm.disable_token_counter = original_disable
litellm.proxy.proxy_server._get_provider_token_counter = (
original_get_provider_token_counter
)
litellm.proxy.proxy_server._get_provider_token_counter = original_get_provider_token_counter
@pytest.mark.asyncio
@ -1169,9 +1168,7 @@ async def test_proxy_token_counter_error_falls_back_when_enabled():
# Save original value and function
original_disable = litellm.disable_token_counter
original_get_provider_token_counter = (
litellm.proxy.proxy_server._get_provider_token_counter
)
original_get_provider_token_counter = litellm.proxy.proxy_server._get_provider_token_counter
try:
litellm.disable_token_counter = False
@ -1185,9 +1182,7 @@ async def test_proxy_token_counter_error_falls_back_when_enabled():
def mock_get_provider_token_counter(deployment, model_to_use):
return (mock_counter, "anthropic.claude-3-sonnet", "bedrock")
litellm.proxy.proxy_server._get_provider_token_counter = (
mock_get_provider_token_counter
)
litellm.proxy.proxy_server._get_provider_token_counter = mock_get_provider_token_counter
# Should not raise, should fall back to local tokenizer
result = await token_counter(
@ -1204,9 +1199,7 @@ async def test_proxy_token_counter_error_falls_back_when_enabled():
assert result.tokenizer_type != "bedrock_api"
finally:
litellm.disable_token_counter = original_disable
litellm.proxy.proxy_server._get_provider_token_counter = (
original_get_provider_token_counter
)
litellm.proxy.proxy_server._get_provider_token_counter = original_get_provider_token_counter
@pytest.mark.asyncio

View file

@ -3,6 +3,7 @@ import json
import httpx
import pytest
import litellm
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
COUNT_TOKENS_URL = "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:countTokens"
@ -60,6 +61,81 @@ async def test_acount_tokens_keeps_contents_body_without_system_or_tools():
assert body == {"contents": [{"role": "user", "parts": [{"text": "hi"}]}]}
@pytest.mark.asyncio
async def test_acount_tokens_sends_generate_content_request_with_system_only():
recorded: list[httpx.Request] = []
def _handler(request: httpx.Request) -> httpx.Response:
recorded.append(request)
return httpx.Response(200, json={"totalTokens": 9})
client = 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",
system_instruction={"parts": [{"text": "be terse"}]},
client=client,
)
body = json.loads(recorded[-1].content)
assert body == {
"generateContentRequest": {
"model": "models/gemini-2.5-flash",
"contents": [{"role": "user", "parts": [{"text": "hi"}]}],
"systemInstruction": {"parts": [{"text": "be terse"}]},
}
}
@pytest.mark.asyncio
async def test_acount_tokens_sends_generate_content_request_with_tools_only():
recorded: list[httpx.Request] = []
def _handler(request: httpx.Request) -> httpx.Response:
recorded.append(request)
return httpx.Response(200, json={"totalTokens": 9})
client = 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",
tools=[{"function_declarations": [{"name": "get_weather"}]}],
client=client,
)
body = json.loads(recorded[-1].content)
assert body == {
"generateContentRequest": {
"model": "models/gemini-2.5-flash",
"contents": [{"role": "user", "parts": [{"text": "hi"}]}],
"tools": [{"function_declarations": [{"name": "get_weather"}]}],
}
}
@pytest.mark.asyncio
async def test_acount_tokens_non_json_body_raises_api_error_with_response_status():
def _handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(200, content=b"<html>proxy error page</html>")
client = httpx.AsyncClient(transport=httpx.MockTransport(_handler))
with pytest.raises(litellm.APIError) as exc_info:
await GoogleAIStudioTokenCounter().acount_tokens(
model="gemini-2.5-flash",
contents=[{"role": "user", "parts": [{"text": "hi"}]}],
api_key="test-key",
client=client,
)
assert exc_info.value.status_code == 200
assert "non-JSON" in exc_info.value.message
@pytest.mark.asyncio
async def test_acount_tokens_wraps_unexpected_error_in_api_error():
import litellm

View file

@ -297,6 +297,20 @@ def test_build_count_tokens_payload_folds_system_into_contents_for_models_withou
assert [part.get("text") for part in payload.contents[0]["parts"]] == ["be nice", "hi"]
def test_build_count_tokens_payload_returns_invalid_for_malformed_anthropic_input():
from litellm.llms.gemini.count_tokens.transformation import InvalidCountTokensRequest
payload = build_count_tokens_payload(
model="gemini-2.5-flash",
messages=[{"role": "user", "content": "hi"}],
system={"not": "a valid system prompt"},
tools=[{"name": "get_weather", "input_schema": {"type": "object"}}],
)
assert isinstance(payload, InvalidCountTokensRequest)
assert payload.message
def test_normalize_count_tokens_tools_handles_each_tool_shape():
from litellm.llms.gemini.count_tokens.transformation import normalize_count_tokens_tools

View file

@ -144,7 +144,13 @@ class TestGoogleAIStudioTokenCounter:
assert result.original_response == mock_response
# Verify the mock was called correctly
mock_acount_tokens.assert_called_once_with(model=model_to_use, contents=contents, client=None)
mock_acount_tokens.assert_called_once_with(
system_instruction=None,
tools=None,
client=None,
model=model_to_use,
contents=contents,
)
@pytest.mark.asyncio
async def test_count_tokens_translates_anthropic_messages_system_and_tools(self):
@ -254,13 +260,13 @@ class TestGoogleAIStudioTokenCounter:
@pytest.mark.asyncio
async def test_count_tokens_translation_error_falls_back(self):
"""A crash translating bad message shapes must surface as an error
TokenCountResponse so the proxy falls back instead of 500ing."""
"""Malformed message shapes surface as a 400 error TokenCountResponse so
the proxy falls back instead of 500ing."""
token_counter = GoogleAIStudioTokenCounter()
result = await token_counter.count_tokens(
model_to_use="gemini-2.5-flash",
messages=[{"role": "tool", "content": "orphaned result", "tool_call_id": "missing-call"}],
messages=[{"role": "user", "content": [{"type": "text", "text": 123}]}],
contents=None,
deployment={"litellm_params": {"api_key": "test-key"}},
request_model="gemini/gemini-2.5-flash",
@ -268,21 +274,56 @@ class TestGoogleAIStudioTokenCounter:
assert result is not None
assert result.error is True
assert result.status_code == 500
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_unexpected_handler_error_returns_error_response(self):
"""A non-litellm exception escaping the handler must still surface as an
error TokenCountResponse so the proxy can fall back."""
async def test_count_tokens_malformed_anthropic_input_returns_400_without_http_call(self):
"""Input that fails Anthropic request validation returns a 400 error
response before any request reaches the provider."""
import httpx
recorded: list = []
def _handler(request):
recorded.append(request)
return httpx.Response(200, json={"totalTokens": 7})
token_counter = GoogleAIStudioTokenCounter()
result = await token_counter.count_tokens(
model_to_use="gemini-2.5-flash",
messages=[{"role": "user", "content": "hi"}],
contents=None,
deployment={"litellm_params": {"api_key": "test-key", "api_base": "https://gemini.example.test"}},
request_model="gemini/gemini-2.5-flash",
system={"not": "a valid system prompt"},
tools=[{"name": "get_weather", "input_schema": {"type": "object"}}],
client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
)
assert result is not None
assert result.error is True
assert result.status_code == 400
assert result.total_tokens == 0
assert recorded == []
@pytest.mark.asyncio
async def test_count_tokens_connection_error_returns_error_response(self):
"""A provider APIConnectionError surfaces as an error TokenCountResponse
so the proxy falls back 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 = RuntimeError("unexpected failure")
mock_acount_tokens.side_effect = litellm.APIConnectionError(
message="connection refused", llm_provider="gemini", model="gemini-2.5-flash"
)
result = await token_counter.count_tokens(
model_to_use="gemini-2.5-flash",
@ -295,7 +336,38 @@ class TestGoogleAIStudioTokenCounter:
assert result is not None
assert result.error is True
assert result.status_code == 500
assert "unexpected failure" in (result.error_message or "")
assert "connection refused" in (result.error_message or "")
@pytest.mark.asyncio
async def test_count_tokens_response_without_total_tokens_returns_502(self):
"""A 200 body missing totalTokens is a malformed provider response, not
a successful count, so it surfaces as a 502 error response."""
import httpx
recorded: list = []
def _handler(request):
recorded.append(request)
return httpx.Response(200, json={"promptTokensDetails": []})
token_counter = GoogleAIStudioTokenCounter()
result = await token_counter.count_tokens(
model_to_use="gemini-2.5-flash",
messages=[{"role": "user", "content": "hello"}],
contents=None,
deployment={"litellm_params": {"api_key": "test-key", "api_base": "https://gemini.example.test"}},
request_model="gemini/gemini-2.5-flash",
client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
)
assert recorded != []
assert result is not None
assert result.error is True
assert result.status_code == 502
assert result.total_tokens == 0
assert "totalTokens" in (result.error_message or "")
assert result.original_response == {"promptTokensDetails": []}
@pytest.mark.asyncio
async def test_count_tokens_returns_none_without_contents_or_messages(self):