From fce9d8f9033b706a4ee49e308313b2556a62d75c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 3 Oct 2026 16:24:33 -0700 Subject: [PATCH] fix(anthropic): count tokens with ANTHROPIC_AUTH_TOKEN through the shared auth header Count-tokens walked its own credential ladder: a static key, else skip minting when ANTHROPIC_AUTH_TOKEN is set, else mint a federated token. With only the auth token set it forwarded nothing and the proxy silently fell back to its local tokenizer while chat on the same deployment authenticated with that token. The handler now takes the auth header that AnthropicModelInfo.aget_auth_header resolves, the same ladder chat, files, batches and skills use, and merges the oauth beta a minted or consumer token carries with the token-counting beta --- .../llms/anthropic/count_tokens/handler.py | 6 +- .../anthropic/count_tokens/token_counter.py | 16 +- .../anthropic/count_tokens/transformation.py | 30 ++-- .../llms/anthropic/prompt_cache_prediction.py | 5 +- ...t_anthropic_count_tokens_transformation.py | 2 +- .../llms/anthropic/test_count_tokens_oauth.py | 167 ++++++++++-------- 6 files changed, 118 insertions(+), 108 deletions(-) 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}