feat(anthropic): workload identity federation via a shared RFC 7523 token exchange engine

Adds a provider-agnostic JWT-bearer token exchange engine (litellm/llms/base_llm/auth/)
with two-tier refresh, single-flight minting, negative caching, response caps, and
RFC 6749 error redaction, plus the Anthropic WIF adapter and wiring: env and
litellm_params config, sync and async facades, and beta-header merging fixes on the
skills, files, batches, messages, and passthrough surfaces

WIF is the lowest credential tier, so a deployment that does not configure it behaves
exactly as before. Verified end to end against the live token endpoint from a container:
both /chat/completions and /v1/messages complete over a minted sk-ant-oat01 token with no
api_key configured anywhere

Two protocol details were confirmed against the live endpoint rather than the docs. The
exchange needs no anthropic-beta header, so none is sent. service_account_id mints fine
when omitted for a single-service-account rule, so it stays optional

Resolves #28607
This commit is contained in:
derhornspieler 2026-08-22 17:34:18 -04:00
parent f005afa146
commit 61a1122420
30 changed files with 4186 additions and 728 deletions

View file

@ -57,7 +57,7 @@
"limit": 5663
},
"reportMissingTypeArgument": {
"limit": 15555
"limit": 15553
},
"reportMissingTypeStubs": {
"limit": 40
@ -99,19 +99,19 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44655
"limit": 44653
},
"reportUnknownLambdaType": {
"limit": 109
},
"reportUnknownMemberType": {
"limit": 39011
"limit": 39003
},
"reportUnknownParameterType": {
"limit": 19885
"limit": 19883
},
"reportUnknownVariableType": {
"limit": 30569
"limit": 30557
},
"reportUnnecessaryCast": {
"limit": 117

View file

@ -24,10 +24,24 @@ AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset(
# The per-deployment Rust opt-in.
RUST_KWARG_KEY: Final = "rust"
# Anthropic workload identity federation config, read from litellm_params by the
# Anthropic auth tier. Registered like `rust`: here so the kwargs funnel carries
# them, and in `all_litellm_params` so they never leak into the provider body.
ANTHROPIC_WIF_KWARGS_KEYS: Final = frozenset(
{
"anthropic_federation_rule_id",
"anthropic_organization_id",
"anthropic_service_account_id",
"anthropic_workspace_id",
"anthropic_identity_token_file",
"anthropic_identity_token",
}
)
# Keys `completion()` forwards from its own kwargs into `get_litellm_params`,
# which are otherwise invisible to it because that call site passes explicit
# named arguments rather than `**kwargs`.
FORWARDED_KWARGS_KEYS: Final = AWS_CREDENTIAL_KWARGS_KEYS | frozenset({RUST_KWARG_KEY})
FORWARDED_KWARGS_KEYS: Final = AWS_CREDENTIAL_KWARGS_KEYS | ANTHROPIC_WIF_KWARGS_KEYS | frozenset({RUST_KWARG_KEY})
# Pre-define optional kwargs keys as frozenset for O(1) lookups
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
@ -62,6 +76,7 @@ OPTIONAL_KWARGS_KEYS: Final = (
}
)
| AWS_CREDENTIAL_KWARGS_KEYS
| ANTHROPIC_WIF_KWARGS_KEYS
)
# Backward-compatible alias for existing imports/tests.

View file

@ -11,6 +11,8 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.llms.openai import AllMessageValues, CreateBatchRequest
from litellm.types.utils import LiteLLMBatch, LlmProviders, ModelResponse
from ..common_utils import merge_anthropic_beta_headers
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -43,23 +45,26 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
api_base: str | None = None,
) -> dict:
"""Validate and prepare environment-specific headers and parameters."""
if api_base is None and isinstance(litellm_params, dict):
api_base = litellm_params.get("api_base")
auth_header: Final = self.anthropic_model_info.get_auth_header(api_key, api_base)
params_mapping: Final = litellm_params if isinstance(litellm_params, dict) else None
if api_base is None and params_mapping is not None:
api_base = params_mapping.get("api_base")
auth_header: Final = self.anthropic_model_info.get_auth_header(api_key, api_base, litellm_params=params_mapping)
if auth_header is None:
raise ValueError(
"Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params"
)
merged_beta: Final = merge_anthropic_beta_headers(
merge_anthropic_beta_headers(headers.get("anthropic-beta"), auth_header.get("anthropic-beta")),
"message-batches-2024-09-24",
)
_headers: Final = {
"accept": "application/json",
"anthropic-version": "2023-06-01",
"content-type": "application/json",
}
_headers.update(auth_header)
# Add beta header for message batches
if "anthropic-beta" not in headers:
headers["anthropic-beta"] = "message-batches-2024-09-24"
headers.update(_headers)
headers["anthropic-beta"] = merged_beta
return headers
def get_complete_batch_url(

View file

@ -20,6 +20,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
from litellm.litellm_core_utils.prompt_templates.factory import (
THOUGHT_SIGNATURE_SEPARATOR,
)
from litellm.llms.anthropic.wif import aget_anthropic_wif_token, get_anthropic_wif_token
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.llms.anthropic import (
@ -56,6 +57,9 @@ def _strip_bedrock_id_suffixes(model: str) -> str:
)
_SERVER_OWNED_AUTH_HEADERS: Final = frozenset({"x-api-key", "authorization"})
def is_anthropic_oauth_key(value: str | None) -> bool:
"""Check if a value contains an Anthropic OAuth token (sk-ant-oat*)."""
if value is None:
@ -65,12 +69,9 @@ def is_anthropic_oauth_key(value: str | None) -> bool:
return value.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX)
def _merge_beta_headers(existing: str | None, new_beta: str) -> str:
"""Merge a new beta value into an existing comma-separated anthropic-beta header."""
if not existing:
return new_beta
betas: Final = {b.strip() for b in existing.split(",") if b.strip()}
betas.add(new_beta)
def merge_anthropic_beta_headers(existing: str | None, new_beta: str | None) -> str:
"""Merge comma-separated anthropic-beta header values, deduplicated and sorted."""
betas: Final = {b.strip() for value in (existing, new_beta) if value for b in value.split(",") if b.strip()}
return ",".join(sorted(betas))
@ -93,14 +94,18 @@ def optionally_handle_anthropic_oauth(headers: dict, api_key: str | None) -> tup
if auth_header and auth_header.startswith(f"Bearer {ANTHROPIC_OAUTH_TOKEN_PREFIX}"):
api_key = auth_header.replace("Bearer ", "")
headers.pop("x-api-key", None)
headers["anthropic-beta"] = _merge_beta_headers(headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER)
headers["anthropic-beta"] = merge_anthropic_beta_headers(
headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER
)
headers["anthropic-dangerous-direct-browser-access"] = "true"
return headers, api_key
# Check api_key directly (standard chat/completion flow)
if api_key and api_key.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX):
headers.pop("x-api-key", None)
headers["authorization"] = f"Bearer {api_key}"
headers["anthropic-beta"] = _merge_beta_headers(headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER)
headers["anthropic-beta"] = merge_anthropic_beta_headers(
headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER
)
headers["anthropic-dangerous-direct-browser-access"] = "true"
return headers, api_key
@ -596,7 +601,9 @@ class AnthropicModelInfo(BaseLLMModelInfo):
return list(set(betas))
@staticmethod
def _make_api_key_auth_header(api_key: str, api_base: str | None, use_bearer_for_custom_base: bool = False) -> dict:
def _make_api_key_auth_header(
api_key: str, api_base: str | None, use_bearer_for_custom_base: bool = False
) -> Mapping[str, str]:
if use_bearer_for_custom_base and (
api_base and "api.anthropic.com" not in api_base and not api_key.startswith("sk-ant-")
):
@ -625,6 +632,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
container_with_skills_used: bool = False,
api_base: str | None = None,
use_bearer_for_custom_base: bool = False,
wif_minted: bool = False,
) -> dict:
betas: Final = set()
# Anthropic no longer requires the prompt-caching beta header
@ -668,7 +676,8 @@ class AnthropicModelInfo(BaseLLMModelInfo):
}
if _is_oauth:
headers["authorization"] = f"Bearer {api_key}"
headers["anthropic-dangerous-direct-browser-access"] = "true"
if not wif_minted:
headers["anthropic-dangerous-direct-browser-access"] = "true"
betas.add(ANTHROPIC_OAUTH_BETA_HEADER)
elif auth_token and not api_key:
headers["authorization"] = f"Bearer {auth_token}"
@ -700,10 +709,11 @@ class AnthropicModelInfo(BaseLLMModelInfo):
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
if api_base is None and isinstance(litellm_params, dict):
api_base = litellm_params.get("api_base")
params_mapping: Final = litellm_params if isinstance(litellm_params, dict) else None
if api_base is None and params_mapping is not None:
api_base = params_mapping.get("api_base")
use_bearer_for_custom_base: Final[bool] = bool(
isinstance(litellm_params, dict) and litellm_params.get("use_bearer_for_custom_base", False)
params_mapping is not None and params_mapping.get("use_bearer_for_custom_base", False)
)
# Check for Anthropic OAuth token in headers
headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key)
@ -712,9 +722,21 @@ class AnthropicModelInfo(BaseLLMModelInfo):
auth_token: str | None = None
if api_key is None:
auth_token = AnthropicModelInfo.get_auth_token()
if api_key is None and auth_token is None:
wif_token: Final = (
get_anthropic_wif_token(params_mapping, api_base, model) if api_key is None and auth_token is None else None
)
wif_minted: Final = wif_token is not None
resolved_api_key: Final = wif_token if wif_token is not None else api_key
if resolved_api_key is None and auth_token is None:
raise litellm.AuthenticationError(
message="Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` or `ANTHROPIC_AUTH_TOKEN` in your environment vars",
message=(
"Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the "
"environment variables or via params. Please set `ANTHROPIC_API_KEY` or `ANTHROPIC_AUTH_TOKEN` "
"in your environment vars, or configure workload identity federation via "
"`ANTHROPIC_FEDERATION_RULE_ID`, `ANTHROPIC_ORGANIZATION_ID`, "
"`ANTHROPIC_SERVICE_ACCOUNT_ID` and "
"`ANTHROPIC_IDENTITY_TOKEN_FILE` (or `ANTHROPIC_IDENTITY_TOKEN`)"
),
llm_provider="anthropic",
model=model,
)
@ -739,7 +761,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
computer_tool_used=computer_tool_used,
prompt_caching_set=prompt_caching_set,
pdf_used=pdf_used,
api_key=api_key,
api_key=resolved_api_key,
auth_token=auth_token,
file_id_used=file_id_used,
web_search_tool_used=web_search_tool_used,
@ -754,11 +776,16 @@ class AnthropicModelInfo(BaseLLMModelInfo):
container_with_skills_used=container_with_skills_used,
api_base=api_base,
use_bearer_for_custom_base=use_bearer_for_custom_base,
wif_minted=wif_minted,
)
headers = {**headers, **anthropic_headers}
caller_headers: Final = (
{name: value for name, value in headers.items() if name.lower() not in _SERVER_OWNED_AUTH_HEADERS}
if wif_minted
else headers
)
return headers
return {**caller_headers, **anthropic_headers}
@staticmethod
def get_api_base(api_base: str | None = None) -> str | None:
@ -793,23 +820,62 @@ class AnthropicModelInfo(BaseLLMModelInfo):
api_key: str | None = None,
api_base: str | None = None,
use_bearer_for_custom_base: bool = False,
) -> dict | None:
litellm_params: Mapping[str, object] | None = None,
) -> Mapping[str, str] | None:
"""Resolve Anthropic credentials and return the appropriate auth header dict.
Checks ANTHROPIC_API_KEY first (-> x-api-key or Bearer depending on
use_bearer_for_custom_base), then ANTHROPIC_AUTH_TOKEN (-> Authorization: Bearer).
Returns None if neither is available.
use_bearer_for_custom_base), then ANTHROPIC_AUTH_TOKEN (-> Authorization: Bearer),
then workload identity federation (-> Authorization: Bearer with a minted
sk-ant-oat01 token, honoring anthropic_* litellm_params when provided). Every
Bearer built from an sk-ant-oat token carries the mandatory oauth anthropic-beta.
Returns None if no credential source is available.
"""
static_header: Final = AnthropicModelInfo._static_auth_header(api_key, api_base, use_bearer_for_custom_base)
if static_header is not None:
return static_header
wif_token: Final = get_anthropic_wif_token(litellm_params, api_base, "")
if wif_token is not None:
return AnthropicModelInfo._oauth_bearer_header(wif_token)
return None
@staticmethod
async def aget_auth_header(
api_key: str | None = None,
api_base: str | None = None,
use_bearer_for_custom_base: bool = False,
litellm_params: Mapping[str, object] | None = None,
) -> Mapping[str, str] | None:
"""Async counterpart of get_auth_header: the WIF tier can block on a token
exchange POST, so async callers await it off the event loop."""
static_header: Final = AnthropicModelInfo._static_auth_header(api_key, api_base, use_bearer_for_custom_base)
if static_header is not None:
return static_header
wif_token: Final = await aget_anthropic_wif_token(litellm_params, api_base, "")
if wif_token is not None:
return AnthropicModelInfo._oauth_bearer_header(wif_token)
return None
@staticmethod
def _static_auth_header(
api_key: str | None,
api_base: str | None,
use_bearer_for_custom_base: bool,
) -> Mapping[str, str] | None:
resolved_key: Final = AnthropicModelInfo.get_api_key(api_key)
if resolved_key is not None:
if is_anthropic_oauth_key(resolved_key):
return {"authorization": f"Bearer {resolved_key}"}
return AnthropicModelInfo._oauth_bearer_header(resolved_key)
return AnthropicModelInfo._make_api_key_auth_header(resolved_key, api_base, use_bearer_for_custom_base)
auth_token: Final = AnthropicModelInfo.get_auth_token()
if auth_token is not None:
return {"authorization": f"Bearer {auth_token}"}
return None
@staticmethod
def _oauth_bearer_header(token: str) -> Mapping[str, str]:
return {"authorization": f"Bearer {token}", "anthropic-beta": ANTHROPIC_OAUTH_BETA_HEADER}
@staticmethod
def get_base_model(model: str | None = None) -> str | None:
return model.replace("anthropic/", "") if model else None
@ -819,7 +885,10 @@ class AnthropicModelInfo(BaseLLMModelInfo):
auth_header: Final = AnthropicModelInfo.get_auth_header(api_key, api_base)
if api_base is None or auth_header is None:
raise ValueError(
"ANTHROPIC_API_BASE/ANTHROPIC_BASE_URL or ANTHROPIC_API_KEY/ANTHROPIC_AUTH_TOKEN is not set. Please set the environment variable, to query Anthropic's `/models` endpoint."
"ANTHROPIC_API_BASE/ANTHROPIC_BASE_URL or ANTHROPIC_API_KEY/ANTHROPIC_AUTH_TOKEN (or workload "
"identity federation via ANTHROPIC_FEDERATION_RULE_ID/ANTHROPIC_ORGANIZATION_ID/"
"ANTHROPIC_IDENTITY_TOKEN_FILE) is not set. Please set the environment variable, to query "
"Anthropic's `/models` endpoint."
)
headers: Final = {"anthropic-version": "2023-06-01"}
headers.update(auth_header)

View file

@ -27,6 +27,7 @@ from litellm.types.router import GenericLiteLLMParams
from ...common_utils import (
AnthropicError,
AnthropicModelInfo,
merge_anthropic_beta_headers,
optionally_handle_anthropic_oauth,
strip_advisor_blocks_from_messages,
)
@ -308,21 +309,80 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key)
if "x-api-key" not in headers and "authorization" not in headers:
auth_header: Final = AnthropicModelInfo.get_auth_header(api_key)
if auth_header is not None:
headers.update(auth_header)
self._apply_env_auth_header(
headers,
AnthropicModelInfo.get_auth_header(
api_key,
api_base=api_base,
litellm_params=self._wif_litellm_params(litellm_params),
),
)
return self._finalize_messages_headers(headers, optional_params), api_base
async def avalidate_anthropic_messages_environment(
self,
headers: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
model: str,
messages: list[Any], # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
optional_params: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
litellm_params: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
api_key: str | None = None,
api_base: str | None = None,
) -> tuple[dict, str | None]: # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
if type(self).validate_anthropic_messages_environment is not (
AnthropicMessagesConfig.validate_anthropic_messages_environment
):
# a subclass sync override must keep winning on the async path
return self.validate_anthropic_messages_environment(
headers=headers,
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
api_key=api_key,
api_base=api_base,
)
oauth_headers, oauth_api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key)
if "x-api-key" not in oauth_headers and "authorization" not in oauth_headers:
self._apply_env_auth_header(
oauth_headers,
await AnthropicModelInfo.aget_auth_header(
oauth_api_key,
api_base=api_base,
litellm_params=self._wif_litellm_params(litellm_params),
),
)
return self._finalize_messages_headers(oauth_headers, optional_params), api_base
@staticmethod
def _apply_env_auth_header(headers: dict, auth_header: Mapping[str, str] | None) -> None: # mutable-ok: out-param
if auth_header is None:
return
merged_beta: Final = merge_anthropic_beta_headers(
headers.get("anthropic-beta"), auth_header.get("anthropic-beta")
)
headers.update(auth_header)
if merged_beta:
headers["anthropic-beta"] = merged_beta
def _wif_litellm_params(self, litellm_params: dict) -> Mapping[str, object] | None: # mutable-ok: sync-contract mirror
"""Subclasses reuse this validate step for their own /v1/messages-compatible providers,
so an Anthropic federation token is only ever minted for Anthropic itself."""
if self._resolved_provider != "anthropic":
return None
return litellm_params if isinstance(litellm_params, dict) else None
def _finalize_messages_headers(self, headers: dict, optional_params: dict) -> dict: # mutable-ok: out-param
if "anthropic-version" not in headers:
headers["anthropic-version"] = DEFAULT_ANTHROPIC_API_VERSION
if "content-type" not in headers:
headers["content-type"] = "application/json"
headers = self._update_headers_with_anthropic_beta(
return self._update_headers_with_anthropic_beta(
headers=headers,
optional_params=optional_params,
)
return headers, api_base
@staticmethod
def _translate_reasoning_effort_to_anthropic(model: str, optional_params: dict, custom_llm_provider: str) -> None:
"""Map OpenAI-style ``reasoning_effort`` to native Anthropic params.

View file

@ -85,7 +85,7 @@ class AnthropicFilesHandler:
# Get Anthropic API credentials
api_base = self.anthropic_model_info.get_api_base(api_base)
auth_header: Final = self.anthropic_model_info.get_auth_header(api_key, api_base)
auth_header: Final = await self.anthropic_model_info.aget_auth_header(api_key, api_base)
if auth_header is None:
raise ValueError("Missing Anthropic API Key")

View file

@ -35,7 +35,7 @@ from litellm.types.llms.openai import (
)
from litellm.types.utils import LlmProviders
from ..common_utils import AnthropicError, AnthropicModelInfo
from ..common_utils import AnthropicError, AnthropicModelInfo, merge_anthropic_beta_headers
ANTHROPIC_FILES_API_BASE: Final = "https://api.anthropic.com"
ANTHROPIC_FILES_BETA_HEADER: Final = "files-api-2025-04-14"
@ -94,18 +94,23 @@ class AnthropicFilesConfig(BaseFilesConfig):
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
if api_base is None and isinstance(litellm_params, dict):
api_base = litellm_params.get("api_base")
auth_header: Final = AnthropicModelInfo.get_auth_header(api_key, api_base)
params_mapping: Final = litellm_params if isinstance(litellm_params, dict) else None
if api_base is None and params_mapping is not None:
api_base = params_mapping.get("api_base")
auth_header: Final = AnthropicModelInfo.get_auth_header(api_key, api_base, litellm_params=params_mapping)
if auth_header is None:
raise ValueError(
"Anthropic API key is required. Set ANTHROPIC_API_KEY or ANTHROPIC_AUTH_TOKEN environment variable or pass api_key parameter."
)
merged_beta: Final = merge_anthropic_beta_headers(
merge_anthropic_beta_headers(headers.get("anthropic-beta"), auth_header.get("anthropic-beta")),
ANTHROPIC_FILES_BETA_HEADER,
)
headers.update(
{
**auth_header,
"anthropic-version": "2023-06-01",
"anthropic-beta": ANTHROPIC_FILES_BETA_HEADER,
"anthropic-beta": merged_beta,
}
)
return headers

View file

@ -2,6 +2,7 @@
Anthropic Skills API configuration and transformations
"""
from types import MappingProxyType
from typing import Any, Final
import httpx
@ -32,37 +33,27 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig):
def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict:
"""Add Anthropic-specific headers"""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.constants import ANTHROPIC_SKILLS_API_BETA_VERSION
from litellm.llms.anthropic.common_utils import (
AnthropicModelInfo,
merge_anthropic_beta_headers,
)
# Get API key from litellm_params if available
api_key = None
api_base = None
if litellm_params is not None:
api_key = litellm_params.api_key
api_base = litellm_params.api_base
auth_header: Final = AnthropicModelInfo.get_auth_header(api_key, api_base)
auth_header: Final = AnthropicModelInfo.get_auth_header(
api_key=litellm_params.api_key if litellm_params is not None else None,
api_base=litellm_params.api_base if litellm_params is not None else None,
litellm_params=MappingProxyType(dict(litellm_params)) if litellm_params is not None else None,
)
if auth_header is None:
raise ValueError("ANTHROPIC_API_KEY or ANTHROPIC_AUTH_TOKEN is required for Skills API")
merged_beta: Final = merge_anthropic_beta_headers(
merge_anthropic_beta_headers(headers.get("anthropic-beta"), auth_header.get("anthropic-beta")),
ANTHROPIC_SKILLS_API_BETA_VERSION,
)
headers.update(auth_header)
headers["anthropic-version"] = "2023-06-01"
# Add beta header for skills API
from litellm.constants import ANTHROPIC_SKILLS_API_BETA_VERSION
if "anthropic-beta" not in headers:
headers["anthropic-beta"] = ANTHROPIC_SKILLS_API_BETA_VERSION
elif isinstance(headers["anthropic-beta"], list):
if ANTHROPIC_SKILLS_API_BETA_VERSION not in headers["anthropic-beta"]:
headers["anthropic-beta"].append(ANTHROPIC_SKILLS_API_BETA_VERSION)
elif isinstance(headers["anthropic-beta"], str):
if ANTHROPIC_SKILLS_API_BETA_VERSION not in headers["anthropic-beta"]:
headers["anthropic-beta"] = [
headers["anthropic-beta"],
ANTHROPIC_SKILLS_API_BETA_VERSION,
]
headers["anthropic-beta"] = merged_beta
headers["content-type"] = "application/json"
return headers

View file

@ -0,0 +1,231 @@
"""Anthropic workload identity federation: exchanges an external OIDC identity
token for a short-lived ``sk-ant-oat01`` token via the shared RFC 7523 engine."""
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final, NoReturn
from pydantic import BaseModel, ConfigDict
from typing_extensions import assert_never
import litellm
from litellm.llms.base_llm.auth.token_exchange import (
JwtBearerTokenExchangeEngine,
default_token_exchange_engine,
)
from litellm.llms.base_llm.auth.types import (
AssertionSourceError,
ExchangeError,
ExchangeResult,
InsecureTokenUrl,
MalformedTokenResponse,
MintedToken,
TokenEndpointError,
TokenExchangeSpec,
TokenTransportError,
)
from litellm.types.llms.anthropic import ANTHROPIC_TOKEN_EXCHANGE_PATH
_JWT_BEARER_GRANT_TYPE: Final = "urn:ietf:params:oauth:grant-type:jwt-bearer"
_DEFAULT_API_BASE: Final = "https://api.anthropic.com"
_INLINE_ENV_VAR: Final = "ANTHROPIC_IDENTITY_TOKEN"
_ACCEPTED_REF_PREFIX: Final = "oidc/"
_REJECTED_REF_PREFIX: Final = "oidc/env_path/"
_WORKSPACE_HINT: Final = (
" If the federation rule is scoped to a workspace, set ANTHROPIC_WORKSPACE_ID"
" (or the anthropic_workspace_id litellm param) to that workspace id."
)
_ALLOWLIST_HINT: Final = (
" Identity token files must sit under an allowed credential directory"
" (/var/run/secrets or /run/secrets by default);"
" set LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS to extend the allowlist."
)
class AnthropicWifParams(BaseModel):
model_config = ConfigDict(frozen=True)
federation_rule_id: str
organization_id: str
service_account_id: str | None = None
workspace_id: str | None = None
assertion_ref: str
def resolve_anthropic_wif_params(litellm_params: Mapping[str, object] | None) -> AnthropicWifParams | None:
federation_rule_id: Final = _config_value(
litellm_params, "anthropic_federation_rule_id", "ANTHROPIC_FEDERATION_RULE_ID"
)
organization_id: Final = _config_value(litellm_params, "anthropic_organization_id", "ANTHROPIC_ORGANIZATION_ID")
if federation_rule_id is None or organization_id is None:
return None
assertion_ref: Final = _resolve_assertion_ref(litellm_params)
if assertion_ref is None:
return None
return AnthropicWifParams(
federation_rule_id=federation_rule_id,
organization_id=organization_id,
service_account_id=_config_value(
litellm_params, "anthropic_service_account_id", "ANTHROPIC_SERVICE_ACCOUNT_ID"
),
workspace_id=_config_value(litellm_params, "anthropic_workspace_id", "ANTHROPIC_WORKSPACE_ID"),
assertion_ref=assertion_ref,
)
def build_anthropic_wif_spec(params: AnthropicWifParams, api_base: str) -> TokenExchangeSpec:
return TokenExchangeSpec(
token_url=api_base.rstrip("/") + ANTHROPIC_TOKEN_EXCHANGE_PATH,
assertion_ref=params.assertion_ref,
assertion_field="assertion",
static_body=MappingProxyType(
{
name: value
for name, value in (
("grant_type", _JWT_BEARER_GRANT_TYPE),
("federation_rule_id", params.federation_rule_id),
("organization_id", params.organization_id),
("service_account_id", params.service_account_id),
("workspace_id", params.workspace_id),
)
if value is not None
}
),
body_encoding="json",
request_headers=MappingProxyType({}),
cache_key_identity=(
params.federation_rule_id,
params.organization_id,
params.service_account_id or "",
params.workspace_id or "",
),
)
def get_anthropic_wif_token(
litellm_params: Mapping[str, object] | None,
api_base: str | None,
model: str,
engine: JwtBearerTokenExchangeEngine = default_token_exchange_engine,
) -> str | None:
params: Final = resolve_anthropic_wif_params(litellm_params)
if params is None:
return None
result: Final = engine.get_token(build_anthropic_wif_spec(params, _token_exchange_base(api_base)))
return _token_from_result(result, model, params)
async def aget_anthropic_wif_token(
litellm_params: Mapping[str, object] | None,
api_base: str | None,
model: str,
engine: JwtBearerTokenExchangeEngine = default_token_exchange_engine,
) -> str | None:
params: Final = resolve_anthropic_wif_params(litellm_params)
if params is None:
return None
result: Final = await engine.aget_token(build_anthropic_wif_spec(params, _token_exchange_base(api_base)))
return _token_from_result(result, model, params)
def _token_from_result(result: ExchangeResult, model: str, params: AnthropicWifParams) -> str:
match result:
case MintedToken():
return result.access_token.get_secret_value()
case _:
_raise_anthropic_wif_error(result, model=model, workspace_id_set=params.workspace_id is not None)
def _token_exchange_base(api_base: str | None) -> str:
"""Exchange base for any caller-supplied form of the deployment base: trailing
slashes and chat-appended ``/v1/messages`` suffixes stripped, so every tier
derives the same token URL (and cache key) for the same deployment."""
return _strip_chat_suffix(api_base if api_base is not None else _resolve_default_api_base())
def _resolve_default_api_base() -> str:
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
return AnthropicModelInfo.get_api_base(None) or _DEFAULT_API_BASE
def _strip_chat_suffix(base: str) -> str:
trimmed: Final = base.rstrip("/")
stripped: Final = trimmed.removesuffix("/v1/messages")
return stripped if stripped == trimmed else _strip_chat_suffix(stripped)
def _config_value(litellm_params: Mapping[str, object] | None, param_key: str, env_name: str) -> str | None:
return _param_str(litellm_params, param_key) or _env_str(env_name)
def _param_str(litellm_params: Mapping[str, object] | None, key: str) -> str | None:
if litellm_params is None:
return None
value: Final = litellm_params.get(key)
return value if isinstance(value, str) and value else None
def _env_str(name: str) -> str | None:
from litellm.secret_managers.main import get_secret_str
value: Final = get_secret_str(name)
return value if isinstance(value, str) and value else None
def _resolve_assertion_ref(litellm_params: Mapping[str, object] | None) -> str | None:
file_param: Final = _param_str(litellm_params, "anthropic_identity_token_file")
if file_param is not None:
return f"oidc/file/{file_param}"
inline_param: Final = _param_str(litellm_params, "anthropic_identity_token")
if inline_param is not None:
return _validated_inline_ref(inline_param)
file_env: Final = _env_str("ANTHROPIC_IDENTITY_TOKEN_FILE")
if file_env is not None:
return f"oidc/file/{file_env}"
if _env_str(_INLINE_ENV_VAR) is not None:
return f"oidc/env/{_INLINE_ENV_VAR}"
return None
def _validated_inline_ref(value: str) -> str:
if value.startswith(_ACCEPTED_REF_PREFIX) and not value.startswith(_REJECTED_REF_PREFIX):
return value
raise litellm.AuthenticationError(
message=(
"anthropic_identity_token must be an oidc/ secret reference such as oidc/env/VAR_NAME,"
" oidc/file//absolute/path, oidc/github/<audience>, or oidc/google/<audience>."
" Raw identity tokens and oidc/env_path/ references are not accepted;"
" to pass a token directly, export it and reference it as oidc/env/VAR_NAME"
),
llm_provider="anthropic",
model="",
)
def _raise_anthropic_wif_error(error: ExchangeError, model: str, workspace_id_set: bool) -> NoReturn:
raise litellm.AuthenticationError(
message=f"Anthropic workload identity federation failed. {_error_detail(error, workspace_id_set)}",
llm_provider="anthropic",
model=model,
)
def _error_detail(error: ExchangeError, workspace_id_set: bool) -> str:
match error:
case AssertionSourceError() if error.kind == "disallowed_path":
return f"Could not read the OIDC identity token from {error.source_ref}.{_ALLOWLIST_HINT}"
case AssertionSourceError():
return f"Could not obtain the OIDC identity token ({error.kind}) from {error.source_ref}."
case InsecureTokenUrl():
return f"The token endpoint must use https; refusing to send the identity token to host {error.host!r}."
case TokenEndpointError() if error.status_code == 401 and not workspace_id_set:
return f"The token endpoint returned HTTP 401: {error.redacted_body}{_WORKSPACE_HINT}"
case TokenEndpointError():
return f"The token endpoint returned HTTP {error.status_code}: {error.redacted_body}"
case TokenTransportError():
return f"Could not reach the token endpoint: {error.detail}."
case MalformedTokenResponse():
return f"The token endpoint returned an unusable response: {error.detail}."
case _:
assert_never(error)

View file

@ -41,6 +41,29 @@ class BaseAnthropicMessagesConfig(ABC):
"""
return headers, api_base
async def avalidate_anthropic_messages_environment(
self,
headers: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
model: str,
messages: list[Any], # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
optional_params: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
litellm_params: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
api_key: str | None = None,
api_base: str | None = None,
) -> tuple[dict, str | None]: # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
"""Async counterpart used by the async handler. The default delegates to the
sync implementation; providers whose sync path can block the event loop
(e.g. a WIF token exchange) override this."""
return self.validate_anthropic_messages_environment(
headers=headers,
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
api_key=api_key,
api_base=api_base,
)
@abstractmethod
def get_complete_url(
self,

View file

@ -0,0 +1,49 @@
from litellm.llms.base_llm.auth.token_exchange import (
ADVISORY_REFRESH_BACKOFF_SECONDS,
ADVISORY_REFRESH_SECONDS,
MANDATORY_REFRESH_SECONDS,
MAX_ASSERTION_BYTES,
MAX_RESPONSE_BYTES,
JwtBearerTokenExchangeEngine,
default_token_exchange_engine,
redact_oauth_error_body,
validate_token_endpoint_url,
)
from litellm.llms.base_llm.auth.types import (
AssertionReader,
AssertionSourceError,
BodyEncoding,
ExchangeError,
ExchangeResult,
InsecureTokenUrl,
MalformedTokenResponse,
MintedToken,
SyncTokenPoster,
TokenEndpointError,
TokenExchangeSpec,
TokenTransportError,
)
__all__ = (
"ADVISORY_REFRESH_BACKOFF_SECONDS",
"ADVISORY_REFRESH_SECONDS",
"MANDATORY_REFRESH_SECONDS",
"MAX_ASSERTION_BYTES",
"MAX_RESPONSE_BYTES",
"AssertionReader",
"AssertionSourceError",
"BodyEncoding",
"ExchangeError",
"ExchangeResult",
"InsecureTokenUrl",
"JwtBearerTokenExchangeEngine",
"MalformedTokenResponse",
"MintedToken",
"SyncTokenPoster",
"TokenEndpointError",
"TokenExchangeSpec",
"TokenTransportError",
"default_token_exchange_engine",
"redact_oauth_error_body",
"validate_token_endpoint_url",
)

View file

@ -0,0 +1,508 @@
"""RFC 7523 JWT-bearer token exchange engine, shared across providers.
One sync state machine per process: bounded engine-owned entry map, two-tier
refresh (advisory background refresh + mandatory single-flight), HTTPS pinning,
response caps, and RFC 6749 5.2 redaction. Providers describe a grant profile as
a ``TokenExchangeSpec`` and map the typed ``ExchangeError`` union to their own
public exception contract.
"""
import asyncio
import hashlib
import json
import threading
import time
from collections.abc import Callable, Mapping
from concurrent.futures import Executor, ThreadPoolExecutor
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, TypeAlias
from urllib.parse import urlencode, urlsplit
import httpx
from pydantic import BaseModel, SecretStr, TypeAdapter, ValidationError
from typing_extensions import assert_never
from litellm._logging import verbose_logger
from litellm.llms.base_llm.auth.types import (
AssertionReader,
AssertionSourceError,
ExchangeError,
ExchangeResult,
InsecureTokenUrl,
MalformedTokenResponse,
MintedToken,
SyncTokenPoster,
TokenEndpointError,
TokenExchangeSpec,
TokenTransportError,
)
if TYPE_CHECKING:
from litellm.llms.custom_httpx.http_handler import HTTPHandler
ADVISORY_REFRESH_SECONDS: Final = 120.0
MANDATORY_REFRESH_SECONDS: Final = 30.0
ADVISORY_REFRESH_BACKOFF_SECONDS: Final = 5.0
FALLBACK_TOKEN_TTL_SECONDS: Final = 60.0
MAX_ASSERTION_BYTES: Final = 16 * 1024
MAX_RESPONSE_BYTES: Final = 1024 * 1024
_REDACTION_CAP: Final = 256
_FOLLOWER_WAIT_GRACE_SECONDS: Final = 5.0
_LOCAL_HOSTS: Final = frozenset({"localhost", "127.0.0.1", "::1"})
_OAUTH_ERROR_FIELDS: Final = ("error", "error_description", "error_uri")
_NESTED_ERROR_FIELDS: Final = ("type", "message")
_CONTENT_TYPES: Final = MappingProxyType({"json": "application/json", "form": "application/x-www-form-urlencoded"})
_OVERSIZED_BODY_MESSAGE: Final = "oversized error response omitted"
_NON_OBJECT_BODY_MESSAGE: Final = "non-object error response omitted"
_NO_OAUTH_FIELDS_MESSAGE: Final = "error response carried no RFC 6749 fields"
class _TokenExchangeResponse(BaseModel):
access_token: str
expires_in: int | None = None
token_type: str | None = None
_RedactableBody: TypeAlias = Mapping[str, object] | list[object] | str | int | float | bool | None
_REDACTABLE_BODY_ADAPTER: Final = TypeAdapter[_RedactableBody](_RedactableBody)
def validate_token_endpoint_url(url: str) -> str | InsecureTokenUrl:
parsed: Final = urlsplit(url)
if parsed.scheme == "https":
return url
if parsed.scheme == "http" and (parsed.hostname or "") in _LOCAL_HOSTS:
return url
return InsecureTokenUrl(host=parsed.hostname or "")
def redact_oauth_error_body(status_code: int, body_text: str) -> TokenEndpointError:
return TokenEndpointError(status_code=status_code, redacted_body=_redact_body_text(body_text))
def _redact_body_text(body_text: str) -> str:
if len(body_text) > MAX_RESPONSE_BYTES:
return _OVERSIZED_BODY_MESSAGE
try:
parsed: Final = _REDACTABLE_BODY_ADAPTER.validate_json(body_text)
except ValidationError:
return body_text[:_REDACTION_CAP]
match parsed:
case str():
return parsed[:_REDACTION_CAP]
case Mapping():
return _format_oauth_error_fields(parsed)
case _:
return _NON_OBJECT_BODY_MESSAGE
def _format_oauth_error_fields(body: Mapping[str, object]) -> str:
fields: Final = tuple(
f"{name}: {_format_oauth_error_value(value)}"
for name in _OAUTH_ERROR_FIELDS
for value in (body.get(name),)
if value is not None
)
return "; ".join(fields) if fields else _NO_OAUTH_FIELDS_MESSAGE
def _format_oauth_error_value(value: object) -> str:
"""RFC 6749 types ``error`` as a string, but Anthropic (and other providers) nest their
own ``{"type": ..., "message": ...}`` envelope there; render that rather than a dict repr."""
if isinstance(value, Mapping):
nested: Final = tuple(
f"{str(part)[:_REDACTION_CAP]}"
for key in _NESTED_ERROR_FIELDS
for part in (value.get(key),)
if part is not None
)
if nested:
return " - ".join(nested)
return str(value)[:_REDACTION_CAP]
def _error_summary(error: ExchangeError) -> str:
match error:
case AssertionSourceError():
return f"AssertionSourceError: assertion {error.kind} from {error.source_ref}"
case InsecureTokenUrl():
return f"InsecureTokenUrl: insecure token endpoint host {error.host}"
case TokenEndpointError():
return f"TokenEndpointError: HTTP {error.status_code}: {error.redacted_body}"
case TokenTransportError():
return f"TokenTransportError: {error.detail}"
case MalformedTokenResponse():
return f"MalformedTokenResponse: {error.detail}"
case _:
assert_never(error)
def _cache_key(spec: TokenExchangeSpec) -> str:
return hashlib.sha256(
"\x1f".join((spec.token_url, spec.assertion_ref, *spec.cache_key_identity)).encode()
).hexdigest()
def _read_assertion(reader: AssertionReader, ref: str) -> SecretStr | AssertionSourceError:
from litellm.secret_managers.main import OidcPathNotAllowedError
try:
raw: Final = reader(ref)
except OidcPathNotAllowedError:
return AssertionSourceError(kind="disallowed_path", source_ref=ref)
except ValueError:
return AssertionSourceError(kind="unreadable", source_ref=ref)
except Exception: # noqa: BLE001 # injected readers (secret managers) raise arbitrarily; all failures become values
return AssertionSourceError(kind="unreadable", source_ref=ref)
if raw is None:
return AssertionSourceError(kind="missing", source_ref=ref)
stripped: Final = raw.strip()
if not stripped:
return AssertionSourceError(kind="empty", source_ref=ref)
if len(stripped.encode("utf-8")) > MAX_ASSERTION_BYTES:
return AssertionSourceError(kind="oversized", source_ref=ref)
return SecretStr(stripped)
def _serialize_body(spec: TokenExchangeSpec, assertion: SecretStr) -> bytes:
if spec.body_encoding == "json":
return json.dumps(
{ # mutable-ok: transient body dict consumed inline by the serializer
**spec.static_body,
spec.assertion_field: assertion.get_secret_value(),
}
).encode()
return urlencode(
{ # mutable-ok: transient body dict consumed inline by the serializer
**spec.static_body,
spec.assertion_field: assertion.get_secret_value(),
}
).encode()
def _sanitize_expires_in(expires_in: int | None) -> float:
if expires_in is None or expires_in <= 0:
return FALLBACK_TOKEN_TTL_SECONDS
return float(expires_in)
def _capped_body_text(response: httpx.Response) -> str:
if len(response.content) > MAX_RESPONSE_BYTES:
return _OVERSIZED_BODY_MESSAGE
return response.text
def _default_assertion_reader(ref: str) -> str | None:
from litellm.secret_managers.main import get_secret_str
return get_secret_str(ref)
class _HttpxSyncTokenPoster:
"""Default poster: a dedicated HTTPHandler (no logging_obj, so litellm's
pre/post-call body logging never sees the exchange POST); returns the
response for any status."""
def __init__(self) -> None:
self._lock: Final = threading.Lock()
self._handler: HTTPHandler | None = None
def _handler_instance(self) -> "HTTPHandler":
from litellm.llms.custom_httpx.http_handler import HTTPHandler
with self._lock:
if self._handler is None:
self._handler = HTTPHandler(timeout=httpx.Timeout(timeout=30.0, connect=5.0))
return self._handler
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
try:
response: Final[httpx.Response | None] = self._handler_instance().post( # pyright: ignore[reportUnknownMemberType] # HTTPHandler.post is legacy-untyped; the result is validated below
url,
content=content,
headers=dict(headers), # mutable-ok: HTTPHandler.post requires a concrete dict
timeout=timeout,
)
except httpx.HTTPStatusError as e:
return e.response
if response is None:
raise httpx.TransportError("token endpoint returned no response")
return response
class _Entry:
"""Single-flight state for one cache key; mutable by design, confined to the
engine, and only ever mutated under the engine lock."""
__slots__ = ("backoff_until", "done", "force_refresh", "in_flight", "last_error", "token")
def __init__(self, force_refresh: bool = False) -> None:
self.token: MintedToken | None = None
self.in_flight: bool = False
self.done: Final = threading.Event()
self.backoff_until: float = float("-inf")
self.force_refresh: bool = force_refresh
self.last_error: ExchangeError | None = None
def arm(self) -> None:
self.in_flight = True
self.last_error = None
self.done.clear()
def publish(self, result: ExchangeResult, backoff_until: float) -> None:
match result:
case MintedToken():
self.token = result
self.last_error = None
case _:
self.last_error = result
self.backoff_until = backoff_until
self.force_refresh = False
self.in_flight = False
self.done.set()
def publish_advisory(self, result: ExchangeResult, backoff_until: float) -> None:
"""A failed advisory refresh records only the backoff, never ``last_error``: a follower whose
cached token expires while this runs must be free to re-lead a fresh mint and recover."""
match result:
case MintedToken():
self.token = result
self.last_error = None
case _:
self.backoff_until = backoff_until
self.in_flight = False
self.done.set()
@dataclass(frozen=True, slots=True)
class _Serve:
token: MintedToken
@dataclass(frozen=True, slots=True)
class _ServeAndRefresh:
token: MintedToken
@dataclass(frozen=True, slots=True)
class _Lead:
pass
@dataclass(frozen=True, slots=True)
class _Follow:
pass
@dataclass(frozen=True, slots=True)
class _Fail:
error: ExchangeError
_Decision: TypeAlias = _Serve | _ServeAndRefresh | _Lead | _Follow | _Fail
@dataclass(frozen=True, slots=True)
class _Unauthorized:
response: httpx.Response
class JwtBearerTokenExchangeEngine:
def __init__(
self,
poster: SyncTokenPoster | None = None,
assertion_reader: AssertionReader | None = None,
clock: Callable[[], float] = time.monotonic,
refresh_executor: Executor | None = None,
max_entries: int = 64,
) -> None:
self._poster: Final[SyncTokenPoster] = poster if poster is not None else _HttpxSyncTokenPoster()
self._assertion_reader: Final[AssertionReader] = (
assertion_reader if assertion_reader is not None else _default_assertion_reader
)
self._clock: Final = clock
self._refresh_executor: Executor | None = refresh_executor
self._max_entries: Final = max_entries
self._lock: Final = threading.Lock()
self._entries: Final[dict[str, _Entry]] = {} # mutable-ok: engine-owned map guarded by _lock
def get_token(self, spec: TokenExchangeSpec) -> ExchangeResult:
with self._lock:
entry: Final = self._get_or_create_entry_locked(spec)
decision: Final = self._classify_and_arm_locked(entry)
match decision:
case _Serve(token=token):
return token
case _ServeAndRefresh(token=token):
self._executor_instance().submit(self._advisory_refresh, spec, entry)
return token
case _Fail(error=error):
return error
case _Lead():
return self._lead(spec, entry)
case _Follow():
followed: Final = self._await_leader(spec, entry)
return followed if followed is not None else self.get_token(spec)
case _:
assert_never(decision)
async def aget_token(self, spec: TokenExchangeSpec) -> ExchangeResult:
return await asyncio.to_thread(self.get_token, spec)
def invalidate(self, spec: TokenExchangeSpec) -> None:
key: Final = _cache_key(spec)
with self._lock:
if key in self._entries:
self._entries[key] = _Entry(force_refresh=True)
def _get_or_create_entry_locked(self, spec: TokenExchangeSpec) -> _Entry:
key: Final = _cache_key(spec)
existing: Final = self._entries.get(key)
if existing is not None:
return existing
if len(self._entries) >= self._max_entries:
self._evict_locked()
created: Final = _Entry()
self._entries[key] = created
return created
def _evict_locked(self) -> None:
now: Final = self._clock()
stale: Final = tuple(
key
for key, entry in self._entries.items()
if not entry.in_flight
and (entry.token is None or (entry.token.expires_at is not None and entry.token.expires_at <= now))
)
for key in stale:
del self._entries[key]
if len(self._entries) < self._max_entries:
return
evictable: Final = tuple(
(entry.token.expires_at, key)
for key, entry in self._entries.items()
if not entry.in_flight and entry.token is not None and entry.token.expires_at is not None
)
if evictable:
del self._entries[min(evictable)[1]]
def _classify_and_arm_locked(self, entry: _Entry) -> _Decision:
token: Final = entry.token
if token is not None and not entry.force_refresh:
if token.expires_at is None:
return _Serve(token=token)
remaining: Final = token.expires_at - self._clock()
if remaining > ADVISORY_REFRESH_SECONDS:
return _Serve(token=token)
if remaining > MANDATORY_REFRESH_SECONDS:
if entry.in_flight or self._clock() < entry.backoff_until:
return _Serve(token=token)
entry.arm()
return _ServeAndRefresh(token=token)
if entry.in_flight:
return _Follow()
if entry.last_error is not None and self._clock() < entry.backoff_until:
return _Fail(error=entry.last_error)
entry.arm()
return _Lead()
def _executor_instance(self) -> Executor:
with self._lock:
if self._refresh_executor is None:
self._refresh_executor = ThreadPoolExecutor(
max_workers=2, thread_name_prefix="litellm-token-exchange-refresh"
)
return self._refresh_executor
def _lead(self, spec: TokenExchangeSpec, entry: _Entry) -> ExchangeResult:
result: Final = self._exchange(spec)
with self._lock:
entry.publish(result, backoff_until=self._clock() + ADVISORY_REFRESH_BACKOFF_SECONDS)
return result
def _await_leader(self, spec: TokenExchangeSpec, entry: _Entry) -> "ExchangeResult | None":
"""None means the finished round left neither a valid token nor an error
(a failed advisory refresh); the caller re-enters and leads a fresh exchange."""
leader_finished: Final = entry.done.wait(2 * spec.timeout_seconds + _FOLLOWER_WAIT_GRACE_SECONDS)
with self._lock:
token: Final = entry.token
if token is not None and (token.expires_at is None or token.expires_at > self._clock()):
return token
if entry.last_error is not None:
return entry.last_error
if leader_finished:
return None
return TokenTransportError(detail="timed out waiting for the token exchange leader")
def _advisory_refresh(self, spec: TokenExchangeSpec, entry: _Entry) -> None:
result: Final = self._exchange(spec)
with self._lock:
entry.publish_advisory(result, backoff_until=self._clock() + ADVISORY_REFRESH_BACKOFF_SECONDS)
stale_expires_at: Final = entry.token.expires_at if entry.token is not None else None
if isinstance(result, MintedToken):
return
seconds_to_mandatory_wall: Final = (
max(stale_expires_at - self._clock() - MANDATORY_REFRESH_SECONDS, 0.0)
if stale_expires_at is not None
else 0.0
)
verbose_logger.warning(
"Advisory token refresh against %s failed (%s); serving the cached token for up to "
"%.0fs before the mandatory refresh wall; next attempt after %.0fs backoff",
urlsplit(spec.token_url).hostname or "",
_error_summary(result),
seconds_to_mandatory_wall,
ADVISORY_REFRESH_BACKOFF_SECONDS,
)
def _exchange(self, spec: TokenExchangeSpec) -> ExchangeResult:
first: Final = self._attempt_exchange(spec)
if not isinstance(first, _Unauthorized):
return first
second: Final = self._attempt_exchange(spec)
if isinstance(second, _Unauthorized):
return redact_oauth_error_body(second.response.status_code, _capped_body_text(second.response))
return second
def _attempt_exchange(self, spec: TokenExchangeSpec) -> "ExchangeResult | _Unauthorized":
assertion: Final = _read_assertion(self._assertion_reader, spec.assertion_ref)
if isinstance(assertion, AssertionSourceError):
return assertion
url_check: Final = validate_token_endpoint_url(spec.token_url)
if isinstance(url_check, InsecureTokenUrl):
return url_check
try:
response: Final = self._poster.post(
spec.token_url,
content=_serialize_body(spec, assertion),
headers=MappingProxyType({"content-type": _CONTENT_TYPES[spec.body_encoding], **spec.request_headers}),
timeout=spec.timeout_seconds,
)
except Exception as e: # noqa: BLE001 # injected posters may raise beyond httpx; transport failures become values
return TokenTransportError(detail=f"{type(e).__name__}: {e}"[:_REDACTION_CAP])
if response.status_code == 401:
return _Unauthorized(response=response)
return self._parse_response(response)
def _parse_response(self, response: httpx.Response) -> ExchangeResult:
if not 200 <= response.status_code < 300:
return redact_oauth_error_body(response.status_code, _capped_body_text(response))
if len(response.content) > MAX_RESPONSE_BYTES:
return MalformedTokenResponse(detail="token response body exceeds the 1 MiB cap")
try:
parsed: Final = _TokenExchangeResponse.model_validate_json(response.content)
except ValidationError:
return MalformedTokenResponse(detail="token response failed RFC 6749 5.1 schema validation")
if parsed.token_type is not None and parsed.token_type.lower() != "bearer":
return MalformedTokenResponse(detail="token response carried a non-bearer token_type")
if not parsed.access_token.strip():
return MalformedTokenResponse(detail="token response carried an empty access_token")
return MintedToken(
access_token=SecretStr(parsed.access_token),
expires_at=self._clock() + _sanitize_expires_in(parsed.expires_in),
)
default_token_exchange_engine: Final = JwtBearerTokenExchangeEngine()

View file

@ -0,0 +1,76 @@
"""Provider-agnostic types for the RFC 7523 JWT-bearer token exchange engine."""
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from typing import Literal, Protocol, TypeAlias
import httpx
from pydantic import SecretStr
BodyEncoding: TypeAlias = Literal["json", "form"]
@dataclass(frozen=True, slots=True)
class TokenExchangeSpec:
"""One grant profile as pure data: one instance per (provider, deployment, identity).
``token_url`` must be derived from deployment config/env only, never per-request caller
input. ``assertion_ref`` is a ``oidc/...`` get_secret ref resolved fresh on every exchange.
"""
token_url: str
assertion_ref: str
assertion_field: str
static_body: Mapping[str, str]
body_encoding: BodyEncoding
request_headers: Mapping[str, str]
cache_key_identity: tuple[str, ...]
timeout_seconds: float = 30.0
@dataclass(frozen=True, slots=True)
class MintedToken:
access_token: SecretStr
expires_at: float | None
@dataclass(frozen=True, slots=True)
class AssertionSourceError:
kind: Literal["missing", "empty", "oversized", "unreadable", "disallowed_path"]
source_ref: str
@dataclass(frozen=True, slots=True)
class InsecureTokenUrl:
host: str
@dataclass(frozen=True, slots=True)
class TokenEndpointError:
status_code: int
redacted_body: str
@dataclass(frozen=True, slots=True)
class TokenTransportError:
detail: str
@dataclass(frozen=True, slots=True)
class MalformedTokenResponse:
detail: str
ExchangeError: TypeAlias = (
AssertionSourceError | InsecureTokenUrl | TokenEndpointError | TokenTransportError | MalformedTokenResponse
)
ExchangeResult: TypeAlias = MintedToken | ExchangeError
class SyncTokenPoster(Protocol):
"""Returns the response for ANY status; never raises for status."""
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: ...
AssertionReader: TypeAlias = Callable[[str], str | None] # mutable-ok: Callable param-list syntax, not a list

View file

@ -2082,7 +2082,7 @@ class BaseLLMHTTPHandler:
(
headers,
api_base,
) = anthropic_messages_provider_config.validate_anthropic_messages_environment(
) = await anthropic_messages_provider_config.avalidate_anthropic_messages_environment(
headers=merged_headers or {},
model=model,
messages=messages,

View file

@ -9,7 +9,7 @@ Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc.
import json
import os
import re
from collections.abc import Callable
from collections.abc import Callable, Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Any, Final, cast
@ -25,7 +25,7 @@ from litellm.constants import (
ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS,
BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES,
)
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.anthropic.common_utils import AnthropicModelInfo, merge_anthropic_beta_headers
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import *
from litellm.proxy.auth.route_checks import RouteChecks
@ -568,6 +568,17 @@ async def is_streaming_request_fn(request: Request) -> bool:
return False
def _anthropic_passthrough_headers(auth_header: Mapping[str, str] | None, client_beta: str | None) -> Mapping[str, str]:
"""Custom headers take priority over forwarded client headers, so merge the
client's anthropic-beta into the auth header's instead of clobbering it."""
if auth_header is None:
return MappingProxyType({})
auth_beta: Final = auth_header.get("anthropic-beta")
if auth_beta is None or client_beta is None:
return auth_header
return MappingProxyType({**auth_header, "anthropic-beta": merge_anthropic_beta_headers(client_beta, auth_beta)})
@router.api_route(
"/anthropic/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
@ -606,11 +617,11 @@ async def anthropic_proxy_route(
is_streaming_request: Final = await is_streaming_request_fn(request)
## CREATE PASS-THROUGH
auth_header: Final = AnthropicModelInfo.get_auth_header(anthropic_api_key or None)
auth_header: Final = await AnthropicModelInfo.aget_auth_header(anthropic_api_key or None)
endpoint_func: Final = create_pass_through_route(
endpoint=endpoint,
target=str(updated_url),
custom_headers=auth_header if auth_header is not None else {},
custom_headers=_anthropic_passthrough_headers(auth_header, request.headers.get("anthropic-beta")),
_forward_headers=True,
is_streaming_request=is_streaming_request,
) # dynamically construct pass-through endpoint based on incoming path

View file

@ -49,6 +49,10 @@ def _oidc_token_cache_ttl(oidc_token: str, max_ttl: int) -> int:
_DEFAULT_OIDC_ALLOWED_CREDENTIAL_DIRS: Final = ("/var/run/secrets", "/run/secrets")
class OidcPathNotAllowedError(ValueError):
"""An ``oidc/file/`` path was rejected by the credential-directory allowlist."""
def _get_oidc_allowed_credential_dirs() -> list[str]:
"""
Return the absolute, normalized list of directories from which
@ -73,7 +77,7 @@ def _resolve_oidc_file_path(requested_path: str) -> str:
credential directories. Raises ``ValueError`` otherwise.
"""
if not os.path.isabs(requested_path):
raise ValueError(
raise OidcPathNotAllowedError(
"oidc/file path must be absolute. Use the format "
"'oidc/file//var/run/secrets/<name>' (note the leading slash "
"after 'oidc/file/')."
@ -87,7 +91,7 @@ def _resolve_oidc_file_path(requested_path: str) -> str:
# commonpath raises when paths are on different drives (Windows);
# treat as not-matching and continue.
continue
raise ValueError(
raise OidcPathNotAllowedError(
"oidc/file path is outside the allowed credential directories. "
"Set LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS to extend the allowlist."
)

View file

@ -721,5 +721,6 @@ ANTHROPIC_EFFORT_BETA_HEADER: Final = "effort-2025-11-24"
# OAuth constants
ANTHROPIC_OAUTH_TOKEN_PREFIX: Final = "sk-ant-oat"
ANTHROPIC_OAUTH_BETA_HEADER: Final = "oauth-2025-04-20"
ANTHROPIC_TOKEN_EXCHANGE_PATH: Final = "/v1/oauth/token"
ANTHROPIC_PROMPT_CACHING_SCOPE_BETA_HEADER: Final = "prompt-caching-scope-2026-01-05"

View file

@ -3476,9 +3476,21 @@ bedrock_batch_litellm_params: Final = (
"bedrock_tags",
)
# Anthropic workload identity federation config, read from litellm_params by the
# Anthropic auth tier. Listed for the same reason as the fields above: an
# unrecognized top-level key is swept into extra_body and sent to /v1/messages.
anthropic_wif_litellm_params: Final = (
"anthropic_federation_rule_id",
"anthropic_organization_id",
"anthropic_service_account_id",
"anthropic_workspace_id",
"anthropic_identity_token_file",
"anthropic_identity_token",
)
all_litellm_params = (
agentic_loop_internal_litellm_params
+ [TRUSTED_CALLBACK_VARS_FIELD, *bedrock_batch_litellm_params]
+ [TRUSTED_CALLBACK_VARS_FIELD, *bedrock_batch_litellm_params, *anthropic_wif_litellm_params]
+ [
"metadata",
"litellm_metadata",

View file

@ -244,3 +244,40 @@ class TestRustOptIn:
from litellm.types.utils import all_litellm_params
assert "rust" in all_litellm_params
class TestAnthropicWifKeys:
"""The six anthropic_* WIF keys need the same dual registration as `rust`:
carried by the kwargs funnel into litellm_params (where the Anthropic auth
tier reads them) AND listed in all_litellm_params (so the extra_body sweep
never sends them to /v1/messages)."""
SIX_KEYS = {
"anthropic_federation_rule_id": "fdrl_1",
"anthropic_organization_id": "org-1",
"anthropic_service_account_id": "svcacct_1",
"anthropic_workspace_id": "wrkspc_1",
"anthropic_identity_token_file": "/var/run/secrets/tok",
"anthropic_identity_token": "oidc/env/TOK",
}
def test_keys_survive_into_litellm_params(self):
params = get_litellm_params(**self.SIX_KEYS)
for key, value in self.SIX_KEYS.items():
assert params[key] == value
def test_keys_are_forwarded_from_completion_kwargs(self):
from litellm.litellm_core_utils.get_litellm_params import FORWARDED_KWARGS_KEYS
assert set(self.SIX_KEYS) <= FORWARDED_KWARGS_KEYS
def test_keys_stay_out_of_the_provider_body(self):
from litellm.types.utils import all_litellm_params
for key in self.SIX_KEYS:
assert key in all_litellm_params
def test_keys_absent_when_not_configured(self):
params = get_litellm_params()
for key in self.SIX_KEYS:
assert key not in params

View file

@ -80,8 +80,11 @@ def test_validate_environment_preserves_existing_beta_header(config):
litellm_params={},
api_key="sk-ant-test",
)
# Existing beta header must NOT be overwritten.
assert headers["anthropic-beta"] == "custom-beta-value"
# Existing beta values are preserved and the batches beta is merged in.
assert set(headers["anthropic-beta"].split(",")) == {
"custom-beta-value",
"message-batches-2024-09-24",
}
def test_validate_environment_oauth_key_uses_bearer(config):

View file

@ -0,0 +1,132 @@
from pathlib import Path
from typing import Final
import pytest
import litellm
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
from litellm.llms.anthropic.wif import get_anthropic_wif_token
from litellm.llms.minimax.messages.transformation import MinimaxMessagesConfig
from litellm.llms.tencent.messages.transformation import TencentAnthropicMessagesConfig
from tests.test_litellm.llms.anthropic.test_anthropic_wif import (
ScriptedPoster,
make_engine,
token_response,
write_token_file,
)
_WIF_PARAMS: Final[dict] = {
"anthropic_federation_rule_id": "fdrl_abc123",
"anthropic_organization_id": "org-uuid-1",
"anthropic_identity_token_file": "/var/run/secrets/identity-token",
}
def test_wif_litellm_params_passthrough_for_anthropic() -> None:
assert AnthropicMessagesConfig()._wif_litellm_params(_WIF_PARAMS) == _WIF_PARAMS
def test_wif_litellm_params_blocked_for_minimax() -> None:
assert MinimaxMessagesConfig()._wif_litellm_params(_WIF_PARAMS) is None
def test_wif_litellm_params_blocked_for_tencent() -> None:
assert TencentAnthropicMessagesConfig()._wif_litellm_params(_WIF_PARAMS) is None
def test_minimax_validate_environment_never_attaches_anthropic_wif_credential(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Regression test: before the fix, an Anthropic-WIF-configured proxy would mint a real
Anthropic federation token inside MiniMax's inherited validate_anthropic_messages_environment
and send it as the Authorization header on the MiniMax-routed request."""
monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_prod")
monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-prod-uuid")
monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False)
monkeypatch.delenv("ANTHROPIC_AUTH_TOKEN", raising=False)
token_file = write_token_file(tmp_path, "jwt-assertion-value")
litellm_params = {"anthropic_identity_token_file": str(token_file)}
headers, _ = MinimaxMessagesConfig().validate_anthropic_messages_environment(
headers={},
model="MiniMax-M2.1",
messages=[],
optional_params={},
litellm_params=litellm_params,
api_key=None,
api_base="https://api.minimax.io/anthropic",
)
assert "authorization" not in headers
assert "x-api-key" not in headers
def test_tencent_validate_environment_never_attaches_anthropic_wif_credential(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_prod")
monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-prod-uuid")
monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False)
monkeypatch.delenv("ANTHROPIC_AUTH_TOKEN", raising=False)
monkeypatch.delenv("TENCENT_API_KEY", raising=False)
token_file = write_token_file(tmp_path, "jwt-assertion-value")
litellm_params = {"anthropic_identity_token_file": str(token_file)}
original_api_key: Final = litellm.api_key
litellm.api_key = None
try:
headers, _ = TencentAnthropicMessagesConfig().validate_anthropic_messages_environment(
headers={},
model="deepseek-v4-pro",
messages=[],
optional_params={},
litellm_params=litellm_params,
api_key=None,
api_base="https://tokenhub-intl.tencentcloudmaas.com",
)
finally:
litellm.api_key = original_api_key
assert "authorization" not in headers
assert "x-api-key" not in headers
def test_wif_token_exchange_reaches_only_anthropic_not_minimax_or_tencent(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
"""get_anthropic_wif_token's engine parameter is the only DI seam in the WIF minting chain;
validate_anthropic_messages_environment always uses the module's default engine, so this
drives that seam directly with the exact litellm_params AnthropicModelInfo.get_auth_header
would receive from each config, proving MiniMax/Tencent never reach the token endpoint even
when a mint would otherwise succeed."""
monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_prod")
monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-prod-uuid")
monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path))
token_file = write_token_file(tmp_path, "jwt-assertion-value")
litellm_params = {"anthropic_identity_token_file": str(token_file)}
poster = ScriptedPoster([token_response("sk-ant-oat01-canary")])
engine = make_engine(poster)
minted: Final = get_anthropic_wif_token(
AnthropicMessagesConfig()._wif_litellm_params(litellm_params),
"https://api.anthropic.com",
"claude-sonnet-4-5",
engine,
)
assert minted == "sk-ant-oat01-canary"
assert len(poster.requests) == 1
for config in (MinimaxMessagesConfig(), TencentAnthropicMessagesConfig()):
assert (
get_anthropic_wif_token(
config._wif_litellm_params(litellm_params),
"https://api.minimax.io/anthropic",
"MiniMax-M2.1",
engine,
)
is None
)
assert len(poster.requests) == 1

View file

@ -70,9 +70,7 @@ class TestAnthropicFilesHandler:
@pytest.fixture
def mock_anthropic_batch_results_canceled(self):
"""Mock Anthropic batch results with canceled status"""
return json.dumps(
{"custom_id": "test-request-3", "result": {"type": "canceled"}}
).encode("utf-8")
return json.dumps({"custom_id": "test-request-3", "result": {"type": "canceled"}}).encode("utf-8")
@pytest.fixture
def mock_anthropic_batch_results_mixed(self):
@ -114,9 +112,7 @@ class TestAnthropicFilesHandler:
return "\n".join(lines).encode("utf-8")
@pytest.mark.asyncio
async def test_afile_content_success(
self, handler, mock_anthropic_batch_results_succeeded
):
async def test_afile_content_success(self, handler, mock_anthropic_batch_results_succeeded):
"""Test successful file content retrieval and transformation"""
file_content_request: FileContentRequest = {
"file_id": "batch_123",
@ -135,16 +131,12 @@ class TestAnthropicFilesHandler:
),
)
with patch(
"litellm.llms.anthropic.files.handler.get_async_httpx_client"
) as mock_get_client:
with patch("litellm.llms.anthropic.files.handler.get_async_httpx_client") as mock_get_client:
mock_client = AsyncMock()
mock_client.get = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
with patch.object(
handler.anthropic_model_info, "get_api_key", return_value="test-api-key"
):
with patch.object(handler.anthropic_model_info, "get_api_key", return_value="test-api-key"):
with patch.object(
handler.anthropic_model_info,
"get_api_base",
@ -161,9 +153,7 @@ class TestAnthropicFilesHandler:
# Verify transformation to OpenAI format
content = result.response.content.decode("utf-8")
lines = [
line for line in content.strip().split("\n") if line.strip()
]
lines = [line for line in content.strip().split("\n") if line.strip()]
assert len(lines) == 1
transformed_result = json.loads(lines[0])
@ -172,18 +162,13 @@ class TestAnthropicFilesHandler:
assert "body" in transformed_result["response"]
# Verify body has required OpenAI format fields
assert "id" in transformed_result["response"]["body"]
assert (
transformed_result["response"]["body"]["object"]
== "chat.completion"
)
assert transformed_result["response"]["body"]["object"] == "chat.completion"
assert "choices" in transformed_result["response"]["body"]
# Verify request_id matches the original message id
assert transformed_result["response"]["request_id"] == "msg_123"
@pytest.mark.asyncio
async def test_afile_content_with_prefix(
self, handler, mock_anthropic_batch_results_succeeded
):
async def test_afile_content_with_prefix(self, handler, mock_anthropic_batch_results_succeeded):
"""Test file content retrieval with anthropic_batch_results: prefix"""
file_content_request: FileContentRequest = {
"file_id": "anthropic_batch_results:batch_123",
@ -201,16 +186,12 @@ class TestAnthropicFilesHandler:
),
)
with patch(
"litellm.llms.anthropic.files.handler.get_async_httpx_client"
) as mock_get_client:
with patch("litellm.llms.anthropic.files.handler.get_async_httpx_client") as mock_get_client:
mock_client = AsyncMock()
mock_client.get = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
with patch.object(
handler.anthropic_model_info, "get_api_key", return_value="test-api-key"
):
with patch.object(handler.anthropic_model_info, "get_api_key", return_value="test-api-key"):
with patch.object(
handler.anthropic_model_info,
"get_api_base",
@ -228,9 +209,7 @@ class TestAnthropicFilesHandler:
assert "batch_123" in call_url
@pytest.mark.asyncio
async def test_afile_content_errored_result(
self, handler, mock_anthropic_batch_results_errored
):
async def test_afile_content_errored_result(self, handler, mock_anthropic_batch_results_errored):
"""Test transformation of errored batch results"""
file_content_request: FileContentRequest = {
"file_id": "batch_123",
@ -248,16 +227,12 @@ class TestAnthropicFilesHandler:
),
)
with patch(
"litellm.llms.anthropic.files.handler.get_async_httpx_client"
) as mock_get_client:
with patch("litellm.llms.anthropic.files.handler.get_async_httpx_client") as mock_get_client:
mock_client = AsyncMock()
mock_client.get = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
with patch.object(
handler.anthropic_model_info, "get_api_key", return_value="test-api-key"
):
with patch.object(handler.anthropic_model_info, "get_api_key", return_value="test-api-key"):
with patch.object(
handler.anthropic_model_info,
"get_api_base",
@ -269,29 +244,17 @@ class TestAnthropicFilesHandler:
)
content = result.response.content.decode("utf-8")
lines = [
line for line in content.strip().split("\n") if line.strip()
]
lines = [line for line in content.strip().split("\n") if line.strip()]
assert len(lines) == 1
transformed_result = json.loads(lines[0])
assert transformed_result["custom_id"] == "test-request-2"
assert (
transformed_result["response"]["status_code"] == 400
) # invalid_request_error maps to 400
assert (
transformed_result["response"]["body"]["error"]["type"]
== "invalid_request_error"
)
assert (
transformed_result["response"]["body"]["error"]["message"]
== "Invalid request"
)
assert transformed_result["response"]["status_code"] == 400 # invalid_request_error maps to 400
assert transformed_result["response"]["body"]["error"]["type"] == "invalid_request_error"
assert transformed_result["response"]["body"]["error"]["message"] == "Invalid request"
@pytest.mark.asyncio
async def test_afile_content_canceled_result(
self, handler, mock_anthropic_batch_results_canceled
):
async def test_afile_content_canceled_result(self, handler, mock_anthropic_batch_results_canceled):
"""Test transformation of canceled batch results"""
file_content_request: FileContentRequest = {
"file_id": "batch_123",
@ -309,16 +272,12 @@ class TestAnthropicFilesHandler:
),
)
with patch(
"litellm.llms.anthropic.files.handler.get_async_httpx_client"
) as mock_get_client:
with patch("litellm.llms.anthropic.files.handler.get_async_httpx_client") as mock_get_client:
mock_client = AsyncMock()
mock_client.get = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
with patch.object(
handler.anthropic_model_info, "get_api_key", return_value="test-api-key"
):
with patch.object(handler.anthropic_model_info, "get_api_key", return_value="test-api-key"):
with patch.object(
handler.anthropic_model_info,
"get_api_base",
@ -330,23 +289,16 @@ class TestAnthropicFilesHandler:
)
content = result.response.content.decode("utf-8")
lines = [
line for line in content.strip().split("\n") if line.strip()
]
lines = [line for line in content.strip().split("\n") if line.strip()]
assert len(lines) == 1
transformed_result = json.loads(lines[0])
assert transformed_result["custom_id"] == "test-request-3"
assert transformed_result["response"]["status_code"] == 400
assert (
"Batch request was canceled"
in transformed_result["response"]["body"]["error"]["message"]
)
assert "Batch request was canceled" in transformed_result["response"]["body"]["error"]["message"]
@pytest.mark.asyncio
async def test_afile_content_mixed_results(
self, handler, mock_anthropic_batch_results_mixed
):
async def test_afile_content_mixed_results(self, handler, mock_anthropic_batch_results_mixed):
"""Test transformation of mixed batch results (succeeded, errored, expired)"""
file_content_request: FileContentRequest = {
"file_id": "batch_123",
@ -364,16 +316,12 @@ class TestAnthropicFilesHandler:
),
)
with patch(
"litellm.llms.anthropic.files.handler.get_async_httpx_client"
) as mock_get_client:
with patch("litellm.llms.anthropic.files.handler.get_async_httpx_client") as mock_get_client:
mock_client = AsyncMock()
mock_client.get = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
with patch.object(
handler.anthropic_model_info, "get_api_key", return_value="test-api-key"
):
with patch.object(handler.anthropic_model_info, "get_api_key", return_value="test-api-key"):
with patch.object(
handler.anthropic_model_info,
"get_api_base",
@ -385,9 +333,7 @@ class TestAnthropicFilesHandler:
)
content = result.response.content.decode("utf-8")
lines = [
line for line in content.strip().split("\n") if line.strip()
]
lines = [line for line in content.strip().split("\n") if line.strip()]
assert len(lines) == 3
# Check first result (succeeded)
@ -396,9 +342,7 @@ class TestAnthropicFilesHandler:
# Check second result (errored)
result2 = json.loads(lines[1])
assert (
result2["response"]["status_code"] == 429
) # rate_limit_error maps to 429
assert result2["response"]["status_code"] == 429 # rate_limit_error maps to 429
# Check third result (expired)
result3 = json.loads(lines[2])
@ -415,12 +359,12 @@ class TestAnthropicFilesHandler:
}
with patch.object(
handler.anthropic_model_info, "get_auth_header", return_value=None
handler.anthropic_model_info,
"aget_auth_header",
new=AsyncMock(return_value=None),
):
with pytest.raises(ValueError, match="Missing Anthropic API Key"):
await handler.afile_content(
file_content_request=file_content_request, api_key=None
)
await handler.afile_content(file_content_request=file_content_request, api_key=None)
@pytest.mark.asyncio
async def test_afile_content_missing_file_id(self, handler):
@ -432,9 +376,7 @@ class TestAnthropicFilesHandler:
}
with pytest.raises(ValueError, match="file_id is required"):
await handler.afile_content(
file_content_request=file_content_request, api_key="test-api-key"
)
await handler.afile_content(file_content_request=file_content_request, api_key="test-api-key")
@pytest.mark.asyncio
async def test_afile_content_http_error(self, handler):
@ -454,21 +396,15 @@ class TestAnthropicFilesHandler:
),
)
mock_response.raise_for_status = MagicMock(
side_effect=httpx.HTTPStatusError(
"Not Found", request=mock_response.request, response=mock_response
)
side_effect=httpx.HTTPStatusError("Not Found", request=mock_response.request, response=mock_response)
)
with patch(
"litellm.llms.anthropic.files.handler.get_async_httpx_client"
) as mock_get_client:
with patch("litellm.llms.anthropic.files.handler.get_async_httpx_client") as mock_get_client:
mock_client = AsyncMock()
mock_client.get = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
with patch.object(
handler.anthropic_model_info, "get_api_key", return_value="test-api-key"
):
with patch.object(handler.anthropic_model_info, "get_api_key", return_value="test-api-key"):
with patch.object(
handler.anthropic_model_info,
"get_api_base",
@ -480,6 +416,84 @@ class TestAnthropicFilesHandler:
api_key="test-api-key",
)
@pytest.mark.asyncio
async def test_afile_content_resolves_wif_via_async_facade(
self, handler, mock_anthropic_batch_results_succeeded, monkeypatch
):
"""Regression: afile_content ran the blocking WIF mint on the event loop
through the sync get_auth_header; it must go through the async facade."""
import threading
from litellm.llms.anthropic import common_utils as anthropic_common_utils
from litellm.llms.anthropic.wif import aget_anthropic_wif_token, get_anthropic_wif_token
from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine
for name in (
"ANTHROPIC_API_KEY",
"ANTHROPIC_AUTH_TOKEN",
"ANTHROPIC_API_BASE",
"ANTHROPIC_BASE_URL",
):
monkeypatch.delenv(name, raising=False)
monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_files")
monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-files")
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "files-inline-jwt")
minted = "sk-ant-oat01-files-minted"
thread_ids = []
class ThreadRecordingPoster:
def post(self, url, *, content, headers, timeout):
thread_ids.append(threading.get_ident())
return httpx.Response(
200,
json={"access_token": minted, "token_type": "Bearer", "expires_in": 3600},
)
engine = JwtBearerTokenExchangeEngine(poster=ThreadRecordingPoster())
sync_calls = []
def sync_shim(litellm_params, api_base, model):
sync_calls.append(model)
return get_anthropic_wif_token(litellm_params, api_base, model, engine)
async def async_shim(litellm_params, api_base, model):
return await aget_anthropic_wif_token(litellm_params, api_base, model, engine)
monkeypatch.setattr(anthropic_common_utils, "get_anthropic_wif_token", sync_shim)
monkeypatch.setattr(anthropic_common_utils, "aget_anthropic_wif_token", async_shim)
mock_response = httpx.Response(
status_code=200,
content=mock_anthropic_batch_results_succeeded,
headers={"content-type": "application/json"},
request=httpx.Request(
method="GET",
url="https://api.anthropic.com/v1/messages/batches/batch_123/results",
),
)
with patch("litellm.llms.anthropic.files.handler.get_async_httpx_client") as mock_get_client:
mock_client = AsyncMock()
mock_client.get = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
await handler.afile_content(
file_content_request={
"file_id": "batch_123",
"extra_headers": None,
"extra_body": None,
},
api_key=None,
)
sent_headers = mock_client.get.call_args.kwargs["headers"]
assert sent_headers["authorization"] == f"Bearer {minted}"
assert "oauth-2025-04-20" in sent_headers["anthropic-beta"]
assert sync_calls == []
assert thread_ids and thread_ids[0] != threading.get_ident()
class TestAnthropicBatchesConfig:
"""Test Anthropic Batches Config for batch retrieval transformation"""
@ -562,15 +576,11 @@ class TestAnthropicBatchesConfig:
)
assert url == "https://api.anthropic.com/v1/messages/batches/batch_123"
def test_transform_retrieve_batch_response_in_progress(
self, config, mock_anthropic_batch_response_in_progress
):
def test_transform_retrieve_batch_response_in_progress(self, config, mock_anthropic_batch_response_in_progress):
"""Test transformation of in_progress batch response"""
mock_response = httpx.Response(
status_code=200,
content=json.dumps(mock_anthropic_batch_response_in_progress).encode(
"utf-8"
),
content=json.dumps(mock_anthropic_batch_response_in_progress).encode("utf-8"),
request=httpx.Request(
method="GET",
url="https://api.anthropic.com/v1/messages/batches/batch_123",
@ -596,9 +606,7 @@ class TestAnthropicBatchesConfig:
assert batch.in_progress_at is not None
assert batch.completed_at is None
def test_transform_retrieve_batch_response_completed(
self, config, mock_anthropic_batch_response_completed
):
def test_transform_retrieve_batch_response_completed(self, config, mock_anthropic_batch_response_completed):
"""Test transformation of completed batch response"""
mock_response = httpx.Response(
status_code=200,
@ -624,9 +632,7 @@ class TestAnthropicBatchesConfig:
assert batch.request_counts.completed == 10
assert batch.request_counts.failed == 0
def test_transform_retrieve_batch_response_canceling(
self, config, mock_anthropic_batch_response_canceling
):
def test_transform_retrieve_batch_response_canceling(self, config, mock_anthropic_batch_response_canceling):
"""Test transformation of canceling batch response"""
mock_response = httpx.Response(
status_code=200,
@ -663,9 +669,7 @@ class TestAnthropicBatchesConfig:
)
logging_obj = MagicMock()
with pytest.raises(
ValueError, match="Failed to parse Anthropic batch response"
):
with pytest.raises(ValueError, match="Failed to parse Anthropic batch response"):
config.transform_retrieve_batch_response(
model="claude-3-5-sonnet-20241022",
raw_response=mock_response,

View file

@ -0,0 +1,637 @@
import concurrent.futures
import json
from collections.abc import Callable, Mapping
from pathlib import Path
from typing import Final
import httpx
import pytest
import litellm
from litellm.llms.anthropic.wif import (
AnthropicWifParams,
_raise_anthropic_wif_error,
build_anthropic_wif_spec,
get_anthropic_wif_token,
resolve_anthropic_wif_params,
)
from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine
from litellm.llms.base_llm.auth.types import (
AssertionSourceError,
ExchangeError,
InsecureTokenUrl,
MalformedTokenResponse,
TokenEndpointError,
TokenTransportError,
)
WIF_ENV_VARS: Final = (
"ANTHROPIC_FEDERATION_RULE_ID",
"ANTHROPIC_ORGANIZATION_ID",
"ANTHROPIC_SERVICE_ACCOUNT_ID",
"ANTHROPIC_WORKSPACE_ID",
"ANTHROPIC_IDENTITY_TOKEN_FILE",
"ANTHROPIC_IDENTITY_TOKEN",
"ANTHROPIC_SCOPE",
"ANTHROPIC_API_BASE",
"ANTHROPIC_BASE_URL",
"LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS",
)
GRANT_TYPE: Final = "urn:ietf:params:oauth:grant-type:jwt-bearer"
@pytest.fixture(autouse=True)
def _clean_wif_env(monkeypatch: pytest.MonkeyPatch) -> None:
for name in WIF_ENV_VARS:
monkeypatch.delenv(name, raising=False)
class FakeClock:
def __init__(self, start: float = 1_000.0) -> None:
self.now = start
def __call__(self) -> float:
return self.now
def advance(self, seconds: float) -> None:
self.now += seconds
class RecordedRequest:
def __init__(self, url: str, content: bytes, headers: Mapping[str, str], timeout: float) -> None:
self.url = url
self.content = content
self.headers = dict(headers)
self.timeout = timeout
def json_body(self) -> dict:
return json.loads(self.content)
class ScriptedPoster:
def __init__(self, responses: list[httpx.Response]) -> None:
self.requests: list[RecordedRequest] = []
self._responses = list(responses)
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
self.requests.append(RecordedRequest(url, content, headers, timeout))
if len(self._responses) > 1:
return self._responses.pop(0)
return self._responses[0]
class ManualExecutor(concurrent.futures.Executor):
def __init__(self) -> None:
self.pending: list[Callable[[], None]] = []
def submit(self, fn, /, *args, **kwargs):
future: concurrent.futures.Future = concurrent.futures.Future()
self.pending.append(lambda: fn(*args, **kwargs))
return future
def token_response(token: str = "sk-ant-oat01-minted", expires_in: int | None = 3600) -> httpx.Response:
body = {"access_token": token, "token_type": "Bearer"}
if expires_in is not None:
body["expires_in"] = expires_in
return httpx.Response(200, json=body)
def make_engine(poster: ScriptedPoster, clock: FakeClock | None = None) -> JwtBearerTokenExchangeEngine:
return JwtBearerTokenExchangeEngine(
poster=poster,
clock=clock if clock is not None else FakeClock(),
refresh_executor=ManualExecutor(),
)
def write_token_file(directory: Path, content: str, name: str = "identity-token") -> Path:
directory.mkdir(parents=True, exist_ok=True)
token_file = directory / name
token_file.write_text(content, encoding="utf-8")
return token_file
class TestWireProtocolExact:
def test_minimal_body_and_headers(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path))
monkeypatch.setenv("ANTHROPIC_SCOPE", "user:inference")
token_file = write_token_file(tmp_path, "jwt-assertion-value\n")
poster = ScriptedPoster([token_response()])
engine = make_engine(poster)
token = get_anthropic_wif_token(
{
"anthropic_federation_rule_id": "fdrl_abc123",
"anthropic_organization_id": "org-uuid-1",
"anthropic_identity_token_file": str(token_file),
},
"https://api.anthropic.com",
"claude-sonnet-4-5",
engine,
)
assert token == "sk-ant-oat01-minted"
assert len(poster.requests) == 1
request = poster.requests[0]
assert request.url == "https://api.anthropic.com/v1/oauth/token"
assert "anthropic-beta" not in request.headers
assert request.headers["content-type"] == "application/json"
assert request.json_body() == {
"grant_type": GRANT_TYPE,
"federation_rule_id": "fdrl_abc123",
"organization_id": "org-uuid-1",
"assertion": "jwt-assertion-value",
}
def test_optional_fields_present_when_set(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path))
token_file = write_token_file(tmp_path, "jwt-assertion-value")
poster = ScriptedPoster([token_response()])
engine = make_engine(poster)
get_anthropic_wif_token(
{
"anthropic_federation_rule_id": "fdrl_abc123",
"anthropic_organization_id": "org-uuid-1",
"anthropic_service_account_id": "svcacct_1",
"anthropic_workspace_id": "wrkspc_1",
"anthropic_identity_token_file": str(token_file),
},
"https://api.anthropic.com",
"claude-sonnet-4-5",
engine,
)
request = poster.requests[0]
assert "anthropic-beta" not in request.headers
assert request.headers["content-type"] == "application/json"
assert request.json_body() == {
"grant_type": GRANT_TYPE,
"federation_rule_id": "fdrl_abc123",
"organization_id": "org-uuid-1",
"service_account_id": "svcacct_1",
"workspace_id": "wrkspc_1",
"assertion": "jwt-assertion-value",
}
def test_spec_cache_key_identity(self):
params = AnthropicWifParams(
federation_rule_id="fdrl_1",
organization_id="org-1",
assertion_ref="oidc/env/ANTHROPIC_IDENTITY_TOKEN",
)
spec = build_anthropic_wif_spec(params, "https://api.anthropic.com")
assert spec.cache_key_identity == ("fdrl_1", "org-1", "", "")
assert spec.body_encoding == "json"
assert spec.assertion_field == "assertion"
def test_full_params_spec_has_no_request_headers(self):
"""The token exchange sends no anthropic-beta header at all (verified against the
live endpoint); this must hold even for a fully populated params set, so a future
edit cannot reintroduce the header gated on service_account_id or workspace_id."""
params = AnthropicWifParams(
federation_rule_id="fdrl_1",
organization_id="org-1",
service_account_id="svcacct_1",
workspace_id="wrkspc_1",
assertion_ref="oidc/env/ANTHROPIC_IDENTITY_TOKEN",
)
spec = build_anthropic_wif_spec(params, "https://api.anthropic.com")
assert dict(spec.request_headers) == {}
class TestBaseUrlDerivation:
def _mint(self, api_base: str | None, monkeypatch: pytest.MonkeyPatch) -> str:
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt")
poster = ScriptedPoster([token_response()])
engine = make_engine(poster)
get_anthropic_wif_token(
{"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"},
api_base,
"claude-sonnet-4-5",
engine,
)
return poster.requests[0].url
def test_explicit_api_base_wins(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ANTHROPIC_API_BASE", "https://env.example.com")
assert self._mint("https://gw.example.com/", monkeypatch) == "https://gw.example.com/v1/oauth/token"
def test_env_api_base(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ANTHROPIC_API_BASE", "https://env.example.com")
assert self._mint(None, monkeypatch) == "https://env.example.com/v1/oauth/token"
def test_env_base_url(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ANTHROPIC_BASE_URL", "https://base.example.com")
assert self._mint(None, monkeypatch) == "https://base.example.com/v1/oauth/token"
def test_default_base(self, monkeypatch: pytest.MonkeyPatch):
assert self._mint(None, monkeypatch) == "https://api.anthropic.com/v1/oauth/token"
@pytest.mark.parametrize(
"api_base",
[
"https://gw.example.com/v1/messages",
"https://gw.example.com/v1/messages/",
"https://gw.example.com/v1/messages//v1/messages",
],
)
def test_chat_appended_bases_normalize_to_clean_token_url(self, api_base: str, monkeypatch: pytest.MonkeyPatch):
"""main.py appends /v1/messages before dispatch (twice for trailing-slash
bases); the exchange must still target the deployment base."""
assert self._mint(api_base, monkeypatch) == "https://gw.example.com/v1/oauth/token"
def test_trailing_slash_env_base_normalizes(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ANTHROPIC_BASE_URL", "https://base.example.com/")
assert self._mint(None, monkeypatch) == "https://base.example.com/v1/oauth/token"
class TestSecretManagerEnvResolution:
"""WIF env vars resolve through get_secret_str so configured secret managers
work, exactly like every sibling Anthropic credential."""
def test_values_resolve_through_get_secret_str(self, monkeypatch: pytest.MonkeyPatch):
secrets: Final = {
"ANTHROPIC_FEDERATION_RULE_ID": "fdrl_sm",
"ANTHROPIC_ORGANIZATION_ID": "org-sm",
"ANTHROPIC_IDENTITY_TOKEN": "sm-inline-jwt",
}
monkeypatch.setattr(
"litellm.secret_managers.main.get_secret_str",
lambda secret_name, default_value=None: secrets.get(secret_name, default_value),
)
params = resolve_anthropic_wif_params(None)
assert params == AnthropicWifParams(
federation_rule_id="fdrl_sm",
organization_id="org-sm",
assertion_ref="oidc/env/ANTHROPIC_IDENTITY_TOKEN",
)
def test_non_str_secret_value_treated_as_unset(self, monkeypatch: pytest.MonkeyPatch):
secrets: Final = {
"ANTHROPIC_FEDERATION_RULE_ID": {"unexpected": "shape"},
"ANTHROPIC_ORGANIZATION_ID": "org-sm",
"ANTHROPIC_IDENTITY_TOKEN": "sm-inline-jwt",
}
monkeypatch.setattr(
"litellm.secret_managers.main.get_secret_str",
lambda secret_name, default_value=None: secrets.get(secret_name, default_value),
)
assert resolve_anthropic_wif_params(None) is None
class TestResolutionMatrix:
def test_params_beat_env_per_field(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_env")
monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-env")
monkeypatch.setenv("ANTHROPIC_SERVICE_ACCOUNT_ID", "svc-env")
monkeypatch.setenv("ANTHROPIC_WORKSPACE_ID", "wrkspc_env")
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN_FILE", "/var/run/secrets/env-token")
params = resolve_anthropic_wif_params(
{
"anthropic_federation_rule_id": "fdrl_param",
"anthropic_organization_id": "org-param",
"anthropic_service_account_id": "svc-param",
"anthropic_workspace_id": "wrkspc_param",
"anthropic_identity_token_file": "/var/run/secrets/param-token",
}
)
assert params == AnthropicWifParams(
federation_rule_id="fdrl_param",
organization_id="org-param",
service_account_id="svc-param",
workspace_id="wrkspc_param",
assertion_ref="oidc/file//var/run/secrets/param-token",
)
def test_env_only_config(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_env")
monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-env")
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "raw-env-jwt")
params = resolve_anthropic_wif_params(None)
assert params is not None
assert params.assertion_ref == "oidc/env/ANTHROPIC_IDENTITY_TOKEN"
assert params.service_account_id is None
assert params.workspace_id is None
def test_file_param_beats_inline_param(self):
params = resolve_anthropic_wif_params(
{
"anthropic_federation_rule_id": "fdrl_1",
"anthropic_organization_id": "org-1",
"anthropic_identity_token_file": "/var/run/secrets/tok",
"anthropic_identity_token": "oidc/env/OTHER",
}
)
assert params is not None
assert params.assertion_ref == "oidc/file//var/run/secrets/tok"
def test_inline_param_beats_env_file(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN_FILE", "/var/run/secrets/env-tok")
params = resolve_anthropic_wif_params(
{
"anthropic_federation_rule_id": "fdrl_1",
"anthropic_organization_id": "org-1",
"anthropic_identity_token": "oidc/env/OTHER",
}
)
assert params is not None
assert params.assertion_ref == "oidc/env/OTHER"
def test_env_file_beats_env_inline(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN_FILE", "/var/run/secrets/env-tok")
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "raw-env-jwt")
params = resolve_anthropic_wif_params(
{"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"}
)
assert params is not None
assert params.assertion_ref == "oidc/file//var/run/secrets/env-tok"
def test_empty_workspace_env_coerced_to_none(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ANTHROPIC_WORKSPACE_ID", "")
params = resolve_anthropic_wif_params(
{
"anthropic_federation_rule_id": "fdrl_1",
"anthropic_organization_id": "org-1",
"anthropic_identity_token": "oidc/env/TOK",
}
)
assert params is not None
assert params.workspace_id is None
spec = build_anthropic_wif_spec(params, "https://api.anthropic.com")
assert "workspace_id" not in spec.static_body
@pytest.mark.parametrize(
"litellm_params",
[
{},
{"anthropic_federation_rule_id": "fdrl_1"},
{"anthropic_organization_id": "org-1"},
{"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"},
{"anthropic_organization_id": "org-1", "anthropic_identity_token": "oidc/env/TOK"},
{"anthropic_federation_rule_id": "fdrl_1", "anthropic_identity_token": "oidc/env/TOK"},
],
)
def test_gate_unmet_returns_none(self, litellm_params: dict):
assert resolve_anthropic_wif_params(litellm_params) is None
def test_gate_unmet_facade_returns_none_without_engine_call(self):
poster = ScriptedPoster([token_response()])
engine = make_engine(poster)
assert get_anthropic_wif_token({}, None, "claude-sonnet-4-5", engine) is None
assert poster.requests == []
class TestServiceAccountIdIsOptional:
"""Anthropic's reference docs mark service_account_id required, but a live exchange
against a federation rule targeting a single service account mints successfully
without it; resolution must not gate activation on it, and the wire body must omit
the key entirely rather than send it as null."""
def test_activates_and_omits_service_account_id_when_unset(
self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
):
monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path))
token_file = write_token_file(tmp_path, "jwt-assertion-value")
litellm_params: Final = {
"anthropic_federation_rule_id": "fdrl_1",
"anthropic_organization_id": "org-1",
"anthropic_identity_token_file": str(token_file),
}
params = resolve_anthropic_wif_params(litellm_params)
assert params is not None
assert params.service_account_id is None
poster = ScriptedPoster([token_response()])
engine = make_engine(poster)
token = get_anthropic_wif_token(litellm_params, "https://api.anthropic.com", "claude-sonnet-4-5", engine)
assert token == "sk-ant-oat01-minted"
assert "service_account_id" not in poster.requests[0].json_body()
class TestInlineRefRestrictions:
RAW_JWT: Final = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ3b3JrbG9hZCJ9.c2lnbmF0dXJl"
@pytest.mark.parametrize("bad_ref", [RAW_JWT, "oidc/env_path/ANTHROPIC_TOKEN_PATH"])
def test_rejected_inline_refs(self, bad_ref: str):
poster = ScriptedPoster([token_response()])
engine = make_engine(poster)
with pytest.raises(litellm.AuthenticationError) as exc_info:
get_anthropic_wif_token(
{
"anthropic_federation_rule_id": "fdrl_1",
"anthropic_organization_id": "org-1",
"anthropic_identity_token": bad_ref,
},
None,
"claude-sonnet-4-5",
engine,
)
assert "oidc/env/" in exc_info.value.message
assert "oidc/file/" in exc_info.value.message
assert self.RAW_JWT not in exc_info.value.message
assert poster.requests == []
class TestFileAllowlistAndSymlink:
SECRET_CONTENT: Final = "super-secret-jwt-content"
def _call(self, token_file: Path, poster: ScriptedPoster) -> str | None:
engine = make_engine(poster)
return get_anthropic_wif_token(
{
"anthropic_federation_rule_id": "fdrl_1",
"anthropic_organization_id": "org-1",
"anthropic_identity_token_file": str(token_file),
},
"https://api.anthropic.com",
"claude-sonnet-4-5",
engine,
)
def test_file_outside_allowlist_rejected(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path / "allowed"))
token_file = write_token_file(tmp_path / "outside", self.SECRET_CONTENT)
poster = ScriptedPoster([token_response()])
with pytest.raises(litellm.AuthenticationError) as exc_info:
self._call(token_file, poster)
assert str(token_file) in exc_info.value.message
assert self.SECRET_CONTENT not in exc_info.value.message
assert poster.requests == []
def test_disallowed_path_message_names_allowlist_and_env_var(
self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
):
"""The disallowed_path error must explain the allowlist and name the env var an
operator would set, not surface as a bare '(disallowed_path)' code dump."""
monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path / "allowed"))
token_file = write_token_file(tmp_path / "outside", self.SECRET_CONTENT)
poster = ScriptedPoster([token_response()])
with pytest.raises(litellm.AuthenticationError) as exc_info:
self._call(token_file, poster)
message = exc_info.value.message
assert "(disallowed_path)" not in message
assert "LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS" in message
assert "allowed credential director" in message
def test_symlink_escape_rejected(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
allowed = tmp_path / "allowed"
allowed.mkdir()
monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(allowed))
outside_file = write_token_file(tmp_path / "outside", self.SECRET_CONTENT)
link = allowed / "identity-token"
link.symlink_to(outside_file)
poster = ScriptedPoster([token_response()])
with pytest.raises(litellm.AuthenticationError) as exc_info:
self._call(link, poster)
assert self.SECRET_CONTENT not in exc_info.value.message
assert poster.requests == []
def test_file_inside_allowlist_succeeds(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path))
token_file = write_token_file(tmp_path, self.SECRET_CONTENT)
poster = ScriptedPoster([token_response()])
assert self._call(token_file, poster) == "sk-ant-oat01-minted"
assert poster.requests[0].json_body()["assertion"] == self.SECRET_CONTENT
class TestErrorMappingExhaustive:
@pytest.mark.parametrize(
"error",
[
AssertionSourceError(kind="missing", source_ref="oidc/env/TOK"),
AssertionSourceError(kind="disallowed_path", source_ref="oidc/file//etc/passwd"),
InsecureTokenUrl(host="token.example"),
TokenEndpointError(status_code=500, redacted_body="error: server_error"),
TokenTransportError(detail="ConnectError: refused"),
MalformedTokenResponse(detail="token response failed RFC 6749 5.1 schema validation"),
],
)
def test_every_variant_maps_to_authentication_error(self, error: ExchangeError):
with pytest.raises(litellm.AuthenticationError) as exc_info:
_raise_anthropic_wif_error(error, model="claude-sonnet-4-5", workspace_id_set=False)
assert exc_info.value.llm_provider == "anthropic"
assert exc_info.value.model == "claude-sonnet-4-5"
def test_endpoint_error_raised_through_facade(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt")
poster = ScriptedPoster([httpx.Response(500, json={"error": "server_error"})])
engine = make_engine(poster)
with pytest.raises(litellm.AuthenticationError) as exc_info:
get_anthropic_wif_token(
{"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"},
None,
"claude-sonnet-4-5",
engine,
)
assert exc_info.value.llm_provider == "anthropic"
assert "HTTP 500" in exc_info.value.message
assert "server_error" in exc_info.value.message
@pytest.mark.parametrize(
"litellm_params,status_code,body",
[
(
{"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"},
401,
{"error": "invalid_grant."},
),
(
{
"anthropic_federation_rule_id": "fdrl_1",
"anthropic_organization_id": "org-1",
"anthropic_workspace_id": "wrkspc_1",
},
500,
{"error": "server_error."},
),
],
)
def test_token_endpoint_error_message_has_no_doubled_period(
self, litellm_params: dict, status_code: int, body: dict, monkeypatch: pytest.MonkeyPatch
):
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt")
poster = ScriptedPoster([httpx.Response(status_code, json=body)])
engine = make_engine(poster)
with pytest.raises(litellm.AuthenticationError) as exc_info:
get_anthropic_wif_token(litellm_params, None, "claude-sonnet-4-5", engine)
assert ".." not in exc_info.value.message
class TestWorkspaceHint:
def _raise_401(self, litellm_params: dict, monkeypatch: pytest.MonkeyPatch) -> str:
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt")
poster = ScriptedPoster([httpx.Response(401, json={"error": "invalid_grant"})])
engine = make_engine(poster)
with pytest.raises(litellm.AuthenticationError) as exc_info:
get_anthropic_wif_token(litellm_params, None, "claude-sonnet-4-5", engine)
assert len(poster.requests) == 2
return exc_info.value.message
def test_hint_when_workspace_unset(self, monkeypatch: pytest.MonkeyPatch):
message = self._raise_401(
{"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"}, monkeypatch
)
assert "ANTHROPIC_WORKSPACE_ID" in message
def test_no_hint_when_workspace_set(self, monkeypatch: pytest.MonkeyPatch):
message = self._raise_401(
{
"anthropic_federation_rule_id": "fdrl_1",
"anthropic_organization_id": "org-1",
"anthropic_workspace_id": "wrkspc_1",
},
monkeypatch,
)
assert "ANTHROPIC_WORKSPACE_ID" not in message
class TestFileRereadOnRefresh:
def test_mandatory_refresh_carries_rotated_assertion(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path))
token_file = write_token_file(tmp_path, "first-assertion")
clock = FakeClock(start=1_000.0)
poster = ScriptedPoster(
[token_response("sk-ant-oat01-first", 3600), token_response("sk-ant-oat01-second", 3600)]
)
engine = make_engine(poster, clock=clock)
litellm_params = {
"anthropic_federation_rule_id": "fdrl_1",
"anthropic_organization_id": "org-1",
"anthropic_identity_token_file": str(token_file),
}
first = get_anthropic_wif_token(litellm_params, "https://api.anthropic.com", "claude-sonnet-4-5", engine)
token_file.write_text("second-assertion", encoding="utf-8")
clock.advance(3600 - 10)
second = get_anthropic_wif_token(litellm_params, "https://api.anthropic.com", "claude-sonnet-4-5", engine)
assert first == "sk-ant-oat01-first"
assert second == "sk-ant-oat01-second"
assert len(poster.requests) == 2
assert poster.requests[1].json_body()["assertion"] == "second-assertion"

View file

@ -0,0 +1,953 @@
import asyncio
import concurrent.futures
import json
import logging
import threading
import time
from collections.abc import Callable, Mapping
from typing import Final
from urllib.parse import parse_qsl
import httpx
import pytest
from litellm.llms.base_llm.auth.token_exchange import (
ADVISORY_REFRESH_BACKOFF_SECONDS,
FALLBACK_TOKEN_TTL_SECONDS,
MAX_ASSERTION_BYTES,
MAX_RESPONSE_BYTES,
JwtBearerTokenExchangeEngine,
redact_oauth_error_body,
)
from litellm.llms.base_llm.auth.types import (
AssertionSourceError,
ExchangeResult,
InsecureTokenUrl,
MalformedTokenResponse,
MintedToken,
TokenEndpointError,
TokenExchangeSpec,
TokenTransportError,
)
from litellm.secret_managers.main import OidcPathNotAllowedError, _resolve_oidc_file_path
DEFAULT_REF: Final = "oidc/env/TEST_ASSERTION"
DEFAULT_ASSERTION: Final = "test-jwt-assertion"
class FakeClock:
def __init__(self, start: float = 1_000.0) -> None:
self.now = start
def __call__(self) -> float:
return self.now
def advance(self, seconds: float) -> None:
self.now += seconds
class RecordedRequest:
def __init__(self, url: str, content: bytes, headers: Mapping[str, str], timeout: float) -> None:
self.url = url
self.content = content
self.headers = dict(headers)
self.timeout = timeout
def json_body(self) -> dict:
return json.loads(self.content)
class ScriptedPoster:
"""Returns scripted responses in order (repeating the last one); records requests."""
def __init__(
self,
responses: list[httpx.Response],
on_request: Callable[[RecordedRequest], None] | None = None,
) -> None:
self.requests: list[RecordedRequest] = []
self._responses = list(responses)
self._on_request = on_request
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
recorded = RecordedRequest(url, content, headers, timeout)
self.requests.append(recorded)
if self._on_request is not None:
self._on_request(recorded)
if len(self._responses) > 1:
return self._responses.pop(0)
return self._responses[0]
class RaisingPoster:
def __init__(self, error: Exception) -> None:
self.calls = 0
self._error = error
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
self.calls += 1
raise self._error
class ManualExecutor(concurrent.futures.Executor):
"""Records submissions; runs them only when the test says so."""
def __init__(self) -> None:
self.pending: list[Callable[[], None]] = []
def submit(self, fn, /, *args, **kwargs):
future: concurrent.futures.Future = concurrent.futures.Future()
self.pending.append(lambda: fn(*args, **kwargs))
return future
def run_all(self) -> None:
drained = list(self.pending)
self.pending.clear()
for job in drained:
job()
class InlineExecutor(concurrent.futures.Executor):
def submit(self, fn, /, *args, **kwargs):
future: concurrent.futures.Future = concurrent.futures.Future()
future.set_result(fn(*args, **kwargs))
return future
def token_response(token: str = "sk-ant-oat01-minted", expires_in: int | None = 3600) -> httpx.Response:
body = {"access_token": token, "token_type": "Bearer"}
if expires_in is not None:
body["expires_in"] = expires_in
return httpx.Response(200, json=body)
def make_spec(**overrides) -> TokenExchangeSpec:
base = {
"token_url": "https://token.example/v1/oauth/token",
"assertion_ref": DEFAULT_REF,
"assertion_field": "assertion",
"static_body": {
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
"federation_rule_id": "fdrl_1",
"organization_id": "org-1",
},
"body_encoding": "json",
"request_headers": {"anthropic-beta": "oauth-2025-04-20,oidc-federation-2026-04-01"},
"cache_key_identity": ("fdrl_1", "org-1", "", ""),
"timeout_seconds": 2.0,
}
base.update(overrides)
return TokenExchangeSpec(**base)
def make_engine(
poster,
reader: Mapping[str, str] | Callable[[str], str | None] | None = None,
clock: FakeClock | None = None,
executor: concurrent.futures.Executor | None = None,
max_entries: int = 64,
) -> JwtBearerTokenExchangeEngine:
resolved_reader = reader if callable(reader) else (reader or {DEFAULT_REF: DEFAULT_ASSERTION}).get
return JwtBearerTokenExchangeEngine(
poster=poster,
assertion_reader=resolved_reader,
clock=clock if clock is not None else FakeClock(),
refresh_executor=executor if executor is not None else ManualExecutor(),
max_entries=max_entries,
)
def mint(engine: JwtBearerTokenExchangeEngine, spec: TokenExchangeSpec) -> MintedToken:
result = engine.get_token(spec)
assert isinstance(result, MintedToken)
return result
class TestFreshMintWireExact:
def test_json_body_and_headers(self):
poster = ScriptedPoster([token_response(expires_in=3600)])
clock = FakeClock(start=1_000.0)
engine = make_engine(poster, clock=clock)
spec = make_spec()
result = mint(engine, spec)
assert result.access_token.get_secret_value() == "sk-ant-oat01-minted"
assert result.expires_at == 1_000.0 + 3600
assert len(poster.requests) == 1
request = poster.requests[0]
assert request.url == "https://token.example/v1/oauth/token"
assert request.timeout == 2.0
assert request.headers == {
"content-type": "application/json",
"anthropic-beta": "oauth-2025-04-20,oidc-federation-2026-04-01",
}
assert request.json_body() == {
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
"federation_rule_id": "fdrl_1",
"organization_id": "org-1",
"assertion": DEFAULT_ASSERTION,
}
def test_form_body_and_content_type(self):
poster = ScriptedPoster([token_response()])
engine = make_engine(poster)
spec = make_spec(body_encoding="form")
mint(engine, spec)
request = poster.requests[0]
assert request.headers["content-type"] == "application/x-www-form-urlencoded"
assert dict(parse_qsl(request.content.decode())) == {
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
"federation_rule_id": "fdrl_1",
"organization_id": "org-1",
"assertion": DEFAULT_ASSERTION,
}
def test_cache_hit_zero_posts():
poster = ScriptedPoster([token_response(expires_in=3600)])
clock = FakeClock()
engine = make_engine(poster, clock=clock)
spec = make_spec()
first = mint(engine, spec)
clock.advance(100.0)
second = mint(engine, spec)
assert len(poster.requests) == 1
assert second.access_token.get_secret_value() == first.access_token.get_secret_value()
@pytest.mark.parametrize(
"remaining,expect_advisory_submit,expect_new_token",
[
(121.0, False, False),
(120.0, True, False),
(119.0, True, False),
(31.0, True, False),
(30.0, False, True),
(29.0, False, True),
],
)
def test_window_boundaries(remaining: float, expect_advisory_submit: bool, expect_new_token: bool):
poster = ScriptedPoster([token_response("old-token", expires_in=3600), token_response("new-token")])
clock = FakeClock(start=1_000.0)
executor = ManualExecutor()
engine = make_engine(poster, clock=clock, executor=executor)
spec = make_spec()
mint(engine, spec)
expires_at = 1_000.0 + 3600
clock.now = expires_at - remaining
result = mint(engine, spec)
assert len(executor.pending) == (1 if expect_advisory_submit else 0)
expected_token = "new-token" if expect_new_token else "old-token"
assert result.access_token.get_secret_value() == expected_token
assert len(poster.requests) == (2 if expect_new_token else 1)
def test_advisory_serve_stale_single_flight_backoff(caplog: pytest.LogCaptureFixture):
poster = ScriptedPoster(
[
token_response("stale-token", expires_in=3600),
httpx.Response(500, json={"error": "server_error"}),
httpx.Response(500, json={"error": "server_error"}),
]
)
clock = FakeClock(start=1_000.0)
executor = ManualExecutor()
engine = make_engine(poster, clock=clock, executor=executor)
spec = make_spec()
mint(engine, spec)
clock.now = 1_000.0 + 3600 - 100.0
first = mint(engine, spec)
second = mint(engine, spec)
assert first.access_token.get_secret_value() == "stale-token"
assert second.access_token.get_secret_value() == "stale-token"
assert len(executor.pending) == 1
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
executor.run_all()
assert len(poster.requests) == 2
warning_records = [r for r in caplog.records if r.levelno == logging.WARNING]
assert any("Advisory token refresh" in r.getMessage() for r in warning_records)
assert "server_error" in caplog.text
assert DEFAULT_ASSERTION not in caplog.text
assert "stale-token" not in caplog.text
within_backoff = mint(engine, spec)
assert within_backoff.access_token.get_secret_value() == "stale-token"
assert len(executor.pending) == 0
clock.advance(ADVISORY_REFRESH_BACKOFF_SECONDS)
after_backoff = mint(engine, spec)
assert after_backoff.access_token.get_secret_value() == "stale-token"
assert len(executor.pending) == 1
executor.run_all()
assert len(poster.requests) == 3
class GatedPoster:
"""Blocks the leader inside post() until the test releases it."""
def __init__(self, response: httpx.Response) -> None:
self.entered = threading.Event()
self.release = threading.Event()
self.calls = 0
self._calls_lock = threading.Lock()
self._response = response
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
with self._calls_lock:
self.calls += 1
self.entered.set()
assert self.release.wait(timeout=10)
return self._response
def _run_concurrent_get_token(
engine: JwtBearerTokenExchangeEngine, spec: TokenExchangeSpec, poster: GatedPoster, thread_count: int
) -> list[ExchangeResult]:
results: list[ExchangeResult] = []
results_lock = threading.Lock()
start_barrier = threading.Barrier(thread_count)
def worker() -> None:
start_barrier.wait()
result = engine.get_token(spec)
with results_lock:
results.append(result)
threads = [threading.Thread(target=worker, daemon=True) for _ in range(thread_count)]
for thread in threads:
thread.start()
assert poster.entered.wait(timeout=10)
time.sleep(0.3)
poster.release.set()
for thread in threads:
thread.join(timeout=10)
assert not thread.is_alive()
return results
def test_mandatory_single_leader():
poster = GatedPoster(token_response("leader-token"))
engine = make_engine(poster)
spec = make_spec()
results = _run_concurrent_get_token(engine, spec, poster, thread_count=5)
assert poster.calls == 1
assert len(results) == 5
for result in results:
assert isinstance(result, MintedToken)
assert result.access_token.get_secret_value() == "leader-token"
def test_mandatory_failure_is_value():
poster = GatedPoster(httpx.Response(500, json={"error": "server_error"}))
engine = make_engine(poster)
spec = make_spec()
results = _run_concurrent_get_token(engine, spec, poster, thread_count=3)
assert len(results) == 3
for result in results:
assert isinstance(result, TokenEndpointError)
assert result.status_code == 500
assert "server_error" in result.redacted_body
def test_lock_released_around_io():
inner_spec = make_spec(
token_url="https://inner.example/v1/oauth/token",
cache_key_identity=("fdrl_inner", "org-1", "", ""),
)
engine_holder: dict[str, JwtBearerTokenExchangeEngine] = {}
inner_results: list[ExchangeResult] = []
class ReentrantPoster:
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
if url == "https://token.example/v1/oauth/token":
inner_results.append(engine_holder["engine"].get_token(inner_spec))
return token_response()
engine = make_engine(ReentrantPoster())
engine_holder["engine"] = engine
outcome: list[ExchangeResult] = []
thread = threading.Thread(target=lambda: outcome.append(engine.get_token(make_spec())), daemon=True)
thread.start()
thread.join(timeout=10)
assert not thread.is_alive(), "engine held its lock across poster I/O and deadlocked"
assert len(outcome) == 1
assert isinstance(outcome[0], MintedToken)
assert len(inner_results) == 1
assert isinstance(inner_results[0], MintedToken)
def test_401_retry_once_with_reread():
assertions = {DEFAULT_REF: "assertion-v1"}
def rotate_on_first_request(request: RecordedRequest) -> None:
assertions[DEFAULT_REF] = "assertion-v2"
poster = ScriptedPoster(
[httpx.Response(401, json={"error": "invalid_grant"}), token_response()],
on_request=rotate_on_first_request,
)
engine = make_engine(poster, reader=assertions.get)
result = mint(engine, make_spec())
assert result.access_token.get_secret_value() == "sk-ant-oat01-minted"
assert len(poster.requests) == 2
assert poster.requests[0].json_body()["assertion"] == "assertion-v1"
assert poster.requests[1].json_body()["assertion"] == "assertion-v2"
def test_401_twice_is_endpoint_error():
poster = ScriptedPoster([httpx.Response(401, json={"error": "invalid_grant"})])
engine = make_engine(poster)
result = engine.get_token(make_spec())
assert isinstance(result, TokenEndpointError)
assert result.status_code == 401
assert "invalid_grant" in result.redacted_body
assert len(poster.requests) == 2
class TestRedactionAndCaps:
def test_object_body_reduced_to_rfc6749_fields(self):
poster = ScriptedPoster(
[
httpx.Response(
400,
json={
"error": "invalid_grant",
"error_description": "d" * 500,
"error_uri": "https://errors.example/e1",
"assertion_echo": "LEAKED-ASSERTION",
},
)
]
)
result = make_engine(poster).get_token(make_spec())
assert isinstance(result, TokenEndpointError)
assert result.status_code == 400
assert "invalid_grant" in result.redacted_body
assert "d" * 256 in result.redacted_body
assert "d" * 257 not in result.redacted_body
assert "https://errors.example/e1" in result.redacted_body
assert "LEAKED-ASSERTION" not in result.redacted_body
def test_nested_error_envelope_renders_readable_text(self):
body = {
"type": "error",
"error": {
"type": "invalid_request_error",
"message": "federation_rule_id is not a well-formed fdrl_ tagged ID",
},
}
result = redact_oauth_error_body(400, json.dumps(body))
assert "invalid_request_error" in result.redacted_body
assert "federation_rule_id is not a well-formed fdrl_ tagged ID" in result.redacted_body
assert "{'" not in result.redacted_body
def test_flat_rfc6749_shape_still_renders(self):
body = {"error": "invalid_grant", "error_description": "bad request"}
result = redact_oauth_error_body(400, json.dumps(body))
assert result.redacted_body == "error: invalid_grant; error_description: bad request"
def test_nested_error_message_is_capped_at_256_chars(self):
body = {"error": {"type": "invalid_request_error", "message": "m" * 500}}
result = redact_oauth_error_body(400, json.dumps(body))
assert "m" * 256 in result.redacted_body
assert "m" * 257 not in result.redacted_body
def test_string_body_truncated(self):
result = redact_oauth_error_body(400, json.dumps("s" * 500))
assert result.redacted_body == "s" * 256
def test_plain_text_body_truncated(self):
result = redact_oauth_error_body(502, "t" * 500)
assert result.redacted_body == "t" * 256
def test_json_array_body_constant_message(self):
result = redact_oauth_error_body(400, json.dumps(["a", "b"]))
assert result.redacted_body == "non-object error response omitted"
def test_oversized_body_never_parsed(self):
poster = ScriptedPoster([httpx.Response(400, content=b'{"error": "' + b"x" * MAX_RESPONSE_BYTES + b'"}')])
result = make_engine(poster).get_token(make_spec())
assert isinstance(result, TokenEndpointError)
assert result.redacted_body == "oversized error response omitted"
def test_oversized_success_body_is_malformed(self):
poster = ScriptedPoster(
[httpx.Response(200, content=b'{"access_token": "' + b"x" * MAX_RESPONSE_BYTES + b'"}')]
)
result = make_engine(poster).get_token(make_spec())
assert not isinstance(result, MintedToken)
assert b"x" * 10 not in str(result).encode()
@pytest.mark.parametrize("access_token", ["", " "])
def test_empty_access_token_is_malformed(access_token: str):
poster = ScriptedPoster(
[httpx.Response(200, json={"access_token": access_token, "token_type": "Bearer", "expires_in": 3600})]
)
result = make_engine(poster).get_token(make_spec())
assert isinstance(result, MalformedTokenResponse)
assert "empty access_token" in result.detail
def test_sentinel_leak_audit(caplog: pytest.LogCaptureFixture):
jwt_sentinel = "JWT-SENTINEL-2c9f1e7ab4"
token_sentinel = "sk-ant-oat01-TOKEN-SENTINEL-90d4c3aa17"
ref = "oidc/env/SENTINEL_ASSERTION"
with caplog.at_level(logging.DEBUG):
success_poster = ScriptedPoster([token_response(token_sentinel, expires_in=3600)])
success_clock = FakeClock()
success_executor = ManualExecutor()
engine = make_engine(success_poster, reader={ref: jwt_sentinel}, clock=success_clock, executor=success_executor)
spec = make_spec(assertion_ref=ref)
minted = mint(engine, spec)
endpoint_error = make_engine(
ScriptedPoster([httpx.Response(400, json={"error": "invalid_grant"})]), reader={ref: jwt_sentinel}
).get_token(spec)
transport_error = make_engine(RaisingPoster(RuntimeError("boom")), reader={ref: jwt_sentinel}).get_token(spec)
malformed_error = make_engine(
ScriptedPoster([httpx.Response(200, json={"unexpected": "shape"})]), reader={ref: jwt_sentinel}
).get_token(spec)
oversized_error = make_engine(
ScriptedPoster([token_response()]), reader={ref: jwt_sentinel + "x" * MAX_ASSERTION_BYTES}
).get_token(spec)
insecure_error = make_engine(ScriptedPoster([token_response()]), reader={ref: jwt_sentinel}).get_token(
make_spec(assertion_ref=ref, token_url="http://token.example/v1/oauth/token")
)
success_poster._responses = [httpx.Response(500, json={"error": "server_error"})]
success_clock.now = success_clock.now + 3600 - 100.0
stale = engine.get_token(spec)
success_executor.run_all()
audited_values = [
str(minted),
repr(minted),
str(minted.access_token),
repr(minted.access_token),
str(endpoint_error),
repr(endpoint_error),
str(transport_error),
repr(transport_error),
str(malformed_error),
repr(malformed_error),
str(oversized_error),
repr(oversized_error),
str(insecure_error),
repr(insecure_error),
str(stale),
repr(stale),
caplog.text,
]
assert isinstance(oversized_error, AssertionSourceError)
assert oversized_error.kind == "oversized"
for value in audited_values:
assert jwt_sentinel not in value
assert token_sentinel not in value
class TestAssertionGuards:
@pytest.mark.parametrize(
"assertion_value,expected_kind",
[
("x" * (MAX_ASSERTION_BYTES + 1), "oversized"),
(" \n\t ", "empty"),
(None, "missing"),
],
)
def test_bad_assertion_values(self, assertion_value: str | None, expected_kind: str):
poster = ScriptedPoster([token_response()])
engine = make_engine(poster, reader=lambda ref: assertion_value)
result = engine.get_token(make_spec())
assert isinstance(result, AssertionSourceError)
assert result.kind == expected_kind
assert result.source_ref == DEFAULT_REF
assert len(poster.requests) == 0
@pytest.mark.parametrize(
"raised,expected_kind",
[
(OidcPathNotAllowedError("path outside allowed credential directories"), "disallowed_path"),
(ValueError("Environment variable ANTHROPIC_IDENTITY_TOKEN not found"), "unreadable"),
(OSError("permission denied"), "unreadable"),
],
)
def test_raising_reader(self, raised: Exception, expected_kind: str):
poster = ScriptedPoster([token_response()])
def reader(ref: str) -> str | None:
raise raised
result = make_engine(poster, reader=reader).get_token(make_spec())
assert isinstance(result, AssertionSourceError)
assert result.kind == expected_kind
assert len(poster.requests) == 0
class TestOidcFilePathAllowlistRaisesTypedError:
"""The engine classifies assertion-source failures by exception type (see
TestAssertionGuards.test_raising_reader); that classification only works if the real
oidc/file allowlist actually raises OidcPathNotAllowedError rather than a bare ValueError."""
def test_out_of_allowlist_absolute_path(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.delenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", raising=False)
with pytest.raises(OidcPathNotAllowedError):
_resolve_oidc_file_path("/etc/not-a-credential-dir/token")
def test_relative_path(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.delenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", raising=False)
with pytest.raises(OidcPathNotAllowedError):
_resolve_oidc_file_path("relative/token/path")
class TestHttpsEnforcement:
def test_plain_http_rejected_host_only_zero_posts(self):
poster = ScriptedPoster([token_response()])
engine = make_engine(poster)
result = engine.get_token(make_spec(token_url="http://token.example/v1/oauth/token"))
assert result == InsecureTokenUrl(host="token.example")
assert "/v1/oauth/token" not in str(result)
assert len(poster.requests) == 0
@pytest.mark.parametrize(
"url",
[
"http://localhost:8080/v1/oauth/token",
"http://127.0.0.1/v1/oauth/token",
"http://[::1]/v1/oauth/token",
],
)
def test_localhost_http_allowed(self, url: str):
poster = ScriptedPoster([token_response()])
engine = make_engine(poster)
result = engine.get_token(make_spec(token_url=url))
assert isinstance(result, MintedToken)
assert len(poster.requests) == 1
def test_cache_key_semantics():
poster = ScriptedPoster([token_response()])
assertions = {DEFAULT_REF: DEFAULT_ASSERTION, "oidc/env/OTHER": "other-assertion"}
engine = make_engine(poster, reader=assertions.get)
base_spec = make_spec()
mint(engine, base_spec)
mint(engine, make_spec(cache_key_identity=("fdrl_1", "org-1", "svc-2", "")))
mint(engine, make_spec(token_url="https://other.example/v1/oauth/token"))
mint(engine, make_spec(assertion_ref="oidc/env/OTHER"))
assert len(poster.requests) == 4
assertions[DEFAULT_REF] = "rotated-assertion"
cached = mint(engine, base_spec)
assert len(poster.requests) == 4
assert cached.access_token.get_secret_value() == "sk-ant-oat01-minted"
def test_bounded_eviction():
clock = FakeClock()
class PerCallPoster:
def __init__(self) -> None:
self.calls = 0
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
self.calls += 1
body = json.loads(content)
expires_in = 3600 + int(body["organization_id"].split("-")[1])
return token_response(f"token-{body['organization_id']}", expires_in=expires_in)
poster = PerCallPoster()
engine = make_engine(poster, clock=clock, max_entries=64)
def spec_for(index: int) -> TokenExchangeSpec:
return make_spec(
static_body={
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
"federation_rule_id": "fdrl_1",
"organization_id": f"org-{index}",
},
cache_key_identity=("fdrl_1", f"org-{index}", "", ""),
)
for index in range(65):
mint(engine, spec_for(index))
assert poster.calls == 65
mint(engine, spec_for(0))
assert poster.calls == 66, "the earliest-expiring entry (index 0) should have been evicted"
mint(engine, spec_for(2))
assert poster.calls == 66, "a later-expiring entry should still be cached"
mint(engine, spec_for(1))
assert poster.calls == 67, "re-inserting index 0 should have evicted the next earliest-expiring entry"
@pytest.mark.parametrize("expires_in", [None, 0, -5])
def test_missing_or_nonsense_expires_in_gets_fallback_ttl(expires_in: int | None):
poster = ScriptedPoster(
[token_response("short-lived", expires_in=expires_in), token_response("reminted", expires_in=3600)]
)
clock = FakeClock(start=1_000.0)
engine = make_engine(poster, clock=clock)
spec = make_spec()
first = mint(engine, spec)
assert first.expires_at == 1_000.0 + FALLBACK_TOKEN_TTL_SECONDS
clock.advance(FALLBACK_TOKEN_TTL_SECONDS + 1.0)
second = mint(engine, spec)
assert second.access_token.get_secret_value() == "reminted"
assert len(poster.requests) == 2, "a token without a sane expires_in must never be cached forever"
async def test_aget_token_loop_responsive():
class SleepingPoster:
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
time.sleep(0.3)
return token_response()
engine = make_engine(SleepingPoster())
spec = make_spec()
ticks = {"count": 0}
stop = asyncio.Event()
async def ticker() -> None:
while not stop.is_set():
ticks["count"] += 1
await asyncio.sleep(0.01)
ticker_task = asyncio.create_task(ticker())
result = await engine.aget_token(spec)
stop.set()
await ticker_task
assert isinstance(result, MintedToken)
assert result.access_token.get_secret_value() == "sk-ant-oat01-minted"
assert ticks["count"] >= 5, "the event loop was blocked during aget_token"
sync_result = engine.get_token(spec)
assert sync_result == result
def test_invalidate_forces_refresh():
poster = ScriptedPoster([token_response("token-1", expires_in=3600), token_response("token-2", expires_in=3600)])
engine = make_engine(poster)
spec = make_spec()
first = mint(engine, spec)
assert first.access_token.get_secret_value() == "token-1"
engine.invalidate(spec)
second = mint(engine, spec)
assert second.access_token.get_secret_value() == "token-2"
assert len(poster.requests) == 2
third = mint(engine, spec)
assert third.access_token.get_secret_value() == "token-2"
assert len(poster.requests) == 2, "force_refresh must be one-shot"
def test_invalidate_unknown_spec_is_noop():
poster = ScriptedPoster([token_response()])
engine = make_engine(poster)
engine.invalidate(make_spec())
assert len(poster.requests) == 0
def test_advisory_failure_wakes_expired_follower_to_re_lead():
poster = ScriptedPoster(
[
token_response("initial-token", expires_in=3600),
httpx.Response(500, json={"error": "server_error"}),
token_response("recovered-token", expires_in=3600),
]
)
clock = FakeClock(start=1_000.0)
executor = ManualExecutor()
engine = make_engine(poster, clock=clock, executor=executor)
spec = make_spec()
mint(engine, spec)
clock.now = 1_000.0 + 3600 - 100.0
mint(engine, spec)
assert len(executor.pending) == 1
clock.advance(200.0)
results: list[ExchangeResult] = []
follower = threading.Thread(target=lambda: results.append(engine.get_token(spec)), daemon=True)
follower.start()
time.sleep(0.3)
executor.run_all()
follower.join(timeout=10)
assert not follower.is_alive()
assert len(results) == 1
result = results[0]
assert isinstance(result, MintedToken), f"follower was handed {result!r} instead of re-leading a fresh mint"
assert result.access_token.get_secret_value() == "recovered-token"
assert len(poster.requests) == 3
class TwoAttemptGatedPoster:
"""401 on the first attempt, then blocks the leader's retry until released."""
def __init__(self) -> None:
self.entered_second = threading.Event()
self.release = threading.Event()
self.calls = 0
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
self.calls += 1
if self.calls == 1:
return httpx.Response(401, json={"error": "invalid_grant"})
self.entered_second.set()
assert self.release.wait(timeout=30)
return token_response("slow-leader-token")
def test_follower_budget_outlasts_slow_two_attempt_leader():
poster = TwoAttemptGatedPoster()
engine = make_engine(poster)
spec = make_spec(timeout_seconds=1.0)
leader_results: list[ExchangeResult] = []
leader = threading.Thread(target=lambda: leader_results.append(engine.get_token(spec)), daemon=True)
leader.start()
assert poster.entered_second.wait(timeout=10)
follower_results: list[ExchangeResult] = []
follower = threading.Thread(target=lambda: follower_results.append(engine.get_token(spec)), daemon=True)
follower.start()
time.sleep(6.5)
poster.release.set()
leader.join(timeout=10)
follower.join(timeout=10)
assert leader_results and isinstance(leader_results[0], MintedToken)
assert follower_results, "follower never returned"
follower_result = follower_results[0]
assert isinstance(follower_result, MintedToken), (
f"follower gave up before the leader's two-attempt worst case: {follower_result!r}"
)
assert follower_result.access_token.get_secret_value() == "slow-leader-token"
class FailThenGatePoster:
"""500 on the first call, then blocks until released before succeeding."""
def __init__(self) -> None:
self.entered_gate = threading.Event()
self.release = threading.Event()
self.calls = 0
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
self.calls += 1
if self.calls == 1:
return httpx.Response(500, json={"error": "server_error"})
self.entered_gate.set()
assert self.release.wait(timeout=30)
return token_response("round-two-token")
def test_new_round_timed_out_follower_never_returns_previous_rounds_error():
poster = FailThenGatePoster()
clock = FakeClock()
engine = make_engine(poster, clock=clock)
spec = make_spec(timeout_seconds=0.05)
first = engine.get_token(spec)
assert isinstance(first, TokenEndpointError)
clock.advance(ADVISORY_REFRESH_BACKOFF_SECONDS + 1.0)
leader = threading.Thread(target=lambda: engine.get_token(spec), daemon=True)
leader.start()
assert poster.entered_gate.wait(timeout=10)
follower_result = engine.get_token(spec)
assert isinstance(follower_result, TokenTransportError), (
f"timed-out follower returned the previous round's error: {follower_result!r}"
)
assert "timed out" in follower_result.detail
poster.release.set()
leader.join(timeout=10)
def test_lead_backoff_fails_fast_within_window_and_expires_after():
poster = ScriptedPoster([httpx.Response(500, json={"error": "server_error"})])
clock = FakeClock()
engine = make_engine(poster, clock=clock)
spec = make_spec()
first = engine.get_token(spec)
assert isinstance(first, TokenEndpointError)
assert len(poster.requests) == 1
clock.advance(ADVISORY_REFRESH_BACKOFF_SECONDS - 1.0)
second = engine.get_token(spec)
assert second == first
assert len(poster.requests) == 1, "a request inside the backoff window must make zero POSTs"
clock.advance(1.0)
third = engine.get_token(spec)
assert isinstance(third, TokenEndpointError)
assert len(poster.requests) == 2
def test_invalidate_bypasses_lead_backoff():
poster = ScriptedPoster(
[httpx.Response(500, json={"error": "server_error"}), token_response("post-invalidate", expires_in=3600)]
)
clock = FakeClock()
engine = make_engine(poster, clock=clock)
spec = make_spec()
first = engine.get_token(spec)
assert isinstance(first, TokenEndpointError)
engine.invalidate(spec)
second = engine.get_token(spec)
assert isinstance(second, MintedToken)
assert second.access_token.get_secret_value() == "post-invalidate"
assert len(poster.requests) == 2

View file

@ -374,7 +374,7 @@ async def test_async_anthropic_messages_handler_extra_headers():
# Mock the config
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
mock_config.avalidate_anthropic_messages_environment = AsyncMock(
return_value=({"x-api-key": "test-key"}, "https://api.anthropic.com")
)
mock_config.transform_anthropic_messages_request = Mock(
@ -421,7 +421,7 @@ async def test_async_anthropic_messages_handler_extra_headers():
captured_headers.update(kwargs.get("headers", {}))
return ({"x-api-key": "test-key"}, "https://api.anthropic.com")
mock_config.validate_anthropic_messages_environment = capture_validate
mock_config.avalidate_anthropic_messages_environment = AsyncMock(side_effect=capture_validate)
try:
await handler.async_anthropic_messages_handler(
@ -677,7 +677,7 @@ async def test_async_anthropic_messages_handler_passes_litellm_metadata():
handler = BaseLLMHTTPHandler()
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
mock_config.avalidate_anthropic_messages_environment = AsyncMock(
return_value=({"x-api-key": "test-key"}, "https://api.anthropic.com")
)
mock_config.transform_anthropic_messages_request = Mock(
@ -756,7 +756,7 @@ async def test_async_anthropic_messages_handler_forwards_router_model_info():
handler = BaseLLMHTTPHandler()
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
mock_config.avalidate_anthropic_messages_environment = AsyncMock(
return_value=({"x-api-key": "test-key"}, "https://api.anthropic.com")
)
mock_config.transform_anthropic_messages_request = Mock(
@ -848,7 +848,7 @@ async def test_async_anthropic_messages_handler_header_priority():
captured_headers.update(kwargs.get("headers", {}))
return ({"x-api-key": "test-key"}, "https://api.anthropic.com")
mock_config.validate_anthropic_messages_environment = capture_validate
mock_config.avalidate_anthropic_messages_environment = AsyncMock(side_effect=capture_validate)
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude-3-opus-20240229", "messages": []}
)
@ -887,7 +887,7 @@ async def test_async_anthropic_messages_handler_drops_top_level_and_nested_param
handler = BaseLLMHTTPHandler()
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
mock_config.avalidate_anthropic_messages_environment = AsyncMock(
return_value=({"x-api-key": "test-key"}, "https://api.anthropic.com")
)
@ -1100,9 +1100,7 @@ def test_sync_delete_responses_sets_json_content_type():
({}, True, None, None),
],
)
def test_resolve_anthropic_messages_timeout(
monkeypatch, litellm_params_kwargs, stream, global_timeout, expected
):
def test_resolve_anthropic_messages_timeout(monkeypatch, litellm_params_kwargs, stream, global_timeout, expected):
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
if global_timeout is None:
@ -1118,9 +1116,7 @@ def test_resolve_anthropic_messages_timeout(
)
else:
monkeypatch.setattr("litellm.request_timeout", global_timeout, raising=False)
monkeypatch.setattr(
"litellm.request_timeout_explicitly_set", True, raising=False
)
monkeypatch.setattr("litellm.request_timeout_explicitly_set", True, raising=False)
resolved = BaseLLMHTTPHandler._resolve_anthropic_messages_timeout(
litellm_params=GenericLiteLLMParams(**litellm_params_kwargs),
@ -1141,13 +1137,11 @@ async def test_async_anthropic_messages_handler_forwards_request_timeout(monkeyp
handler = BaseLLMHTTPHandler()
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
mock_config.avalidate_anthropic_messages_environment = AsyncMock(
return_value=({"x-api-key": "k"}, "https://api.anthropic.com")
)
mock_config.should_filter_anthropic_beta_headers = Mock(return_value=False)
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude", "messages": []}
)
mock_config.transform_anthropic_messages_request = Mock(return_value={"model": "claude", "messages": []})
mock_config.get_complete_url = Mock(return_value="https://api.anthropic.com/v1/messages")
mock_config.sign_request = Mock(return_value=({"x-api-key": "k"}, None))
mock_config.max_retry_on_anthropic_messages_http_error = 1
@ -1189,13 +1183,11 @@ async def test_async_anthropic_messages_handler_forwards_stream_timeout(monkeypa
handler = BaseLLMHTTPHandler()
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
mock_config.avalidate_anthropic_messages_environment = AsyncMock(
return_value=({"x-api-key": "k"}, "https://api.anthropic.com")
)
mock_config.should_filter_anthropic_beta_headers = Mock(return_value=False)
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude", "messages": []}
)
mock_config.transform_anthropic_messages_request = Mock(return_value={"model": "claude", "messages": []})
mock_config.get_complete_url = Mock(return_value="https://api.anthropic.com/v1/messages")
mock_config.sign_request = Mock(return_value=({"x-api-key": "k"}, None))
mock_config.max_retry_on_anthropic_messages_http_error = 1
@ -1597,7 +1589,7 @@ async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks(
handler = BaseLLMHTTPHandler()
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
mock_config.avalidate_anthropic_messages_environment = AsyncMock(
return_value=({"x-api-key": "sk-test"}, "https://api.anthropic.com")
)
mock_config.transform_anthropic_messages_request = Mock(
@ -1605,7 +1597,13 @@ async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks(
)
mock_config.sign_request = Mock(return_value=({}, None))
fake_raw_response = {"id": "msg_1", "type": "message", "role": "assistant", "content": [], "stop_reason": "end_turn"}
fake_raw_response = {
"id": "msg_1",
"type": "message",
"role": "assistant",
"content": [],
"stop_reason": "end_turn",
}
mock_config.transform_anthropic_messages_response = Mock(return_value=fake_raw_response)
mock_logging_obj = Mock()
@ -1625,10 +1623,17 @@ async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks(
mock_httpx_response.status_code = 200
with (
patch.object(handler, "_async_post_anthropic_messages_with_http_error_retry", new=AsyncMock(return_value=mock_httpx_response)),
patch.object(
handler,
"_async_post_anthropic_messages_with_http_error_retry",
new=AsyncMock(return_value=mock_httpx_response),
),
patch.object(handler, "_call_agentic_completion_hooks", side_effect=fake_agentic_hooks),
patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client"),
patch("litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers", return_value=None),
patch(
"litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers",
return_value=None,
),
):
result = await handler.async_anthropic_messages_handler(
model="claude-haiku",
@ -1870,7 +1875,9 @@ def test_audio_transcriptions_sends_dict_data_as_json_body():
form-encodes it and silently ignores json=; JSON-body providers (e.g.
Google Speech-to-Text) need an application/json body."""
captured = {}
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_json_transcription_request(captured))))
client = HTTPHandler(
client=httpx.Client(transport=httpx.MockTransport(_capture_json_transcription_request(captured)))
)
response = BaseLLMHTTPHandler().audio_transcriptions(
client=client,
@ -2106,9 +2113,7 @@ async def test_anthropic_invalid_thinking_signature_retry_resigns_bedrock_reques
ok_response = httpx.Response(200, json={"id": "msg_1"}, request=httpx.Request("POST", request_url))
class FakeAsyncClient:
async def post(
self, url, headers, data, stream=False, logging_obj=None, timeout=None
):
async def post(self, url, headers, data, stream=False, logging_obj=None, timeout=None):
posts.append({"headers": dict(headers), "data": data})
return invalid_signature_response if len(posts) == 1 else ok_response
@ -2335,7 +2340,7 @@ async def test_async_anthropic_messages_handler_carries_deployment_vertex_locati
custom_llm_provider="vertex_ai",
)
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
mock_config.avalidate_anthropic_messages_environment = AsyncMock(
return_value=({"authorization": "Bearer t"}, "https://us-east5-aiplatform.googleapis.com")
)
mock_config.transform_anthropic_messages_request = Mock(

View file

@ -27,9 +27,7 @@ FAKE_API_KEY = "sk-ant-test-key-1234"
FAKE_API_BASE = "https://api.anthropic.com"
def _make_mock_response(
json_data: dict, status_code: int = 200, method: str = "POST"
) -> httpx.Response:
def _make_mock_response(json_data: dict, status_code: int = 200, method: str = "POST") -> httpx.Response:
return httpx.Response(
status_code=status_code,
json=json_data,
@ -111,9 +109,7 @@ class TestAnthropicSkillsConfigHeaderValidation:
"litellm.llms.anthropic.common_utils.AnthropicModelInfo.get_api_key",
return_value=FAKE_API_KEY,
):
headers = self.config.validate_environment(
headers={}, litellm_params=self._make_litellm_params()
)
headers = self.config.validate_environment(headers={}, litellm_params=self._make_litellm_params())
assert headers["x-api-key"] == FAKE_API_KEY
def test_sets_anthropic_version_header(self):
@ -121,9 +117,7 @@ class TestAnthropicSkillsConfigHeaderValidation:
"litellm.llms.anthropic.common_utils.AnthropicModelInfo.get_api_key",
return_value=FAKE_API_KEY,
):
headers = self.config.validate_environment(
headers={}, litellm_params=self._make_litellm_params()
)
headers = self.config.validate_environment(headers={}, litellm_params=self._make_litellm_params())
assert headers["anthropic-version"] == "2023-06-01"
def test_sets_skills_beta_header(self):
@ -131,12 +125,12 @@ class TestAnthropicSkillsConfigHeaderValidation:
"litellm.llms.anthropic.common_utils.AnthropicModelInfo.get_api_key",
return_value=FAKE_API_KEY,
):
headers = self.config.validate_environment(
headers={}, litellm_params=self._make_litellm_params()
)
headers = self.config.validate_environment(headers={}, litellm_params=self._make_litellm_params())
assert headers["anthropic-beta"] == ANTHROPIC_SKILLS_API_BETA_VERSION
def test_merges_existing_beta_header_string(self):
def test_merges_existing_beta_header_into_string(self):
"""The merged value must stay a comma-separated string: a list value makes
httpx.Headers raise TypeError when the request is built."""
with patch(
"litellm.llms.anthropic.common_utils.AnthropicModelInfo.get_api_key",
return_value=FAKE_API_KEY,
@ -145,21 +139,26 @@ class TestAnthropicSkillsConfigHeaderValidation:
headers={"anthropic-beta": "other-beta-2024-01-01"},
litellm_params=self._make_litellm_params(),
)
assert isinstance(headers["anthropic-beta"], list)
assert "other-beta-2024-01-01" in headers["anthropic-beta"]
assert ANTHROPIC_SKILLS_API_BETA_VERSION in headers["anthropic-beta"]
assert isinstance(headers["anthropic-beta"], str)
betas = set(headers["anthropic-beta"].split(","))
assert {"other-beta-2024-01-01", ANTHROPIC_SKILLS_API_BETA_VERSION} <= betas
httpx.Headers(headers)
def test_merges_existing_beta_header_list(self):
def test_oauth_key_beta_merges_without_crashing_httpx(self):
"""Regression: an sk-ant-oat/WIF auth header carries its own anthropic-beta;
the old list-building merge produced a Python list that crashed httpx."""
with patch(
"litellm.llms.anthropic.common_utils.AnthropicModelInfo.get_api_key",
return_value=FAKE_API_KEY,
return_value="sk-ant-oat01-fake-skills-token",
):
headers = self.config.validate_environment(
headers={"anthropic-beta": ["other-beta-2024-01-01"]},
litellm_params=self._make_litellm_params(),
headers={}, litellm_params=self._make_litellm_params(api_key=None)
)
assert ANTHROPIC_SKILLS_API_BETA_VERSION in headers["anthropic-beta"]
assert "other-beta-2024-01-01" in headers["anthropic-beta"]
assert headers["authorization"] == "Bearer sk-ant-oat01-fake-skills-token"
assert isinstance(headers["anthropic-beta"], str)
betas = set(headers["anthropic-beta"].split(","))
assert {"oauth-2025-04-20", ANTHROPIC_SKILLS_API_BETA_VERSION} <= betas
httpx.Headers(headers)
def test_does_not_duplicate_beta_header(self):
with patch(
@ -170,11 +169,7 @@ class TestAnthropicSkillsConfigHeaderValidation:
headers={"anthropic-beta": ANTHROPIC_SKILLS_API_BETA_VERSION},
litellm_params=self._make_litellm_params(),
)
beta = headers["anthropic-beta"]
if isinstance(beta, list):
assert beta.count(ANTHROPIC_SKILLS_API_BETA_VERSION) == 1
else:
assert beta == ANTHROPIC_SKILLS_API_BETA_VERSION
assert headers["anthropic-beta"] == ANTHROPIC_SKILLS_API_BETA_VERSION
def test_raises_without_api_key(self):
with patch(
@ -182,9 +177,7 @@ class TestAnthropicSkillsConfigHeaderValidation:
return_value=None,
):
with pytest.raises(ValueError, match="ANTHROPIC_API_KEY"):
self.config.validate_environment(
headers={}, litellm_params=self._make_litellm_params(api_key=None)
)
self.config.validate_environment(headers={}, litellm_params=self._make_litellm_params(api_key=None))
class TestAnthropicSkillsConfigCreateRequestTransformation:
@ -275,9 +268,7 @@ class TestAnthropicSkillsConfigResponseTransformation:
def test_create_skill_response_parses_skill(self):
payload = _make_skill_payload()
raw = _make_mock_response(payload)
skill = self.config.transform_create_skill_response(
raw_response=raw, logging_obj=self.logging_obj
)
skill = self.config.transform_create_skill_response(raw_response=raw, logging_obj=self.logging_obj)
assert isinstance(skill, Skill)
assert skill.id == "skill_abc123"
assert skill.source == "custom"
@ -286,9 +277,7 @@ class TestAnthropicSkillsConfigResponseTransformation:
def test_get_skill_response_parses_skill(self):
payload = _make_skill_payload(id="skill_xyz", display_title="Another")
raw = _make_mock_response(payload, method="GET")
skill = self.config.transform_get_skill_response(
raw_response=raw, logging_obj=self.logging_obj
)
skill = self.config.transform_get_skill_response(raw_response=raw, logging_obj=self.logging_obj)
assert isinstance(skill, Skill)
assert skill.id == "skill_xyz"
assert skill.display_title == "Another"
@ -300,9 +289,7 @@ class TestAnthropicSkillsConfigResponseTransformation:
"next_page": None,
}
raw = _make_mock_response(payload, method="GET")
result = self.config.transform_list_skills_response(
raw_response=raw, logging_obj=self.logging_obj
)
result = self.config.transform_list_skills_response(raw_response=raw, logging_obj=self.logging_obj)
assert isinstance(result, ListSkillsResponse)
assert len(result.data) == 2
assert result.data[0].id == "skill_abc123"
@ -316,18 +303,14 @@ class TestAnthropicSkillsConfigResponseTransformation:
"next_page": "page_token_xyz",
}
raw = _make_mock_response(payload, method="GET")
result = self.config.transform_list_skills_response(
raw_response=raw, logging_obj=self.logging_obj
)
result = self.config.transform_list_skills_response(raw_response=raw, logging_obj=self.logging_obj)
assert result.has_more is True
assert result.next_page == "page_token_xyz"
def test_delete_skill_response_parses_correctly(self):
payload = {"id": "skill_abc123", "type": "skill_deleted"}
raw = _make_mock_response(payload, method="DELETE")
result = self.config.transform_delete_skill_response(
raw_response=raw, logging_obj=self.logging_obj
)
result = self.config.transform_delete_skill_response(raw_response=raw, logging_obj=self.logging_obj)
assert isinstance(result, DeleteSkillResponse)
assert result.id == "skill_abc123"
assert result.type == "skill_deleted"
@ -341,8 +324,6 @@ class TestAnthropicSkillsConfigResponseTransformation:
"type": "skill",
}
raw = _make_mock_response(payload)
skill = self.config.transform_create_skill_response(
raw_response=raw, logging_obj=self.logging_obj
)
skill = self.config.transform_create_skill_response(raw_response=raw, logging_obj=self.logging_obj)
assert skill.display_title is None
assert skill.latest_version is None

View file

@ -1,6 +1,6 @@
{
"LIT001": {
"limit": 22805
"limit": 22803
},
"LIT002": {
"limit": 26873