fix(anthropic): derive the count-tokens URL from the deployment api_base

This commit is contained in:
mateo-berri 2026-09-02 13:06:03 -07:00
parent c813cc71ae
commit 412890ff52
3 changed files with 62 additions and 4 deletions

View file

@ -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)

View file

@ -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": <number>}
"""
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,

View file

@ -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}