mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
f005afa146
commit
61a1122420
30 changed files with 4186 additions and 728 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
231
litellm/llms/anthropic/wif.py
Normal file
231
litellm/llms/anthropic/wif.py
Normal 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)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
49
litellm/llms/base_llm/auth/__init__.py
Normal file
49
litellm/llms/base_llm/auth/__init__.py
Normal 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",
|
||||
)
|
||||
508
litellm/llms/base_llm/auth/token_exchange.py
Normal file
508
litellm/llms/base_llm/auth/token_exchange.py
Normal 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()
|
||||
76
litellm/llms/base_llm/auth/types.py
Normal file
76
litellm/llms/base_llm/auth/types.py
Normal 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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -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,
|
||||
|
|
|
|||
637
tests/test_litellm/llms/anthropic/test_anthropic_wif.py
Normal file
637
tests/test_litellm/llms/anthropic/test_anthropic_wif.py
Normal 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"
|
||||
0
tests/test_litellm/llms/base_llm/auth/__init__.py
Normal file
0
tests/test_litellm/llms/base_llm/auth/__init__.py
Normal file
953
tests/test_litellm/llms/base_llm/auth/test_token_exchange.py
Normal file
953
tests/test_litellm/llms/base_llm/auth/test_token_exchange.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22805
|
||||
"limit": 22803
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26873
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue