diff --git a/litellm/llms/anthropic/count_tokens/handler.py b/litellm/llms/anthropic/count_tokens/handler.py index 38cd429d99a..39968ac6fb3 100644 --- a/litellm/llms/anthropic/count_tokens/handler.py +++ b/litellm/llms/anthropic/count_tokens/handler.py @@ -41,7 +41,7 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig): model: The model identifier (e.g., "claude-3-5-sonnet-20241022") messages: The messages to count tokens for api_key: The Anthropic API key - api_base: Optional custom API base URL + api_base: Optional deployment api_base the count-tokens path is appended to timeout: Optional timeout for the request (defaults to litellm.request_timeout) Returns: @@ -67,7 +67,7 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig): verbose_logger.debug("Transformed request: %s", request_body) # Get endpoint URL - endpoint_url: Final = api_base or self.get_anthropic_count_tokens_endpoint() + endpoint_url: Final = self.get_anthropic_count_tokens_endpoint(api_base) verbose_logger.debug("Making request to: %s", endpoint_url) diff --git a/litellm/llms/anthropic/count_tokens/transformation.py b/litellm/llms/anthropic/count_tokens/transformation.py index 12581b9f658..7db2ea67830 100644 --- a/litellm/llms/anthropic/count_tokens/transformation.py +++ b/litellm/llms/anthropic/count_tokens/transformation.py @@ -7,6 +7,7 @@ This module handles the transformation of requests to Anthropic's CountTokens AP from typing import Any, Final from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION +from litellm.llms.anthropic.wif import anthropic_base_without_chat_suffix class AnthropicCountTokensConfig: @@ -19,14 +20,21 @@ class AnthropicCountTokensConfig: - Response: {"input_tokens": } """ - def get_anthropic_count_tokens_endpoint(self) -> str: + def get_anthropic_count_tokens_endpoint(self, api_base: str | None = None) -> str: """ Get the Anthropic CountTokens API endpoint. + Args: + api_base: The deployment's api_base, which names the chat surface (a host, or a + base already carrying ``/v1`` or ``/v1/messages``); the count-tokens path is + appended to it, so it is never the full count-tokens URL + Returns: The endpoint URL for the CountTokens API """ - return "https://api.anthropic.com/v1/messages/count_tokens" + if api_base is None: + return "https://api.anthropic.com/v1/messages/count_tokens" + return anthropic_base_without_chat_suffix(api_base) + "/v1/messages/count_tokens" def transform_request_to_count_tokens( self, diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_count_tokens_transformation.py b/tests/test_litellm/llms/anthropic/test_anthropic_count_tokens_transformation.py index a436744f648..8953b351acd 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_count_tokens_transformation.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_count_tokens_transformation.py @@ -1,3 +1,9 @@ +import httpx +import pytest +import respx + +import litellm +from litellm.llms.anthropic.count_tokens.handler import AnthropicCountTokensHandler from litellm.llms.anthropic.count_tokens.transformation import ( AnthropicCountTokensConfig, ) @@ -87,3 +93,47 @@ def test_transform_no_system_no_tools(): assert "system" not in result assert "tools" not in result + + +@pytest.mark.parametrize( + ("api_base", "expected"), + [ + (None, "https://api.anthropic.com/v1/messages/count_tokens"), + ("https://gateway.example", "https://gateway.example/v1/messages/count_tokens"), + ("https://gateway.example/", "https://gateway.example/v1/messages/count_tokens"), + ("https://gateway.example/v1", "https://gateway.example/v1/messages/count_tokens"), + ("https://gateway.example/anthropic/v1/messages", "https://gateway.example/anthropic/v1/messages/count_tokens"), + ], +) +def test_endpoint_appends_count_tokens_path_to_deployment_api_base(api_base, expected): + assert AnthropicCountTokensConfig().get_anthropic_count_tokens_endpoint(api_base) == expected + + +@pytest.fixture +def httpx_transport_clients(monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + client_cache = getattr(litellm, "in_memory_llm_clients_cache", None) + if client_cache is not None: + client_cache.flush_cache() + yield + if client_cache is not None: + client_cache.flush_cache() + + +@pytest.mark.asyncio +async def test_handler_posts_to_count_tokens_path_under_deployment_api_base(httpx_transport_clients): + """A deployment api_base names the chat host, so a handler that posts to it verbatim lands on + the host root, gets a 404, and the official count silently degrades to the local tokenizer.""" + with respx.mock: + route = respx.post("https://gateway.example/v1/messages/count_tokens").mock( + return_value=httpx.Response(200, json={"input_tokens": 7}) + ) + result = await AnthropicCountTokensHandler().handle_count_tokens_request( + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "hi"}], + api_key="sk-ant-api03-test-key", + api_base="https://gateway.example", + ) + + assert route.called + assert result == {"input_tokens": 7}