diff --git a/litellm/llms/gemini/common_utils.py b/litellm/llms/gemini/common_utils.py index 0ad95ee57db..dd905d5a72a 100644 --- a/litellm/llms/gemini/common_utils.py +++ b/litellm/llms/gemini/common_utils.py @@ -491,6 +491,7 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter): request_model: str = "", tools: list[dict[str, object]] | None = None, system: object | None = None, + client: httpx.AsyncClient | None = None, ) -> TokenCountResponse | None: import copy @@ -507,23 +508,26 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter): if contents is None else None ) + system_instruction: Final = payload.system_instruction if payload is not None else system + gemini_tools: Final = payload.tools if payload is not None else tools count_tokens_params: Final = { "model": model_to_use, "contents": payload.contents if payload is not None else contents, **( - {"system_instruction": payload.system_instruction} # mutable-ok: kwargs dict for acount_tokens - if payload is not None and payload.system_instruction is not None + {"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": payload.tools} # mutable-ok: kwargs dict for acount_tokens - if payload is not None and payload.tools is not None + {"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) try: result: Final = await GoogleAIStudioTokenCounter().acount_tokens( + client=client, **count_tokens_params_request, ) except (litellm.APIError, litellm.APIConnectionError) as e: diff --git a/litellm/llms/gemini/count_tokens/handler.py b/litellm/llms/gemini/count_tokens/handler.py index 8827a812085..142bafc89f2 100644 --- a/litellm/llms/gemini/count_tokens/handler.py +++ b/litellm/llms/gemini/count_tokens/handler.py @@ -88,6 +88,7 @@ class GoogleAIStudioTokenCounter: timeout: float | httpx.Timeout | None = None, system_instruction: SystemInstructions | None = None, tools: list[Tools] | None = None, + client: httpx.AsyncClient | None = None, **kwargs: object, ) -> dict[str, Any]: """ @@ -159,7 +160,7 @@ class GoogleAIStudioTokenCounter: } ) - async_httpx_client: Final = get_async_httpx_client( + async_httpx_client: Final = client or get_async_httpx_client( llm_provider=LlmProviders.GEMINI, ) @@ -185,5 +186,9 @@ class GoogleAIStudioTokenCounter: 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 + raise litellm.APIError( + message=f"Unexpected error during token counting: {e}", + llm_provider="gemini", + model=model, + status_code=500, + ) from e diff --git a/tests/test_litellm/llms/gemini/count_tokens/test_handler.py b/tests/test_litellm/llms/gemini/count_tokens/test_handler.py index a952f4f4392..3a3068fba4b 100644 --- a/tests/test_litellm/llms/gemini/count_tokens/test_handler.py +++ b/tests/test_litellm/llms/gemini/count_tokens/test_handler.py @@ -9,17 +9,14 @@ COUNT_TOKENS_URL = "https://generativelanguage.googleapis.com/v1beta/models/gemi @pytest.mark.asyncio -async def test_acount_tokens_sends_generate_content_request_when_system_or_tools_present(monkeypatch): +async def test_acount_tokens_sends_generate_content_request_when_system_or_tools_present(): recorded: list[httpx.Request] = [] def _handler(request: httpx.Request) -> httpx.Response: recorded.append(request) return httpx.Response(200, json={"totalTokens": 42}) - monkeypatch.setattr( - "litellm.llms.gemini.count_tokens.handler.get_async_httpx_client", - lambda **kwargs: httpx.AsyncClient(transport=httpx.MockTransport(_handler)), - ) + client = httpx.AsyncClient(transport=httpx.MockTransport(_handler)) result = await GoogleAIStudioTokenCounter().acount_tokens( model="gemini-2.5-flash", @@ -27,6 +24,7 @@ async def test_acount_tokens_sends_generate_content_request_when_system_or_tools api_key="test-key", system_instruction={"parts": [{"text": "You are a helpful assistant"}]}, tools=[{"function_declarations": [{"name": "get_weather"}]}], + client=client, ) assert result == {"totalTokens": 42} @@ -42,22 +40,20 @@ async def test_acount_tokens_sends_generate_content_request_when_system_or_tools @pytest.mark.asyncio -async def test_acount_tokens_keeps_contents_body_without_system_or_tools(monkeypatch): +async def test_acount_tokens_keeps_contents_body_without_system_or_tools(): recorded: list[httpx.Request] = [] def _handler(request: httpx.Request) -> httpx.Response: recorded.append(request) return httpx.Response(200, json={"totalTokens": 4}) - monkeypatch.setattr( - "litellm.llms.gemini.count_tokens.handler.get_async_httpx_client", - lambda **kwargs: httpx.AsyncClient(transport=httpx.MockTransport(_handler)), - ) + 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", + client=client, ) body = json.loads(recorded[-1].content) diff --git a/tests/test_litellm/llms/gemini/test_gemini_common_utils.py b/tests/test_litellm/llms/gemini/test_gemini_common_utils.py index 09273c8d1ad..5262adbb5f4 100644 --- a/tests/test_litellm/llms/gemini/test_gemini_common_utils.py +++ b/tests/test_litellm/llms/gemini/test_gemini_common_utils.py @@ -1,3 +1,4 @@ +import json from unittest.mock import AsyncMock, patch import pytest @@ -158,81 +159,114 @@ 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 ) @pytest.mark.asyncio async def test_count_tokens_translates_anthropic_messages_system_and_tools(self): """Anthropic-format messages are converted to gemini contents/system/tools before hitting the countTokens endpoint.""" + import httpx + + recorded: list = [] + + def _handler(request): + recorded.append(request) + return httpx.Response(200, json={"totalTokens": 12}) + token_counter = GoogleAIStudioTokenCounter() - with patch( - "litellm.llms.gemini.count_tokens.handler.GoogleAIStudioTokenCounter.acount_tokens", - new_callable=AsyncMock, - ) as mock_acount_tokens: - mock_acount_tokens.return_value = {"totalTokens": 12} + result = await token_counter.count_tokens( + model_to_use="gemini-2.5-flash", + messages=[{"role": "user", "content": "hello world"}], + contents=None, + deployment={"litellm_params": {"api_key": "test-key", "api_base": "https://gemini.example.test"}}, + request_model="gemini/gemini-2.5-flash", + tools=[ + { + "name": "get_weather", + "description": "Get the current weather for a city.", + "input_schema": { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], + }, + } + ], + system="You are a helpful assistant", + client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)), + ) - result = await token_counter.count_tokens( - model_to_use="gemini-2.5-flash", - messages=[{"role": "user", "content": "hello world"}], - contents=None, - deployment=None, - request_model="gemini/gemini-2.5-flash", - tools=[ - { - "name": "get_weather", - "description": "Get the current weather for a city.", - "input_schema": { - "type": "object", - "properties": {"city": {"type": "string"}}, - "required": ["city"], - }, - } - ], - system="You are a helpful assistant", - ) + assert result is not None + assert result.total_tokens == 12 + body = json.loads(recorded[-1].content) + generate_content_request = body["generateContentRequest"] + assert generate_content_request["contents"] + assert generate_content_request["contents"][0]["parts"][0].get("text") == "hello world" + assert generate_content_request["systemInstruction"]["parts"][0].get("text") == "You are a helpful assistant" + assert generate_content_request["tools"][0]["function_declarations"][0]["name"] == "get_weather" - assert result is not None - assert result.total_tokens == 12 - kwargs = mock_acount_tokens.call_args.kwargs - assert kwargs["contents"] - assert kwargs["contents"][0]["parts"][0].get("text") == "hello world" - assert kwargs["system_instruction"]["parts"][0].get("text") == "You are a helpful assistant" - assert kwargs["tools"][0]["function_declarations"][0]["name"] == "get_weather" + @pytest.mark.asyncio + async def test_count_tokens_passes_system_and_tools_with_native_contents(self): + """A request that already carries gemini contents still counts the + caller-supplied system instruction and tools.""" + import httpx + + recorded: list = [] + + def _handler(request): + recorded.append(request) + return httpx.Response(200, json={"totalTokens": 20}) + + token_counter = GoogleAIStudioTokenCounter() + + result = await token_counter.count_tokens( + model_to_use="gemini-2.5-flash", + messages=None, + contents=[{"role": "user", "parts": [{"text": "hello world"}]}], + deployment={"litellm_params": {"api_key": "test-key", "api_base": "https://gemini.example.test"}}, + request_model="gemini/gemini-2.5-flash", + system={"parts": [{"text": "You are a helpful assistant"}]}, + tools=[{"function_declarations": [{"name": "get_weather"}]}], + client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)), + ) + + assert result is not None + assert result.total_tokens == 20 + body = json.loads(recorded[-1].content) + generate_content_request = body["generateContentRequest"] + assert generate_content_request["contents"] == [{"role": "user", "parts": [{"text": "hello world"}]}] + assert generate_content_request["systemInstruction"] == {"parts": [{"text": "You are a helpful assistant"}]} + assert generate_content_request["tools"][0]["function_declarations"][0]["name"] == "get_weather" @pytest.mark.asyncio async def test_count_tokens_provider_error_returns_error_response(self): """A provider APIError must surface as an error TokenCountResponse so the proxy falls back to the local tokenizer instead of 500ing.""" - import litellm + import httpx + + def _handler(request): + return httpx.Response( + 400, + json={"error": {"code": 400, "message": "bad request", "status": "INVALID_ARGUMENT"}}, + ) token_counter = GoogleAIStudioTokenCounter() - with patch( - "litellm.llms.gemini.count_tokens.handler.GoogleAIStudioTokenCounter.acount_tokens", - new_callable=AsyncMock, - ) as mock_acount_tokens: - mock_acount_tokens.side_effect = litellm.APIError( - status_code=400, - message="Google Gen AI Studio API error: 400", - llm_provider="gemini", - model="gemini-2.5-flash", - ) + result = await token_counter.count_tokens( + model_to_use="gemini-2.5-flash", + messages=[{"role": "user", "content": "hello world"}], + contents=None, + deployment={"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)), + ) - result = await token_counter.count_tokens( - model_to_use="gemini-2.5-flash", - messages=[{"role": "user", "content": "hello world"}], - contents=None, - deployment=None, - request_model="gemini/gemini-2.5-flash", - ) - - assert result is not None - assert result.error is True - assert result.status_code == 400 - assert result.total_tokens == 0 - assert result.error_message is not None + assert result is not None + assert result.error is True + assert result.status_code == 400 + assert result.total_tokens == 0 + assert result.error_message is not None @pytest.mark.asyncio async def test_count_tokens_returns_none_without_contents_or_messages(self):