diff --git a/litellm/llms/anthropic/count_tokens/handler.py b/litellm/llms/anthropic/count_tokens/handler.py index 6594209a39a..b4f107cdef4 100644 --- a/litellm/llms/anthropic/count_tokens/handler.py +++ b/litellm/llms/anthropic/count_tokens/handler.py @@ -32,7 +32,7 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig): self, model: str, messages: list[dict[str, JsonValue]], - api_key: str, + auth_header: Mapping[str, str], api_base: str | None = None, timeout: float | httpx.Timeout | None = None, tools: list[dict[str, JsonValue]] | None = None, @@ -45,7 +45,7 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig): Args: model: The model identifier (e.g., "claude-3-5-sonnet-20241022") messages: The messages to count tokens for - api_key: The Anthropic API key + auth_header: The resolved Anthropic auth header (``AnthropicModelInfo.get_auth_header``) api_base: Optional deployment api_base the count-tokens path is appended to timeout: Optional timeout for the request (defaults to litellm.request_timeout) @@ -78,7 +78,7 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig): verbose_logger.debug("Making request to: %s", endpoint_url) # Get required headers - headers: Final = self.get_required_headers(api_key) + headers: Final = self.get_count_tokens_headers(auth_header) # Use LiteLLM's async httpx client async_client: Final = get_async_httpx_client(llm_provider=litellm.LlmProviders.ANTHROPIC) diff --git a/litellm/llms/anthropic/count_tokens/token_counter.py b/litellm/llms/anthropic/count_tokens/token_counter.py index d5b8667a2e1..920916726f2 100644 --- a/litellm/llms/anthropic/count_tokens/token_counter.py +++ b/litellm/llms/anthropic/count_tokens/token_counter.py @@ -47,7 +47,6 @@ class AnthropicTokenCounter(BaseTokenCounter): TokenCountResponse with token count, or None if counting fails """ from litellm.llms.anthropic.common_utils import AnthropicError, AnthropicModelInfo - from litellm.llms.anthropic.wif import aget_anthropic_wif_token if not messages: return None @@ -55,23 +54,22 @@ class AnthropicTokenCounter(BaseTokenCounter): deployment = deployment or {} litellm_params: Final = deployment.get("litellm_params", {}) api_base: Final = litellm_params.get("api_base") - static_key: Final = AnthropicModelInfo.get_api_key(litellm_params.get("api_key")) - auth_token_configured: Final = AnthropicModelInfo.get_auth_token() is not None try: - api_key: Final = ( - static_key - if static_key or auth_token_configured - else await aget_anthropic_wif_token(litellm_params, api_base, model_to_use) + auth_header: Final = await AnthropicModelInfo.aget_auth_header( + api_key=litellm_params.get("api_key"), + api_base=api_base, + litellm_params=litellm_params, + allow_workload_identity=True, ) - if not api_key: + if auth_header is None: verbose_logger.warning("No Anthropic credential found for token counting") return None result: Final = await anthropic_count_tokens_handler.handle_count_tokens_request( model=model_to_use, messages=messages, - api_key=api_key, + auth_header=auth_header, api_base=api_base, tools=tools, system=system, diff --git a/litellm/llms/anthropic/count_tokens/transformation.py b/litellm/llms/anthropic/count_tokens/transformation.py index acff7fb417b..1e97d56913a 100644 --- a/litellm/llms/anthropic/count_tokens/transformation.py +++ b/litellm/llms/anthropic/count_tokens/transformation.py @@ -11,6 +11,7 @@ from typing import Final from pydantic import JsonValue, TypeAdapter from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION +from litellm.llms.anthropic.common_utils import merge_anthropic_beta_headers from litellm.llms.anthropic.wif import resolve_anthropic_base _COUNT_REQUEST: Final = TypeAdapter(dict[str, JsonValue]) @@ -72,28 +73,19 @@ class AnthropicCountTokensConfig: ) ) - def get_required_headers(self, api_key: str) -> dict[str, str]: - """ - Get the required headers for the CountTokens API. - - Args: - api_key: The Anthropic API key - - Returns: - Dictionary of required headers - """ - from litellm.llms.anthropic.common_utils import ( - optionally_handle_anthropic_oauth, - ) - - headers: dict[str, str] = { + def get_count_tokens_headers(self, auth_header: Mapping[str, str]) -> dict[str, str]: + """The count-tokens headers around a resolved Anthropic auth header + (``AnthropicModelInfo.get_auth_header``): x-api-key for a static key, an Authorization + bearer for ``ANTHROPIC_AUTH_TOKEN`` and for sk-ant-oat tokens, whose mandatory oauth beta + merges with the token-counting beta instead of replacing it.""" + return { "Content-Type": "application/json", - "x-api-key": api_key, "anthropic-version": "2023-06-01", - "anthropic-beta": ANTHROPIC_TOKEN_COUNTING_BETA_VERSION, + **auth_header, + "anthropic-beta": merge_anthropic_beta_headers( + auth_header.get("anthropic-beta"), ANTHROPIC_TOKEN_COUNTING_BETA_VERSION + ), } - headers, _ = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key) - return headers def validate_request( self, diff --git a/litellm/llms/anthropic/prompt_cache_prediction.py b/litellm/llms/anthropic/prompt_cache_prediction.py index 84310ce4ec3..ff304e845f7 100644 --- a/litellm/llms/anthropic/prompt_cache_prediction.py +++ b/litellm/llms/anthropic/prompt_cache_prediction.py @@ -524,6 +524,9 @@ async def count_prompt_tokens( body: Mapping[str, JsonValue], api_base: str | None = None, ) -> int | None: + auth_header: Final = AnthropicModelInfo.get_auth_header(api_key=api_key, api_base=api_base) + if auth_header is None: + return None try: native: Final = _CountBody.model_validate(body) result: Final = _CountResult.model_validate( @@ -532,7 +535,7 @@ async def count_prompt_tokens( messages=_count_objects(native.messages), tools=_count_objects(native.tools) if native.tools is not None else None, system=_JSON_OBJECT.validate_python(MappingProxyType({"system": native.system}))["system"], - api_key=api_key, + auth_header=auth_header, api_base=api_base, optional_params=_JSON_OBJECT.validate_python( MappingProxyType({key: body[key] for key in COUNT_TOKEN_OPTION_NAMES if key in body}) diff --git a/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py b/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py index 5b6b4251e4d..6a7ec13ec4c 100644 --- a/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py +++ b/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py @@ -156,7 +156,7 @@ async def test_handler_posts_to_count_tokens_path_under_deployment_api_base(http result = await AnthropicCountTokensHandler().handle_count_tokens_request( model="claude-sonnet-4-5", messages=[{"role": "user", "content": "hi"}], - api_key="sk-ant-api03-test-key", + auth_header={"x-api-key": "sk-ant-api03-test-key"}, api_base="https://gateway.example", ) diff --git a/tests/unit/llms/anthropic/test_count_tokens_oauth.py b/tests/unit/llms/anthropic/test_count_tokens_oauth.py index 542f8f55b59..96d909a4b3f 100644 --- a/tests/unit/llms/anthropic/test_count_tokens_oauth.py +++ b/tests/unit/llms/anthropic/test_count_tokens_oauth.py @@ -1,85 +1,109 @@ """ -Tests for Anthropic CountTokens API OAuth token handling. +Tests for the credential every Anthropic count-tokens request carries. -Verifies that get_required_headers() correctly handles OAuth tokens -(sk-ant-oat*) by delegating to optionally_handle_anthropic_oauth(). +The count-tokens handler receives the auth header that ``AnthropicModelInfo.get_auth_header`` +resolved, so a static key, an OAuth token (sk-ant-oat*), ``ANTHROPIC_AUTH_TOKEN`` and a minted +workload-identity token all reach Anthropic exactly the way chat on the same deployment does. -Regression test for https://github.com/BerriAI/litellm/issues/22040 +Regression tests for https://github.com/BerriAI/litellm/issues/22040 and for the +``ANTHROPIC_AUTH_TOKEN`` gap where count-tokens skipped minting but forwarded no credential. """ import os import sys +import httpx import pytest +import respx sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../.."))) +import litellm +from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION +from litellm.llms.anthropic.common_utils import AnthropicModelInfo from litellm.llms.anthropic.count_tokens.transformation import ( AnthropicCountTokensConfig, ) +from litellm.types.llms.anthropic import ANTHROPIC_OAUTH_BETA_HEADER # Fake tokens for testing (not real secrets) FAKE_OAUTH_TOKEN = "sk-ant-oat01-fake-token-for-testing-123456789abcdef" FAKE_REGULAR_KEY = "sk-ant-api03-regular-key-for-testing-123456789" +FEDERATED_DEPLOYMENT = { + "litellm_params": { + "model": "anthropic/claude-sonnet-4-5", + "anthropic_federation_rule_id": "fdrl_x", + "anthropic_organization_id": "org-x", + } +} + + +def count_tokens_headers_for(api_key: str) -> dict[str, str]: + auth_header = AnthropicModelInfo.get_auth_header(api_key=api_key) + assert auth_header is not None + return AnthropicCountTokensConfig().get_count_tokens_headers(auth_header) + + +@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() + class TestCountTokensOAuthHeaders: """Tests that count_tokens headers are correct for both regular and OAuth keys.""" def test_regular_api_key_uses_x_api_key(self): """Regular API keys should be sent via x-api-key header.""" - config = AnthropicCountTokensConfig() - headers = config.get_required_headers(FAKE_REGULAR_KEY) + headers = count_tokens_headers_for(FAKE_REGULAR_KEY) assert headers["x-api-key"] == FAKE_REGULAR_KEY assert "authorization" not in headers def test_oauth_key_uses_bearer_authorization(self): """OAuth tokens (sk-ant-oat*) should be sent via Authorization: Bearer.""" - config = AnthropicCountTokensConfig() - headers = config.get_required_headers(FAKE_OAUTH_TOKEN) + headers = count_tokens_headers_for(FAKE_OAUTH_TOKEN) assert headers.get("authorization") == f"Bearer {FAKE_OAUTH_TOKEN}" assert "x-api-key" not in headers def test_oauth_key_sets_oauth_beta_header(self): """OAuth tokens should trigger the anthropic-beta oauth header.""" - config = AnthropicCountTokensConfig() - headers = config.get_required_headers(FAKE_OAUTH_TOKEN) + headers = count_tokens_headers_for(FAKE_OAUTH_TOKEN) - assert "oauth-2025-04-20" in headers.get("anthropic-beta", "") + assert ANTHROPIC_OAUTH_BETA_HEADER in headers.get("anthropic-beta", "").split(",") def test_regular_key_preserves_token_counting_beta(self): """Regular keys should keep the token-counting beta header.""" - config = AnthropicCountTokensConfig() - headers = config.get_required_headers(FAKE_REGULAR_KEY) + headers = count_tokens_headers_for(FAKE_REGULAR_KEY) - assert "token-counting" in headers.get("anthropic-beta", "") + assert headers.get("anthropic-beta") == ANTHROPIC_TOKEN_COUNTING_BETA_VERSION def test_headers_always_have_content_type(self): """Both regular and OAuth paths should have Content-Type.""" - config = AnthropicCountTokensConfig() - for key in [FAKE_REGULAR_KEY, FAKE_OAUTH_TOKEN]: - headers = config.get_required_headers(key) + headers = count_tokens_headers_for(key) assert headers["Content-Type"] == "application/json" def test_headers_always_have_anthropic_version(self): """Both paths should have anthropic-version.""" - config = AnthropicCountTokensConfig() - for key in [FAKE_REGULAR_KEY, FAKE_OAUTH_TOKEN]: - headers = config.get_required_headers(key) + headers = count_tokens_headers_for(key) assert headers["anthropic-version"] == "2023-06-01" def test_oauth_key_preserves_token_counting_beta(self): """OAuth tokens must preserve the token-counting beta alongside the OAuth beta.""" - config = AnthropicCountTokensConfig() - headers = config.get_required_headers(FAKE_OAUTH_TOKEN) + headers = count_tokens_headers_for(FAKE_OAUTH_TOKEN) - beta_value = headers.get("anthropic-beta", "") - assert "token-counting" in beta_value, f"token-counting beta missing from OAuth headers: {beta_value}" - assert "oauth-2025-04-20" in beta_value, f"oauth beta missing from OAuth headers: {beta_value}" + betas = headers.get("anthropic-beta", "").split(",") + assert ANTHROPIC_TOKEN_COUNTING_BETA_VERSION in betas, f"token-counting beta missing: {betas}" + assert ANTHROPIC_OAUTH_BETA_HEADER in betas, f"oauth beta missing: {betas}" class TestCountTokensUsesWorkloadIdentity: @@ -92,13 +116,13 @@ class TestCountTokensUsesWorkloadIdentity: from litellm.llms.anthropic.count_tokens import token_counter as token_counter_module monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) + monkeypatch.delenv("ANTHROPIC_AUTH_TOKEN", raising=False) minted = "sk-ant-oat01-minted-for-count" async def fake_mint(_params, _api_base, _model): return minted - monkeypatch.setattr(token_counter_module, "aget_anthropic_wif_token", fake_mint, raising=False) - monkeypatch.setattr("litellm.llms.anthropic.wif.aget_anthropic_wif_token", fake_mint, raising=False) + monkeypatch.setattr("litellm.llms.anthropic.common_utils.aget_anthropic_wif_token", fake_mint) seen: dict[str, object] = {} @@ -117,54 +141,58 @@ class TestCountTokensUsesWorkloadIdentity: model_to_use="claude-sonnet-4-5", messages=[{"role": "user", "content": "hi"}], contents=None, - deployment={ - "litellm_params": { - "model": "anthropic/claude-sonnet-4-5", - "anthropic_federation_rule_id": "fdrl_x", - "anthropic_organization_id": "org-x", - } - }, + deployment=FEDERATED_DEPLOYMENT, request_model="claude-sonnet-4-5", ) assert result is not None assert result.total_tokens == 42 - assert seen["api_key"] == minted + assert seen["auth_header"] == { + "authorization": f"Bearer {minted}", + "anthropic-beta": ANTHROPIC_OAUTH_BETA_HEADER, + } @pytest.mark.asyncio - async def test_an_auth_token_deployment_never_mints(self, monkeypatch): + async def test_an_auth_token_deployment_counts_with_a_bearer_and_never_mints( + self, monkeypatch, httpx_transport_clients + ): + """With only ``ANTHROPIC_AUTH_TOKEN`` set, chat on a federated deployment authenticates with + that token, so count-tokens must send the same Bearer instead of silently returning None.""" from litellm.llms.anthropic.count_tokens import token_counter as token_counter_module - monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) + for name in ("ANTHROPIC_API_KEY", "ANTHROPIC_API_BASE", "ANTHROPIC_BASE_URL"): + monkeypatch.delenv(name, raising=False) monkeypatch.setenv("ANTHROPIC_AUTH_TOKEN", "bearer-token-for-testing") - mint_calls: list[str] = [] - async def fake_mint(_params, _api_base, model): - mint_calls.append(model) - return "sk-ant-oat01-should-not-be-minted" + async def fake_mint(_params, _api_base, _model): + raise AssertionError("an auth-token deployment must never mint a federated token") - monkeypatch.setattr("litellm.llms.anthropic.wif.aget_anthropic_wif_token", fake_mint, raising=False) + monkeypatch.setattr("litellm.llms.anthropic.common_utils.aget_anthropic_wif_token", fake_mint) - result = await token_counter_module.AnthropicTokenCounter().count_tokens( - model_to_use="claude-sonnet-4-5", - messages=[{"role": "user", "content": "hi"}], - contents=None, - deployment={ - "litellm_params": { - "model": "anthropic/claude-sonnet-4-5", - "anthropic_federation_rule_id": "fdrl_x", - "anthropic_organization_id": "org-x", - } - }, - request_model="claude-sonnet-4-5", - ) + with respx.mock(assert_all_called=True) as router: + route = router.post("https://api.anthropic.com/v1/messages/count_tokens").mock( + return_value=httpx.Response(200, json={"input_tokens": 11}) + ) + result = await token_counter_module.AnthropicTokenCounter().count_tokens( + model_to_use="claude-sonnet-4-5", + messages=[{"role": "user", "content": "hi"}], + contents=None, + deployment=FEDERATED_DEPLOYMENT, + request_model="claude-sonnet-4-5", + ) - assert result is None - assert mint_calls == [] + assert result is not None + assert result.total_tokens == 11 + assert result.tokenizer_type == "anthropic_api" + sent = route.calls.last.request.headers + assert sent["authorization"] == "Bearer bearer-token-for-testing" + assert "x-api-key" not in sent + betas = sent["anthropic-beta"].split(",") + assert ANTHROPIC_TOKEN_COUNTING_BETA_VERSION in betas + assert ANTHROPIC_OAUTH_BETA_HEADER not in betas @pytest.mark.asyncio async def test_a_failed_mint_degrades_like_an_anthropic_error(self, monkeypatch): - import litellm from litellm.llms.anthropic.count_tokens import token_counter as token_counter_module monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) @@ -177,7 +205,7 @@ class TestCountTokensUsesWorkloadIdentity: model=model, ) - monkeypatch.setattr("litellm.llms.anthropic.wif.aget_anthropic_wif_token", failing_mint, raising=False) + monkeypatch.setattr("litellm.llms.anthropic.common_utils.aget_anthropic_wif_token", failing_mint) result = await token_counter_module.AnthropicTokenCounter().count_tokens( model_to_use="claude-sonnet-4-5", @@ -212,14 +240,10 @@ class TestCountTokensUsesWorkloadIdentity: monkeypatch.setattr("litellm.secret_managers.main.get_secret_str", vault_only, raising=False) - mint_calls: list[str] = [] + async def fake_mint(_params, _api_base, _model): + raise AssertionError("a static key must never mint a federated token") - async def fake_mint(_params, _api_base, model): - mint_calls.append(model) - return "sk-ant-oat01-should-not-be-minted" - - monkeypatch.setattr(token_counter_module, "aget_anthropic_wif_token", fake_mint, raising=False) - monkeypatch.setattr("litellm.llms.anthropic.wif.aget_anthropic_wif_token", fake_mint, raising=False) + monkeypatch.setattr("litellm.llms.anthropic.common_utils.aget_anthropic_wif_token", fake_mint) seen: dict[str, object] = {} @@ -238,17 +262,10 @@ class TestCountTokensUsesWorkloadIdentity: model_to_use="claude-sonnet-4-5", messages=[{"role": "user", "content": "hi"}], contents=None, - deployment={ - "litellm_params": { - "model": "anthropic/claude-sonnet-4-5", - "anthropic_federation_rule_id": "fdrl_x", - "anthropic_organization_id": "org-x", - } - }, + deployment=FEDERATED_DEPLOYMENT, request_model="claude-sonnet-4-5", ) assert result is not None assert result.total_tokens == 7 - assert seen["api_key"] == vault_key - assert mint_calls == [] + assert seen["auth_header"] == {"x-api-key": vault_key}