fix(gemini): fall back to the local estimate when Gemini's count fails

The Gemini counter now returns a failed count for API errors, connection
errors, timeouts, non-JSON bodies and responses without an integer
totalTokens, so the proxy falls back to its local estimate instead of
answering 500
This commit is contained in:
shrey kharbanda 2026-09-27 04:14:46 +00:00
parent 303434d573
commit 613caa5e39
No known key found for this signature in database
7 changed files with 293 additions and 37 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 TYPE_CHECKING, Any, Final
import httpx
@ -15,6 +15,9 @@ from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import TokenCountResponse
if TYPE_CHECKING:
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
GEMINI_IMAGE_ASPECT_RATIOS: Final[dict[str, float]] = {
"1:1": 1 / 1,
"1:4": 1 / 4,
@ -491,11 +494,26 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
request_model: str = "",
tools: list[dict[str, object]] | None = None,
system: object | None = None,
client: "httpx.AsyncClient | AsyncHTTPHandler | None" = None,
) -> TokenCountResponse | None:
import copy
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
def failed(
message: str, status_code: int, original_response: dict[str, object] | None = None
) -> TokenCountResponse:
return TokenCountResponse(
total_tokens=0,
request_model=request_model,
model_used=model_to_use,
tokenizer_type="gemini_api",
error=True,
error_message=message,
status_code=status_code,
original_response=original_response,
)
deployment = deployment or {}
count_tokens_params_request: Final = copy.deepcopy(deployment.get("litellm_params", {}))
count_tokens_params: Final = {
@ -503,17 +521,24 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
"contents": contents,
}
count_tokens_params_request.update(count_tokens_params)
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
**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=result.get("tokenizer_used", ""),
original_response=result,
try:
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
client=client,
**count_tokens_params_request,
)
return None
except (litellm.APIError, litellm.APIConnectionError) as e:
return failed(e.message, e.status_code)
total_tokens: Final = result.get("totalTokens") if isinstance(result, dict) else None
if not isinstance(total_tokens, int) or isinstance(total_tokens, bool):
return failed(
"Google Gen AI Studio countTokens response has no totalTokens",
502,
result if isinstance(result, dict) else None,
)
return TokenCountResponse(
total_tokens=total_tokens,
request_model=request_model,
model_used=model_to_use,
tokenizer_type=result.get("tokenizer_used", ""),
original_response=result,
)

View file

@ -3,7 +3,7 @@ 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.utils import LlmProviders
if TYPE_CHECKING:
@ -84,6 +84,7 @@ class GoogleAIStudioTokenCounter:
api_key: str | None = None,
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
client: httpx.AsyncClient | AsyncHTTPHandler | None = None,
**kwargs: object,
) -> dict[str, Any]:
"""
@ -96,6 +97,7 @@ 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
client: Optional HTTP client to send the request with
**kwargs: Additional parameters
Returns:
@ -113,10 +115,8 @@ class GoogleAIStudioTokenCounter:
}
Raises:
ValueError: If API key is missing
litellm.APIError: If the API call fails
litellm.APIConnectionError: If the connection fails
Exception: For any other unexpected errors
litellm.APIError: If the API returns an error status or a body that is not JSON
litellm.APIConnectionError: If the request fails or times out
"""
# Prepare headers
@ -132,31 +132,31 @@ class GoogleAIStudioTokenCounter:
cleaned_contents: Final = self._clean_contents_for_gemini_api(contents)
request_body: Final = {"contents": cleaned_contents}
async_httpx_client: Final = get_async_httpx_client(
llm_provider=LlmProviders.GEMINI,
)
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(
message=error_msg,
message=f"Google Gen AI Studio API error: {e.response.status_code} - {e.response.text}",
llm_provider="gemini",
model=model,
status_code=e.response.status_code,
) from e
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
except (httpx.RequestError, litellm.Timeout) as e:
raise litellm.APIConnectionError(
message=f"Request to Google Gen AI Studio failed: {e}", llm_provider="gemini", model=model
) from e
try:
return 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=502,
) from e

View file

@ -13676,7 +13676,8 @@ async def _try_provider_token_count(
code=result.status_code or 500,
)
verbose_proxy_logger.warning(
"Provider token counting failed (%s): %s. Falling back to local tokenizer.",
"Provider token counting for model %s failed (%s): %s. Falling back to local tokenizer.",
model_to_use,
result.status_code,
result.error_message,
)

View file

@ -0,0 +1,65 @@
import httpx
import pytest
import litellm
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
@pytest.mark.asyncio
async def test_acount_tokens_non_json_body_raises_api_error_with_502():
def _handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(200, content=b"<html>proxy error page</html>")
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=httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
)
assert exc_info.value.status_code == 502
assert "non-JSON" in exc_info.value.message
@pytest.mark.asyncio
async def test_acount_tokens_lets_internal_errors_propagate():
def _handler(request: httpx.Request) -> httpx.Response:
raise RuntimeError("transport exploded")
with pytest.raises(RuntimeError, match="transport exploded"):
await GoogleAIStudioTokenCounter().acount_tokens(
model="gemini-2.5-flash",
contents=[{"role": "user", "parts": [{"text": "hello"}]}],
api_key="test-key",
client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
)
def _timing_out(request: httpx.Request) -> httpx.Response:
raise httpx.ReadTimeout("timed out", request=request)
def _litellm_handler_timing_out() -> AsyncHTTPHandler:
handler = AsyncHTTPHandler()
handler.client = httpx.AsyncClient(transport=httpx.MockTransport(_timing_out))
return handler
@pytest.mark.asyncio
@pytest.mark.parametrize(
"client",
(
pytest.param(httpx.AsyncClient(transport=httpx.MockTransport(_timing_out)), id="httpx-client"),
pytest.param(_litellm_handler_timing_out(), id="litellm-http-handler"),
),
)
async def test_acount_tokens_raises_connection_error_on_timeout(client):
with pytest.raises(litellm.APIConnectionError):
await GoogleAIStudioTokenCounter().acount_tokens(
model="gemini-2.5-flash",
contents=[{"role": "user", "parts": [{"text": "hello"}]}],
api_key="test-key",
client=client,
)

View file

@ -89,6 +89,103 @@ class TestGeminiModelInfo:
class TestGoogleAIStudioTokenCounter:
async def _count(self, handler, litellm_params=None):
import httpx
return await GoogleAIStudioTokenCounter().count_tokens(
model_to_use="gemini-2.5-flash",
messages=None,
contents=[{"role": "user", "parts": [{"text": "hello"}]}],
deployment={"litellm_params": litellm_params or {"api_key": "test-key"}},
request_model="gemini/gemini-2.5-flash",
client=httpx.AsyncClient(transport=httpx.MockTransport(handler)),
)
@pytest.mark.asyncio
async def test_count_tokens_provider_error_returns_error_response(self):
import httpx
result = await self._count(
lambda request: httpx.Response(
400, json={"error": {"code": 400, "message": "bad request", "status": "INVALID_ARGUMENT"}}
)
)
assert result is not None
assert result.error is True
assert result.status_code == 400
assert result.total_tokens == 0
assert "bad request" in (result.error_message or "")
@pytest.mark.asyncio
async def test_count_tokens_without_api_key_returns_provider_error_response(self, monkeypatch):
import httpx
import litellm
monkeypatch.delenv("GEMINI_API_KEY", raising=False)
monkeypatch.delenv("GOOGLE_API_KEY", raising=False)
monkeypatch.setattr(litellm, "api_key", None)
recorded = []
def _handler(request):
recorded.append(request)
return httpx.Response(
403, json={"error": {"code": 403, "message": "API key not valid", "status": "PERMISSION_DENIED"}}
)
result = await self._count(_handler, litellm_params={"model": "gemini/gemini-2.5-flash"})
assert result is not None
assert result.error is True
assert result.status_code == 403
assert len(recorded) == 1 and "x-goog-api-key" not in recorded[0].headers
@pytest.mark.asyncio
async def test_count_tokens_connection_error_returns_error_response(self):
import httpx
def _handler(request: httpx.Request) -> httpx.Response:
raise httpx.ConnectError("connection refused", request=request)
result = await self._count(_handler)
assert result is not None
assert result.error is True
assert result.status_code == 500
assert "connection refused" in (result.error_message or "")
@pytest.mark.asyncio
@pytest.mark.parametrize(
"upstream_json",
[
{"totalTokens": "abc"},
{"totalTokens": True},
{"promptTokensDetails": []},
[{"totalTokens": 5}],
],
)
async def test_count_tokens_malformed_provider_response_returns_502(self, upstream_json):
import httpx
result = await self._count(lambda request: httpx.Response(200, json=upstream_json))
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 "")
@pytest.mark.asyncio
async def test_count_tokens_valid_response_returns_the_count(self):
import httpx
result = await self._count(lambda request: httpx.Response(200, json={"totalTokens": 7}))
assert result is not None
assert result.error is not True
assert result.total_tokens == 7
"""Test suite for GoogleAIStudioTokenCounter class"""
def test_should_use_token_counting_api(self):
@ -158,7 +255,7 @@ 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
)
def test_clean_contents_for_gemini_api_removes_id_field(self):

View file

@ -1245,3 +1245,71 @@ async def test_anthropic_endpoint_429_rate_limit_error_format():
finally:
anthropic_endpoints._read_request_body = original_read_request_body
proxy_server.token_counter = original_token_counter
def _gemini_router() -> Router:
return Router(
model_list=[
{
"model_name": "gemini-count",
"litellm_params": {"model": "gemini/gemini-2.5-flash", "api_key": "fake-gemini-key"},
}
]
)
_GEMINI_COUNT_TOKENS_URL = "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:countTokens"
@pytest.mark.asyncio
async def test_gemini_count_error_falls_back_to_the_local_estimate(monkeypatch, respx_mock):
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", _gemini_router())
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setattr(litellm, "disable_token_counter", False)
count_route = respx_mock.post(_GEMINI_COUNT_TOKENS_URL).mock(
return_value=httpx.Response(400, json={"error": {"code": 400, "message": "API key not valid"}})
)
response = await token_counter(
request=TokenCountRequest(model="gemini-count", messages=[{"role": "user", "content": "hello world"}]),
call_endpoint=True,
)
assert count_route.called
assert response.error is not True
assert response.total_tokens > 0
assert response.tokenizer_type != "gemini_api"
@pytest.mark.asyncio
async def test_gemini_count_error_is_returned_when_fallback_is_disabled(monkeypatch, respx_mock):
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", _gemini_router())
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setattr(litellm, "disable_token_counter", True)
respx_mock.post(_GEMINI_COUNT_TOKENS_URL).mock(
return_value=httpx.Response(400, json={"error": {"code": 400, "message": "API key not valid"}})
)
with pytest.raises(ProxyException) as exc_info:
await token_counter(
request=TokenCountRequest(model="gemini-count", messages=[{"role": "user", "content": "hi"}]),
call_endpoint=True,
)
assert exc_info.value.code == "400"
assert "API key not valid" in exc_info.value.message
@pytest.mark.asyncio
async def test_gemini_non_json_success_body_surfaces_as_bad_gateway_when_fallback_disabled(monkeypatch, respx_mock):
monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", _gemini_router())
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setattr(litellm, "disable_token_counter", True)
respx_mock.post(_GEMINI_COUNT_TOKENS_URL).mock(return_value=httpx.Response(200, content=b"<html>portal</html>"))
with pytest.raises(ProxyException) as exc_info:
await token_counter(
request=TokenCountRequest(model="gemini-count", messages=[{"role": "user", "content": "hi"}]),
call_endpoint=True,
)
assert exc_info.value.code == "502"