From d9b612e1e2f247829cfca2de90b4ee0d46a23a25 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 12:47:12 +0000 Subject: [PATCH] fix(gemini): count Anthropic messages, system and tools on /v1/messages/count_tokens and fall back on provider errors Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/gemini/common_utils.py | 34 ++++++- litellm/llms/gemini/count_tokens/handler.py | 19 +++- .../gemini/count_tokens/transformation.py | 61 ++++++++++++ .../test_messages_count_tokens_e2e.py | 93 +++++++++++++++++++ tests/e2e/models.py | 2 + .../llms/gemini/count_tokens/__init__.py | 0 .../llms/gemini/count_tokens/test_handler.py | 64 +++++++++++++ .../count_tokens/test_transformation.py | 62 +++++++++++++ .../llms/gemini/test_gemini_common_utils.py | 87 +++++++++++++++++ 9 files changed, 416 insertions(+), 6 deletions(-) create mode 100644 litellm/llms/gemini/count_tokens/transformation.py create mode 100644 tests/e2e/llm_translation/test_messages_count_tokens_e2e.py create mode 100644 tests/test_litellm/llms/gemini/count_tokens/__init__.py create mode 100644 tests/test_litellm/llms/gemini/count_tokens/test_handler.py create mode 100644 tests/test_litellm/llms/gemini/count_tokens/test_transformation.py diff --git a/litellm/llms/gemini/common_utils.py b/litellm/llms/gemini/common_utils.py index 78e6e6aaf82..f20f904596d 100644 --- a/litellm/llms/gemini/common_utils.py +++ b/litellm/llms/gemini/common_utils.py @@ -495,17 +495,43 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter): import copy from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter + from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload + + if contents is None and not messages: + return None deployment = deployment or {} count_tokens_params_request: Final = copy.deepcopy(deployment.get("litellm_params", {})) + payload: Final = ( + build_count_tokens_payload(model=model_to_use, messages=messages, system=system, tools=tools) + if contents is None + else None + ) count_tokens_params: Final = { "model": model_to_use, - "contents": contents, + "contents": payload.contents if payload is not None else contents, + **( + {"system_instruction": payload.system_instruction} + if payload is not None and payload.system_instruction is not None + else {} + ), + **({"tools": payload.tools} if payload is not None and payload.tools is not None else {}), } count_tokens_params_request.update(count_tokens_params) - result: Final = await GoogleAIStudioTokenCounter().acount_tokens( - **count_tokens_params_request, - ) + try: + result: Final = await GoogleAIStudioTokenCounter().acount_tokens( + **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 result is not None: return TokenCountResponse( diff --git a/litellm/llms/gemini/count_tokens/handler.py b/litellm/llms/gemini/count_tokens/handler.py index c2f0ef473ae..d055b7ce5ff 100644 --- a/litellm/llms/gemini/count_tokens/handler.py +++ b/litellm/llms/gemini/count_tokens/handler.py @@ -4,6 +4,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.utils import LlmProviders if TYPE_CHECKING: @@ -51,7 +53,7 @@ class GoogleAIStudioTokenCounter: """ Construct the URL for the Google Gen AI Studio countTokens endpoint. """ - base_url: Final = api_base or "https://generativelanguage.googleapis.com" + base_url: Final = api_base or get_secret_str("GEMINI_API_BASE") or "https://generativelanguage.googleapis.com" return f"{base_url}/v1beta/models/{model}:countTokens" async def validate_environment( @@ -84,6 +86,8 @@ class GoogleAIStudioTokenCounter: api_key: str | None = None, api_base: str | None = None, timeout: float | httpx.Timeout | None = None, + system_instruction: SystemInstructions | None = None, + tools: list[Tools] | None = None, **kwargs: object, ) -> dict[str, Any]: """ @@ -130,7 +134,18 @@ 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} + 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 {}), + } + } + ) async_httpx_client: Final = get_async_httpx_client( llm_provider=LlmProviders.GEMINI, diff --git a/litellm/llms/gemini/count_tokens/transformation.py b/litellm/llms/gemini/count_tokens/transformation.py new file mode 100644 index 00000000000..25380ab625a --- /dev/null +++ b/litellm/llms/gemini/count_tokens/transformation.py @@ -0,0 +1,61 @@ +"""Translate an Anthropic /v1/messages/count_tokens request into a Gemini +countTokens payload (contents + systemInstruction + tools).""" + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from typing import Final, cast + +from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import ( + LiteLLMAnthropicMessagesAdapter, +) +from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, # pyright: ignore[reportPrivateUsage] # shared helper already used by gemini/chat, context_caching, and vertex_and_google_ai_studio_gemini + _transform_system_message, +) +from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig +from litellm.types.llms.anthropic import AnthropicMessagesRequest +from litellm.types.llms.vertex_ai import ContentType, SystemInstructions, Tools + + +@dataclass(frozen=True, slots=True) +class GeminiCountTokensPayload: + contents: list[ContentType] + system_instruction: SystemInstructions | None + tools: list[Tools] | None + + +def build_count_tokens_payload( + model: str, + messages: Sequence[Mapping[str, object]], + system: object | None, + tools: Sequence[Mapping[str, object]] | None, +) -> GeminiCountTokensPayload: + anthropic_request: Final[AnthropicMessagesRequest] = cast( + AnthropicMessagesRequest, # cast-ok: untrusted client payload, adapter reads the anthropic-shape keys only + { + "model": model, + "messages": list(messages), + **({"system": system} if system else {}), + **({"tools": list(tools)} if tools else {}), + }, + ) + openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai( + anthropic_request, custom_llm_provider="gemini" + ) + system_instruction, remaining_messages = _transform_system_message( + supports_system_message=True, + messages=list(openai_request["messages"]), + ) + contents: Final = _gemini_convert_messages_with_history( + messages=remaining_messages, + model=model, + custom_llm_provider="gemini", + ) + openai_tools: Final = openai_request.get("tools") + return GeminiCountTokensPayload( + contents=contents, + system_instruction=system_instruction, + tools=VertexGeminiConfig()._map_function(value=[dict(tool) for tool in openai_tools], optional_params={}) + if openai_tools + else None, + ) diff --git a/tests/e2e/llm_translation/test_messages_count_tokens_e2e.py b/tests/e2e/llm_translation/test_messages_count_tokens_e2e.py new file mode 100644 index 00000000000..d732f190ab2 --- /dev/null +++ b/tests/e2e/llm_translation/test_messages_count_tokens_e2e.py @@ -0,0 +1,93 @@ +"""Live e2e: POST /v1/messages/count_tokens against a gemini/ deployment. + +Registers a fresh gemini/gemini-2.5-flash deployment via /model/new (deleted on +teardown) and drives the endpoint through the shared transport, since no +official provider SDK covers this route. +""" + +import pytest +from e2e_config import unique_marker +from e2e_http import unwrap +from lifecycle import ResourceManager +from models import ( + AnthropicCustomTool, + ChatMessage, + CountTokensBody, + JsonSchemaProperty, + LiteLLMParamsBody, + ToolInputSchema, +) +from proxy_client import ProxyClient + +pytestmark = pytest.mark.e2e + +BACKEND_MODEL = "gemini/gemini-2.5-flash" +GEMINI_API_KEY = "os.environ/GEMINI_API_KEY" + +WEATHER_TOOL = AnthropicCustomTool( + name="get_weather", + description="Get the current weather for a city.", + input_schema=ToolInputSchema( + properties={"city": JsonSchemaProperty(type="string")}, + required=["city"], + ), +) + + +def _provision(proxy: ProxyClient, resources: ResourceManager) -> str: + model_name = f"e2e-count-tokens-gemini-{unique_marker()}" + model_id = proxy.create_model( + model_name, + LiteLLMParamsBody(model=BACKEND_MODEL, api_key=GEMINI_API_KEY), + ) + resources.defer(lambda: proxy.delete_model(model_id)) + return model_name + + +class TestMessagesCountTokens: + def test_count_tokens_gemini_returns_input_tokens( + self, proxy: ProxyClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = _provision(proxy, resources) + + response = unwrap( + proxy.count_tokens( + scoped_key, + CountTokensBody( + model=model, + messages=[ChatMessage(role="user", content=f"hello world {unique_marker()}")], + ), + ) + ) + assert response.input_tokens > 0, f"input_tokens not positive: {response.input_tokens}" + + def test_count_tokens_gemini_with_system_and_tools( + self, proxy: ProxyClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = _provision(proxy, resources) + prompt = f"hello world {unique_marker()}" + + plain = unwrap( + proxy.count_tokens( + scoped_key, + CountTokensBody( + model=model, + messages=[ChatMessage(role="user", content=prompt)], + ), + ) + ) + with_system_and_tools = unwrap( + proxy.count_tokens( + scoped_key, + CountTokensBody( + model=model, + messages=[ChatMessage(role="user", content=prompt)], + system="You are a helpful assistant", + tools=[WEATHER_TOOL], + ), + ) + ) + assert with_system_and_tools.input_tokens > plain.input_tokens, ( + f"system + tools count {with_system_and_tools.input_tokens} did not exceed " + f"the plain message count {plain.input_tokens}; the extra prompt was not counted" + ) diff --git a/tests/e2e/models.py b/tests/e2e/models.py index bd7f5171172..a26676148f7 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -574,6 +574,8 @@ class CountTokensBody(BaseModel): model: str messages: list[ChatMessage] + system: str | None = None + tools: list[AnthropicTool] | None = None class AnthropicMessagesResponse(BaseModel): diff --git a/tests/test_litellm/llms/gemini/count_tokens/__init__.py b/tests/test_litellm/llms/gemini/count_tokens/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/gemini/count_tokens/test_handler.py b/tests/test_litellm/llms/gemini/count_tokens/test_handler.py new file mode 100644 index 00000000000..a952f4f4392 --- /dev/null +++ b/tests/test_litellm/llms/gemini/count_tokens/test_handler.py @@ -0,0 +1,64 @@ +import json + +import httpx +import pytest + +from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter + +COUNT_TOKENS_URL = "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:countTokens" + + +@pytest.mark.asyncio +async def test_acount_tokens_sends_generate_content_request_when_system_or_tools_present(monkeypatch): + 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)), + ) + + result = await GoogleAIStudioTokenCounter().acount_tokens( + model="gemini-2.5-flash", + contents=[{"role": "user", "parts": [{"text": "hello world"}]}], + api_key="test-key", + system_instruction={"parts": [{"text": "You are a helpful assistant"}]}, + tools=[{"function_declarations": [{"name": "get_weather"}]}], + ) + + assert result == {"totalTokens": 42} + request = recorded[-1] + assert request.url == COUNT_TOKENS_URL + body = json.loads(request.content) + assert "contents" not in body + generate_content_request = body["generateContentRequest"] + assert generate_content_request["model"] == "models/gemini-2.5-flash" + 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_acount_tokens_keeps_contents_body_without_system_or_tools(monkeypatch): + 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)), + ) + + await GoogleAIStudioTokenCounter().acount_tokens( + model="gemini-2.5-flash", + contents=[{"role": "user", "parts": [{"text": "hi"}]}], + api_key="test-key", + ) + + body = json.loads(recorded[-1].content) + assert body == {"contents": [{"role": "user", "parts": [{"text": "hi"}]}]} diff --git a/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py b/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py new file mode 100644 index 00000000000..829f8d49e29 --- /dev/null +++ b/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py @@ -0,0 +1,62 @@ +from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload + + +def test_build_count_tokens_payload_translates_anthropic_request(): + payload = build_count_tokens_payload( + model="gemini-2.5-flash", + messages=[{"role": "user", "content": "hello world"}], + system="You are a helpful assistant", + tools=[ + { + "name": "get_weather", + "description": "Get the current weather for a city.", + "input_schema": { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], + }, + } + ], + ) + + assert payload.contents + assert payload.contents[0]["parts"][0].get("text") == "hello world" + assert payload.system_instruction is not None + assert payload.system_instruction["parts"][0].get("text") == "You are a helpful assistant" + assert payload.tools is not None + function_declarations = payload.tools[0]["function_declarations"] + assert function_declarations[0]["name"] == "get_weather" + assert function_declarations[0].get("parameters", {}).get("required") == ["city"] + + +def test_build_count_tokens_payload_without_system_or_tools(): + payload = build_count_tokens_payload( + model="gemini-2.5-flash", + messages=[{"role": "user", "content": "hi"}], + system=None, + tools=None, + ) + + assert payload.contents + assert payload.system_instruction is None + assert payload.tools is None + + +def test_build_count_tokens_payload_passes_openai_tools_through(): + payload = build_count_tokens_payload( + model="gemini-2.5-flash", + messages=[{"role": "user", "content": "hi"}], + system=None, + tools=[ + { + "type": "function", + "function": { + "name": "get_weather", + "parameters": {"type": "object", "properties": {"city": {"type": "string"}}}, + }, + } + ], + ) + + assert payload.tools is not None + assert payload.tools[0]["function_declarations"][0]["name"] == "get_weather" 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 70946e10590..09273c8d1ad 100644 --- a/tests/test_litellm/llms/gemini/test_gemini_common_utils.py +++ b/tests/test_litellm/llms/gemini/test_gemini_common_utils.py @@ -161,6 +161,93 @@ class TestGoogleAIStudioTokenCounter: model=model_to_use, contents=contents ) + @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.""" + 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=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 + 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_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 + + 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=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 + + @pytest.mark.asyncio + async def test_count_tokens_returns_none_without_contents_or_messages(self): + token_counter = GoogleAIStudioTokenCounter() + + result = await token_counter.count_tokens( + model_to_use="gemini-2.5-flash", + messages=None, + contents=None, + deployment=None, + request_model="gemini/gemini-2.5-flash", + ) + + assert result is None + def test_clean_contents_for_gemini_api_removes_id_field(self): """Test that _clean_contents_for_gemini_api removes unsupported 'id' field from function responses""" from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter