From 613caa5e399e81546a9ff400654759a08d13784d Mon Sep 17 00:00:00 2001 From: shrey kharbanda Date: Sun, 27 Sep 2026 04:14:46 +0000 Subject: [PATCH] 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 --- litellm/llms/gemini/common_utils.py | 53 +++++++--- litellm/llms/gemini/count_tokens/handler.py | 42 ++++---- litellm/proxy/proxy_server.py | 3 +- .../unit/llms/gemini/count_tokens/__init__.py | 0 .../llms/gemini/count_tokens/test_handler.py | 65 ++++++++++++ .../llms/gemini/test_gemini_common_utils.py | 99 ++++++++++++++++++- tests/unit/proxy/test_proxy_token_counter.py | 68 +++++++++++++ 7 files changed, 293 insertions(+), 37 deletions(-) create mode 100644 tests/unit/llms/gemini/count_tokens/__init__.py create mode 100644 tests/unit/llms/gemini/count_tokens/test_handler.py diff --git a/litellm/llms/gemini/common_utils.py b/litellm/llms/gemini/common_utils.py index 78e6e6aaf82..0bde51c1384 100644 --- a/litellm/llms/gemini/common_utils.py +++ b/litellm/llms/gemini/common_utils.py @@ -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, + ) diff --git a/litellm/llms/gemini/count_tokens/handler.py b/litellm/llms/gemini/count_tokens/handler.py index c2f0ef473ae..5513be79e7a 100644 --- a/litellm/llms/gemini/count_tokens/handler.py +++ b/litellm/llms/gemini/count_tokens/handler.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f842f2e1e4a..7a6068f3b6f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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, ) diff --git a/tests/unit/llms/gemini/count_tokens/__init__.py b/tests/unit/llms/gemini/count_tokens/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/gemini/count_tokens/test_handler.py b/tests/unit/llms/gemini/count_tokens/test_handler.py new file mode 100644 index 00000000000..2147a02f9ac --- /dev/null +++ b/tests/unit/llms/gemini/count_tokens/test_handler.py @@ -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"proxy error page") + + 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, + ) diff --git a/tests/unit/llms/gemini/test_gemini_common_utils.py b/tests/unit/llms/gemini/test_gemini_common_utils.py index 70946e10590..edae8e2ec2b 100644 --- a/tests/unit/llms/gemini/test_gemini_common_utils.py +++ b/tests/unit/llms/gemini/test_gemini_common_utils.py @@ -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): diff --git a/tests/unit/proxy/test_proxy_token_counter.py b/tests/unit/proxy/test_proxy_token_counter.py index 8590e959961..b39759eca70 100644 --- a/tests/unit/proxy/test_proxy_token_counter.py +++ b/tests/unit/proxy/test_proxy_token_counter.py @@ -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"portal")) + + 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"