mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(anthropic): derive the count-tokens URL from the deployment api_base
This commit is contained in:
parent
c813cc71ae
commit
412890ff52
3 changed files with 62 additions and 4 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue