diff --git a/litellm/llms/gemini/common_utils.py b/litellm/llms/gemini/common_utils.py index d2e1f0ad294..b19e36d456b 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 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, + ) diff --git a/litellm/llms/gemini/count_tokens/handler.py b/litellm/llms/gemini/count_tokens/handler.py index 3894a6f4837..73b9018a1d9 100644 --- a/litellm/llms/gemini/count_tokens/handler.py +++ b/litellm/llms/gemini/count_tokens/handler.py @@ -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( diff --git a/litellm/llms/gemini/count_tokens/transformation.py b/litellm/llms/gemini/count_tokens/transformation.py index c4098adb823..aac2915cff5 100644 --- a/litellm/llms/gemini/count_tokens/transformation.py +++ b/litellm/llms/gemini/count_tokens/transformation.py @@ -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, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 87992d8da07..b1aa0c90404 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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( diff --git a/litellm/types/llms/gemini.py b/litellm/types/llms/gemini.py index 57fb8b5b0cd..32a6766d604 100644 --- a/litellm/types/llms/gemini.py +++ b/litellm/types/llms/gemini.py @@ -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" diff --git a/tests/proxy_unit_tests/test_proxy_token_counter.py b/tests/proxy_unit_tests/test_proxy_token_counter.py index 39ec4bb1887..81e62eb77ba 100644 --- a/tests/proxy_unit_tests/test_proxy_token_counter.py +++ b/tests/proxy_unit_tests/test_proxy_token_counter.py @@ -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 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 c245e02141e..cc09fd215d4 100644 --- a/tests/test_litellm/llms/gemini/count_tokens/test_handler.py +++ b/tests/test_litellm/llms/gemini/count_tokens/test_handler.py @@ -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"proxy error page") + + 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 diff --git a/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py b/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py index 254e4b96858..447d4d4dda2 100644 --- a/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py +++ b/tests/test_litellm/llms/gemini/count_tokens/test_transformation.py @@ -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 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 062bf207357..0c666fe601f 100644 --- a/tests/test_litellm/llms/gemini/test_gemini_common_utils.py +++ b/tests/test_litellm/llms/gemini/test_gemini_common_utils.py @@ -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):