mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
feat(anthropic): support workload identity federation tokens
Exchanges workload identity tokens for anthropic token at runtime.
This commit is contained in:
parent
73e9071311
commit
7b5f2bcf2b
3 changed files with 729 additions and 9 deletions
|
|
@ -11,6 +11,10 @@ import litellm
|
|||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_file_ids_from_messages,
|
||||
)
|
||||
from litellm.llms.anthropic.workload_identity_federation import (
|
||||
AnthropicWorkloadIdentityFederationError,
|
||||
exchange_anthropic_workload_identity_federation_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 (
|
||||
|
|
@ -374,7 +378,8 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
"computer_20241022": "computer-use-2024-10-22",
|
||||
}
|
||||
return computer_tool_beta_mapping.get(
|
||||
computer_tool_version, "computer-use-2024-10-22" # Default fallback
|
||||
computer_tool_version,
|
||||
"computer-use-2024-10-22", # Default fallback
|
||||
)
|
||||
|
||||
def get_anthropic_beta_list(
|
||||
|
|
@ -523,7 +528,14 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
# Resolve auth_token from ANTHROPIC_AUTH_TOKEN if api_key is not set
|
||||
auth_token: Optional[str] = None
|
||||
if api_key is None:
|
||||
auth_token = AnthropicModelInfo.get_auth_token()
|
||||
try:
|
||||
auth_token = AnthropicModelInfo.get_auth_token(api_base=api_base)
|
||||
except AnthropicWorkloadIdentityFederationError as exc:
|
||||
raise litellm.AuthenticationError(
|
||||
message=str(exc),
|
||||
llm_provider="anthropic",
|
||||
model=model,
|
||||
)
|
||||
if 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",
|
||||
|
|
@ -594,30 +606,45 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
return api_key or get_secret_str("ANTHROPIC_API_KEY")
|
||||
|
||||
@staticmethod
|
||||
def get_auth_token(auth_token: Optional[str] = None) -> Optional[str]:
|
||||
"""Get auth token from ANTHROPIC_AUTH_TOKEN env var.
|
||||
def get_auth_token(
|
||||
auth_token: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
"""Get auth token from ANTHROPIC_AUTH_TOKEN env var, or exchange one
|
||||
via Workload Identity Federation when the ANTHROPIC_FEDERATION_*
|
||||
env vars are set.
|
||||
|
||||
Unlike api_key (which uses X-Api-Key header), auth_token uses
|
||||
Authorization: Bearer header, matching the official Anthropic SDK behavior.
|
||||
"""
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
return auth_token or get_secret_str("ANTHROPIC_AUTH_TOKEN")
|
||||
return (
|
||||
auth_token
|
||||
or get_secret_str("ANTHROPIC_AUTH_TOKEN")
|
||||
or exchange_anthropic_workload_identity_federation_token(
|
||||
api_base=AnthropicModelInfo.get_api_base(api_base),
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def get_auth_header(api_key: Optional[str] = None) -> Optional[dict]:
|
||||
"""Resolve Anthropic credentials and return the appropriate auth header dict.
|
||||
|
||||
Checks ANTHROPIC_API_KEY first (-> x-api-key), then
|
||||
ANTHROPIC_AUTH_TOKEN (-> Authorization: Bearer).
|
||||
Returns None if neither is available.
|
||||
Checks ANTHROPIC_API_KEY first (-> x-api-key), then ANTHROPIC_AUTH_TOKEN
|
||||
/ exchanged Workload Identity Federation token via get_auth_token
|
||||
(-> Authorization: Bearer).
|
||||
Returns None if no credential source is available.
|
||||
"""
|
||||
resolved_key = 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 {"x-api-key": resolved_key}
|
||||
auth_token = AnthropicModelInfo.get_auth_token()
|
||||
try:
|
||||
auth_token = AnthropicModelInfo.get_auth_token()
|
||||
except AnthropicWorkloadIdentityFederationError:
|
||||
return None
|
||||
if auth_token is not None:
|
||||
return {"authorization": f"Bearer {auth_token}"}
|
||||
return None
|
||||
|
|
|
|||
276
litellm/llms/anthropic/workload_identity_federation.py
Normal file
276
litellm/llms/anthropic/workload_identity_federation.py
Normal file
|
|
@ -0,0 +1,276 @@
|
|||
"""Anthropic Workload Identity Federation token exchange.
|
||||
|
||||
See https://platform.claude.com/docs/en/manage-claude/wif-reference.
|
||||
"""
|
||||
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
ANTHROPIC_WORKLOAD_IDENTITY_FEDERATION_GRANT_TYPE = (
|
||||
"urn:ietf:params:oauth:grant-type:jwt-bearer"
|
||||
)
|
||||
ANTHROPIC_WORKLOAD_IDENTITY_FEDERATION_TOKEN_PATH = "/v1/oauth/token"
|
||||
|
||||
# Two-tier refresh thresholds, matching the official Anthropic SDKs.
|
||||
ANTHROPIC_WORKLOAD_IDENTITY_FEDERATION_ADVISORY_REFRESH_SECONDS = 120
|
||||
ANTHROPIC_WORKLOAD_IDENTITY_FEDERATION_MANDATORY_REFRESH_SECONDS = 30
|
||||
|
||||
|
||||
class AnthropicWorkloadIdentityFederationError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class AnthropicWorkloadIdentityFederationCredentials:
|
||||
federation_rule_id: str
|
||||
organization_id: str
|
||||
service_account_id: str
|
||||
workspace_id: Optional[str] = None
|
||||
identity_token: Optional[str] = None
|
||||
identity_token_file: Optional[str] = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
required = {
|
||||
"federation_rule_id": self.federation_rule_id,
|
||||
"organization_id": self.organization_id,
|
||||
"service_account_id": self.service_account_id,
|
||||
}
|
||||
missing = [name for name, value in required.items() if not value]
|
||||
if missing:
|
||||
raise AnthropicWorkloadIdentityFederationError(
|
||||
f"Anthropic workload identity federation credentials missing required fields: {missing}"
|
||||
)
|
||||
if not self.identity_token and not self.identity_token_file:
|
||||
raise AnthropicWorkloadIdentityFederationError(
|
||||
"Anthropic workload identity federation credentials must set either "
|
||||
"`identity_token` or `identity_token_file`."
|
||||
)
|
||||
|
||||
|
||||
def get_workload_identity_federation_credentials_from_env() -> (
|
||||
Optional[AnthropicWorkloadIdentityFederationCredentials]
|
||||
):
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
federation_rule_id = get_secret_str("ANTHROPIC_FEDERATION_RULE_ID")
|
||||
organization_id = get_secret_str("ANTHROPIC_ORGANIZATION_ID")
|
||||
service_account_id = get_secret_str("ANTHROPIC_SERVICE_ACCOUNT_ID")
|
||||
identity_token = get_secret_str("ANTHROPIC_IDENTITY_TOKEN")
|
||||
identity_token_file = get_secret_str("ANTHROPIC_IDENTITY_TOKEN_FILE")
|
||||
workspace_id = get_secret_str("ANTHROPIC_WORKSPACE_ID")
|
||||
|
||||
if not all((federation_rule_id, organization_id, service_account_id)):
|
||||
return None
|
||||
if not identity_token and not identity_token_file:
|
||||
return None
|
||||
|
||||
return AnthropicWorkloadIdentityFederationCredentials(
|
||||
federation_rule_id=federation_rule_id, # type: ignore[arg-type]
|
||||
organization_id=organization_id, # type: ignore[arg-type]
|
||||
service_account_id=service_account_id, # type: ignore[arg-type]
|
||||
workspace_id=workspace_id,
|
||||
identity_token=identity_token,
|
||||
identity_token_file=identity_token_file,
|
||||
)
|
||||
|
||||
|
||||
class AnthropicWorkloadIdentityFederationTokenProvider:
|
||||
def __init__(
|
||||
self,
|
||||
credentials: AnthropicWorkloadIdentityFederationCredentials,
|
||||
api_base: str = "https://api.anthropic.com",
|
||||
http_client: Optional[httpx.Client] = None,
|
||||
timeout: float = 30.0,
|
||||
) -> None:
|
||||
self._credentials = credentials
|
||||
self._api_base = api_base.rstrip("/")
|
||||
self._http_client = http_client
|
||||
self._timeout = timeout
|
||||
self._lock = threading.Lock()
|
||||
self._access_token: Optional[str] = None
|
||||
self._expires_at: float = 0.0
|
||||
|
||||
@property
|
||||
def credentials(self) -> AnthropicWorkloadIdentityFederationCredentials:
|
||||
return self._credentials
|
||||
|
||||
def get_token(self, assertion: Optional[str] = None) -> str:
|
||||
with self._lock:
|
||||
now = time.time()
|
||||
time_to_expiry = self._expires_at - now
|
||||
|
||||
if (
|
||||
assertion is None
|
||||
and self._access_token
|
||||
and time_to_expiry
|
||||
> ANTHROPIC_WORKLOAD_IDENTITY_FEDERATION_ADVISORY_REFRESH_SECONDS
|
||||
):
|
||||
return self._access_token
|
||||
|
||||
mandatory = (
|
||||
assertion is not None
|
||||
or self._access_token is None
|
||||
or time_to_expiry
|
||||
<= ANTHROPIC_WORKLOAD_IDENTITY_FEDERATION_MANDATORY_REFRESH_SECONDS
|
||||
)
|
||||
|
||||
try:
|
||||
self._refresh_locked(assertion=assertion)
|
||||
except Exception as exc:
|
||||
if mandatory:
|
||||
raise
|
||||
verbose_logger.warning(
|
||||
"Anthropic workload identity federation advisory refresh failed; reusing cached token: %s",
|
||||
exc,
|
||||
)
|
||||
|
||||
assert self._access_token is not None
|
||||
return self._access_token
|
||||
|
||||
def _resolve_assertion(self, override: Optional[str]) -> str:
|
||||
if override:
|
||||
return override
|
||||
if self._credentials.identity_token:
|
||||
return self._credentials.identity_token
|
||||
path = self._credentials.identity_token_file
|
||||
if not path:
|
||||
raise AnthropicWorkloadIdentityFederationError(
|
||||
"Anthropic workload identity federation credentials have neither "
|
||||
"`identity_token` nor `identity_token_file` set."
|
||||
)
|
||||
# Re-read on each exchange: cloud workload identity providers
|
||||
# rotate the projected JWT in place, so caching its contents
|
||||
# would silently send a stale assertion.
|
||||
try:
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
assertion = f.read().strip()
|
||||
except OSError as exc:
|
||||
raise AnthropicWorkloadIdentityFederationError(
|
||||
f"Failed to read ANTHROPIC_IDENTITY_TOKEN_FILE at {path!r}: {exc}"
|
||||
) from exc
|
||||
if not assertion:
|
||||
raise AnthropicWorkloadIdentityFederationError(
|
||||
f"ANTHROPIC_IDENTITY_TOKEN_FILE at {path!r} is empty"
|
||||
)
|
||||
return assertion
|
||||
|
||||
def _refresh_locked(self, assertion: Optional[str] = None) -> None:
|
||||
resolved_assertion = self._resolve_assertion(assertion)
|
||||
url = f"{self._api_base}{ANTHROPIC_WORKLOAD_IDENTITY_FEDERATION_TOKEN_PATH}"
|
||||
data = {
|
||||
"grant_type": ANTHROPIC_WORKLOAD_IDENTITY_FEDERATION_GRANT_TYPE,
|
||||
"assertion": resolved_assertion,
|
||||
"federation_rule_id": self._credentials.federation_rule_id,
|
||||
"organization_id": self._credentials.organization_id,
|
||||
"service_account_id": self._credentials.service_account_id,
|
||||
}
|
||||
if self._credentials.workspace_id:
|
||||
data["workspace_id"] = self._credentials.workspace_id
|
||||
headers = {
|
||||
"content-type": "application/x-www-form-urlencoded",
|
||||
"accept": "application/json",
|
||||
}
|
||||
|
||||
try:
|
||||
if self._http_client is not None:
|
||||
response = self._http_client.post(url=url, data=data, headers=headers)
|
||||
else:
|
||||
with httpx.Client(timeout=self._timeout) as client:
|
||||
response = client.post(url=url, data=data, headers=headers)
|
||||
except httpx.HTTPError as exc:
|
||||
raise AnthropicWorkloadIdentityFederationError(
|
||||
f"Anthropic workload identity federation token exchange failed: {exc}"
|
||||
) from exc
|
||||
|
||||
if response.status_code >= 400:
|
||||
raise AnthropicWorkloadIdentityFederationError(
|
||||
f"Anthropic workload identity federation token exchange returned HTTP {response.status_code}: {response.text}"
|
||||
)
|
||||
|
||||
try:
|
||||
payload = response.json()
|
||||
except ValueError as exc:
|
||||
raise AnthropicWorkloadIdentityFederationError(
|
||||
f"Anthropic workload identity federation token exchange returned non-JSON response: {response.text}"
|
||||
) from exc
|
||||
|
||||
access_token = payload.get("access_token")
|
||||
if not access_token:
|
||||
raise AnthropicWorkloadIdentityFederationError(
|
||||
f"Anthropic workload identity federation token exchange response missing 'access_token': {payload}"
|
||||
)
|
||||
|
||||
# Fall back to a short window if expires_in is missing/invalid —
|
||||
# otherwise we'd cache forever on a misbehaving response.
|
||||
expires_in = payload.get("expires_in")
|
||||
try:
|
||||
expires_in_int = int(expires_in) if expires_in is not None else 0
|
||||
except (TypeError, ValueError):
|
||||
expires_in_int = 0
|
||||
if expires_in_int <= 0:
|
||||
expires_in_int = (
|
||||
ANTHROPIC_WORKLOAD_IDENTITY_FEDERATION_MANDATORY_REFRESH_SECONDS * 2
|
||||
)
|
||||
|
||||
self._access_token = access_token
|
||||
self._expires_at = time.time() + expires_in_int
|
||||
|
||||
|
||||
_PROVIDER_CACHE: dict = {}
|
||||
_PROVIDER_CACHE_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def _provider_cache_key(
|
||||
credentials: AnthropicWorkloadIdentityFederationCredentials, api_base: str
|
||||
) -> tuple:
|
||||
return (
|
||||
credentials.federation_rule_id,
|
||||
credentials.organization_id,
|
||||
credentials.service_account_id,
|
||||
credentials.workspace_id,
|
||||
credentials.identity_token,
|
||||
credentials.identity_token_file,
|
||||
api_base,
|
||||
)
|
||||
|
||||
|
||||
def get_or_create_workload_identity_federation_provider(
|
||||
credentials: AnthropicWorkloadIdentityFederationCredentials,
|
||||
api_base: str = "https://api.anthropic.com",
|
||||
) -> AnthropicWorkloadIdentityFederationTokenProvider:
|
||||
key = _provider_cache_key(credentials, api_base)
|
||||
with _PROVIDER_CACHE_LOCK:
|
||||
provider = _PROVIDER_CACHE.get(key)
|
||||
if provider is None:
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=credentials, api_base=api_base
|
||||
)
|
||||
_PROVIDER_CACHE[key] = provider
|
||||
return provider
|
||||
|
||||
|
||||
def exchange_anthropic_workload_identity_federation_token(
|
||||
api_base: Optional[str] = None,
|
||||
credentials: Optional[AnthropicWorkloadIdentityFederationCredentials] = None,
|
||||
assertion: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
resolved = credentials or get_workload_identity_federation_credentials_from_env()
|
||||
if resolved is None:
|
||||
return None
|
||||
provider = get_or_create_workload_identity_federation_provider(
|
||||
credentials=resolved,
|
||||
api_base=(api_base or "https://api.anthropic.com").rstrip("/"),
|
||||
)
|
||||
return provider.get_token(assertion=assertion)
|
||||
|
||||
|
||||
def reset_workload_identity_federation_provider_cache() -> None:
|
||||
# Intended for tests — prevents stale providers leaking across cases.
|
||||
with _PROVIDER_CACHE_LOCK:
|
||||
_PROVIDER_CACHE.clear()
|
||||
|
|
@ -0,0 +1,417 @@
|
|||
"""Tests for Anthropic Workload Identity Federation support."""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from typing import Optional
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../..")))
|
||||
|
||||
from litellm.llms.anthropic.workload_identity_federation import (
|
||||
ANTHROPIC_WORKLOAD_IDENTITY_FEDERATION_GRANT_TYPE,
|
||||
AnthropicWorkloadIdentityFederationCredentials,
|
||||
AnthropicWorkloadIdentityFederationError,
|
||||
AnthropicWorkloadIdentityFederationTokenProvider,
|
||||
get_workload_identity_federation_credentials_from_env,
|
||||
exchange_anthropic_workload_identity_federation_token,
|
||||
reset_workload_identity_federation_provider_cache,
|
||||
)
|
||||
|
||||
|
||||
FAKE_REQUIRED_CREDS_KWARGS = dict(
|
||||
federation_rule_id="fdrl_test_rule",
|
||||
organization_id="org_test",
|
||||
service_account_id="svac_test",
|
||||
)
|
||||
|
||||
|
||||
def _make_creds(
|
||||
*,
|
||||
identity_token_file: Optional[str] = None,
|
||||
identity_token: Optional[str] = None,
|
||||
workspace_id: Optional[str] = "wrkspc_test",
|
||||
) -> AnthropicWorkloadIdentityFederationCredentials:
|
||||
return AnthropicWorkloadIdentityFederationCredentials(
|
||||
identity_token_file=identity_token_file,
|
||||
identity_token=identity_token,
|
||||
workspace_id=workspace_id,
|
||||
**FAKE_REQUIRED_CREDS_KWARGS,
|
||||
)
|
||||
|
||||
|
||||
def _build_response(*, access_token: str = "anth_at_test", expires_in: Optional[int] = 3600) -> httpx.Response:
|
||||
payload: dict = {"access_token": access_token}
|
||||
if expires_in is not None:
|
||||
payload["expires_in"] = expires_in
|
||||
return httpx.Response(status_code=200, json=payload, request=httpx.Request("POST", "https://x"))
|
||||
|
||||
|
||||
class _StubClient:
|
||||
def __init__(
|
||||
self,
|
||||
response: Optional[httpx.Response] = None,
|
||||
exc: Optional[Exception] = None,
|
||||
):
|
||||
self.calls: list = []
|
||||
self._response = response or _build_response()
|
||||
self._exc = exc
|
||||
|
||||
def post(self, url: str, data: dict, headers: dict) -> httpx.Response:
|
||||
self.calls.append({"url": url, "data": data, "headers": headers})
|
||||
if self._exc is not None:
|
||||
raise self._exc
|
||||
return self._response
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_provider_cache():
|
||||
reset_workload_identity_federation_provider_cache()
|
||||
yield
|
||||
reset_workload_identity_federation_provider_cache()
|
||||
|
||||
|
||||
def _clear_anthropic_env(monkeypatch):
|
||||
for var in (
|
||||
"ANTHROPIC_API_KEY",
|
||||
"ANTHROPIC_AUTH_TOKEN",
|
||||
"ANTHROPIC_FEDERATION_RULE_ID",
|
||||
"ANTHROPIC_ORGANIZATION_ID",
|
||||
"ANTHROPIC_SERVICE_ACCOUNT_ID",
|
||||
"ANTHROPIC_WORKSPACE_ID",
|
||||
"ANTHROPIC_IDENTITY_TOKEN",
|
||||
"ANTHROPIC_IDENTITY_TOKEN_FILE",
|
||||
):
|
||||
monkeypatch.delenv(var, raising=False)
|
||||
|
||||
|
||||
class TestAnthropicWorkloadIdentityFederationCredentials:
|
||||
def test_missing_required_field_raises(self):
|
||||
with pytest.raises(AnthropicWorkloadIdentityFederationError) as exc:
|
||||
AnthropicWorkloadIdentityFederationCredentials(
|
||||
federation_rule_id="fdrl_x",
|
||||
organization_id="org_x",
|
||||
service_account_id="",
|
||||
identity_token_file="/tmp/jwt",
|
||||
)
|
||||
assert "service_account_id" in str(exc.value)
|
||||
|
||||
def test_missing_both_identity_sources_raises(self):
|
||||
with pytest.raises(AnthropicWorkloadIdentityFederationError) as exc:
|
||||
AnthropicWorkloadIdentityFederationCredentials(
|
||||
federation_rule_id="fdrl_x",
|
||||
organization_id="org_x",
|
||||
service_account_id="svac_x",
|
||||
)
|
||||
assert "identity_token" in str(exc.value)
|
||||
|
||||
def test_workspace_id_is_optional(self):
|
||||
creds = AnthropicWorkloadIdentityFederationCredentials(
|
||||
federation_rule_id="fdrl_x",
|
||||
organization_id="org_x",
|
||||
service_account_id="svac_x",
|
||||
identity_token="raw.jwt",
|
||||
)
|
||||
assert creds.workspace_id is None
|
||||
|
||||
|
||||
class TestGetWorkloadIdentityFederationCredentialsFromEnv:
|
||||
def test_returns_none_when_required_var_missing(self, monkeypatch):
|
||||
_clear_anthropic_env(monkeypatch)
|
||||
monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_x")
|
||||
monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org_x")
|
||||
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "raw.jwt")
|
||||
assert get_workload_identity_federation_credentials_from_env() is None
|
||||
|
||||
def test_returns_none_when_no_identity_source(self, monkeypatch):
|
||||
_clear_anthropic_env(monkeypatch)
|
||||
monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_x")
|
||||
monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org_x")
|
||||
monkeypatch.setenv("ANTHROPIC_SERVICE_ACCOUNT_ID", "svac_x")
|
||||
assert get_workload_identity_federation_credentials_from_env() is None
|
||||
|
||||
def test_returns_credentials_with_identity_token_env(self, monkeypatch):
|
||||
_clear_anthropic_env(monkeypatch)
|
||||
monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_x")
|
||||
monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org_x")
|
||||
monkeypatch.setenv("ANTHROPIC_SERVICE_ACCOUNT_ID", "svac_x")
|
||||
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline.jwt")
|
||||
creds = get_workload_identity_federation_credentials_from_env()
|
||||
assert creds is not None
|
||||
assert creds.identity_token == "inline.jwt"
|
||||
assert creds.identity_token_file is None
|
||||
assert creds.workspace_id is None
|
||||
|
||||
def test_returns_credentials_with_identity_token_file_env(self, monkeypatch, tmp_path):
|
||||
_clear_anthropic_env(monkeypatch)
|
||||
token_file = tmp_path / "jwt"
|
||||
token_file.write_text("file.jwt")
|
||||
monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_x")
|
||||
monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org_x")
|
||||
monkeypatch.setenv("ANTHROPIC_SERVICE_ACCOUNT_ID", "svac_x")
|
||||
monkeypatch.setenv("ANTHROPIC_WORKSPACE_ID", "wrkspc_x")
|
||||
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN_FILE", str(token_file))
|
||||
creds = get_workload_identity_federation_credentials_from_env()
|
||||
assert creds is not None
|
||||
assert creds.identity_token_file == str(token_file)
|
||||
assert creds.workspace_id == "wrkspc_x"
|
||||
|
||||
|
||||
class TestAnthropicWorkloadIdentityFederationTokenProviderExchange:
|
||||
def test_exchange_posts_jwt_bearer_form(self, tmp_path):
|
||||
token_file = tmp_path / "jwt"
|
||||
token_file.write_text("assertion.v1")
|
||||
stub = _StubClient()
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=_make_creds(identity_token_file=str(token_file)),
|
||||
api_base="https://api.anthropic.com",
|
||||
http_client=stub, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
token = provider.get_token()
|
||||
assert token == "anth_at_test"
|
||||
assert len(stub.calls) == 1
|
||||
call = stub.calls[0]
|
||||
assert call["url"] == "https://api.anthropic.com/v1/oauth/token"
|
||||
assert call["data"]["grant_type"] == ANTHROPIC_WORKLOAD_IDENTITY_FEDERATION_GRANT_TYPE
|
||||
assert call["data"]["assertion"] == "assertion.v1"
|
||||
assert call["data"]["federation_rule_id"] == "fdrl_test_rule"
|
||||
assert call["data"]["organization_id"] == "org_test"
|
||||
assert call["data"]["service_account_id"] == "svac_test"
|
||||
assert call["data"]["workspace_id"] == "wrkspc_test"
|
||||
assert call["headers"]["content-type"] == "application/x-www-form-urlencoded"
|
||||
|
||||
def test_workspace_id_omitted_when_unset(self):
|
||||
stub = _StubClient()
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=_make_creds(identity_token="raw.jwt", workspace_id=None),
|
||||
http_client=stub, # type: ignore[arg-type]
|
||||
)
|
||||
provider.get_token()
|
||||
assert "workspace_id" not in stub.calls[0]["data"]
|
||||
|
||||
def test_identity_token_used_directly(self):
|
||||
stub = _StubClient()
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=_make_creds(identity_token="inline.assertion"),
|
||||
http_client=stub, # type: ignore[arg-type]
|
||||
)
|
||||
provider.get_token()
|
||||
assert stub.calls[0]["data"]["assertion"] == "inline.assertion"
|
||||
|
||||
def test_assertion_re_read_on_each_exchange(self, tmp_path):
|
||||
token_file = tmp_path / "jwt"
|
||||
token_file.write_text("assertion.v1")
|
||||
stub = _StubClient(response=_build_response(expires_in=3600))
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=_make_creds(identity_token_file=str(token_file)),
|
||||
http_client=stub, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
provider.get_token()
|
||||
token_file.write_text("assertion.v2")
|
||||
provider._access_token = None
|
||||
provider._expires_at = 0.0
|
||||
provider.get_token()
|
||||
|
||||
assert [c["data"]["assertion"] for c in stub.calls] == [
|
||||
"assertion.v1",
|
||||
"assertion.v2",
|
||||
]
|
||||
|
||||
def test_assertion_kwarg_overrides_credentials(self, tmp_path):
|
||||
token_file = tmp_path / "jwt"
|
||||
token_file.write_text("file.assertion")
|
||||
stub = _StubClient(response=_build_response(expires_in=3600))
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=_make_creds(identity_token_file=str(token_file)),
|
||||
http_client=stub, # type: ignore[arg-type]
|
||||
)
|
||||
provider.get_token(assertion="caller.assertion")
|
||||
assert stub.calls[0]["data"]["assertion"] == "caller.assertion"
|
||||
|
||||
def test_assertion_kwarg_forces_refresh(self, tmp_path):
|
||||
token_file = tmp_path / "jwt"
|
||||
token_file.write_text("file.assertion")
|
||||
stub = _StubClient(response=_build_response(expires_in=3600))
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=_make_creds(identity_token_file=str(token_file)),
|
||||
http_client=stub, # type: ignore[arg-type]
|
||||
)
|
||||
provider.get_token()
|
||||
provider.get_token(assertion="explicit.assertion")
|
||||
assert len(stub.calls) == 2
|
||||
assert stub.calls[1]["data"]["assertion"] == "explicit.assertion"
|
||||
|
||||
def test_missing_access_token_in_response_raises(self, tmp_path):
|
||||
token_file = tmp_path / "jwt"
|
||||
token_file.write_text("assertion")
|
||||
bad_response = httpx.Response(
|
||||
status_code=200,
|
||||
json={"expires_in": 3600},
|
||||
request=httpx.Request("POST", "https://x"),
|
||||
)
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=_make_creds(identity_token_file=str(token_file)),
|
||||
http_client=_StubClient(response=bad_response), # type: ignore[arg-type]
|
||||
)
|
||||
with pytest.raises(AnthropicWorkloadIdentityFederationError) as exc:
|
||||
provider.get_token()
|
||||
assert "access_token" in str(exc.value)
|
||||
|
||||
def test_http_error_response_raises(self, tmp_path):
|
||||
token_file = tmp_path / "jwt"
|
||||
token_file.write_text("assertion")
|
||||
err_response = httpx.Response(
|
||||
status_code=401,
|
||||
text="unauthorized",
|
||||
request=httpx.Request("POST", "https://x"),
|
||||
)
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=_make_creds(identity_token_file=str(token_file)),
|
||||
http_client=_StubClient(response=err_response), # type: ignore[arg-type]
|
||||
)
|
||||
with pytest.raises(AnthropicWorkloadIdentityFederationError) as exc:
|
||||
provider.get_token()
|
||||
assert "HTTP 401" in str(exc.value)
|
||||
|
||||
def test_missing_jwt_file_raises(self, tmp_path):
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=_make_creds(identity_token_file=str(tmp_path / "does-not-exist")),
|
||||
http_client=_StubClient(), # type: ignore[arg-type]
|
||||
)
|
||||
with pytest.raises(AnthropicWorkloadIdentityFederationError) as exc:
|
||||
provider.get_token()
|
||||
assert "ANTHROPIC_IDENTITY_TOKEN_FILE" in str(exc.value)
|
||||
|
||||
|
||||
class TestAnthropicWorkloadIdentityFederationTokenProviderRefresh:
|
||||
def test_serves_cached_token_when_fresh(self, tmp_path):
|
||||
token_file = tmp_path / "jwt"
|
||||
token_file.write_text("assertion")
|
||||
stub = _StubClient(response=_build_response(expires_in=3600))
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=_make_creds(identity_token_file=str(token_file)),
|
||||
http_client=stub, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
first = provider.get_token()
|
||||
second = provider.get_token()
|
||||
assert first == second
|
||||
assert len(stub.calls) == 1
|
||||
|
||||
def test_advisory_refresh_failure_returns_cached_token(self, tmp_path):
|
||||
token_file = tmp_path / "jwt"
|
||||
token_file.write_text("assertion")
|
||||
stub = _StubClient(response=_build_response(expires_in=3600))
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=_make_creds(identity_token_file=str(token_file)),
|
||||
http_client=stub, # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
provider.get_token()
|
||||
original_token = provider._access_token
|
||||
import time as _time
|
||||
|
||||
provider._expires_at = _time.time() + 90
|
||||
stub._exc = AnthropicWorkloadIdentityFederationError("network blip")
|
||||
|
||||
served = provider.get_token()
|
||||
assert served == original_token
|
||||
|
||||
def test_mandatory_refresh_failure_raises(self, tmp_path):
|
||||
token_file = tmp_path / "jwt"
|
||||
token_file.write_text("assertion")
|
||||
stub = _StubClient(response=_build_response(expires_in=3600))
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=_make_creds(identity_token_file=str(token_file)),
|
||||
http_client=stub, # type: ignore[arg-type]
|
||||
)
|
||||
provider.get_token()
|
||||
import time as _time
|
||||
|
||||
provider._expires_at = _time.time() + 10
|
||||
stub._exc = AnthropicWorkloadIdentityFederationError("hard fail")
|
||||
|
||||
with pytest.raises(AnthropicWorkloadIdentityFederationError):
|
||||
provider.get_token()
|
||||
|
||||
|
||||
class TestExchangeAnthropicWorkloadIdentityFederationToken:
|
||||
def test_returns_none_when_env_vars_missing(self, monkeypatch):
|
||||
_clear_anthropic_env(monkeypatch)
|
||||
assert exchange_anthropic_workload_identity_federation_token() is None
|
||||
|
||||
def test_explicit_credentials_and_assertion_passthrough(self, monkeypatch):
|
||||
_clear_anthropic_env(monkeypatch)
|
||||
creds = AnthropicWorkloadIdentityFederationCredentials(
|
||||
federation_rule_id="fdrl_x",
|
||||
organization_id="org_x",
|
||||
service_account_id="svac_x",
|
||||
identity_token="placeholder",
|
||||
)
|
||||
stub = _StubClient()
|
||||
from litellm.llms.anthropic import workload_identity_federation as wif_module
|
||||
|
||||
provider = AnthropicWorkloadIdentityFederationTokenProvider(
|
||||
credentials=creds,
|
||||
http_client=stub, # type: ignore[arg-type]
|
||||
)
|
||||
with patch.object(
|
||||
wif_module,
|
||||
"get_or_create_workload_identity_federation_provider",
|
||||
return_value=provider,
|
||||
):
|
||||
token = exchange_anthropic_workload_identity_federation_token(
|
||||
credentials=creds, assertion="caller.assertion"
|
||||
)
|
||||
assert token == "anth_at_test"
|
||||
assert stub.calls[0]["data"]["assertion"] == "caller.assertion"
|
||||
|
||||
|
||||
class TestValidateEnvironmentWiring:
|
||||
def test_validate_environment_uses_workload_identity_federation_when_no_api_key(self, monkeypatch, tmp_path):
|
||||
_clear_anthropic_env(monkeypatch)
|
||||
monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_x")
|
||||
monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org_x")
|
||||
monkeypatch.setenv("ANTHROPIC_SERVICE_ACCOUNT_ID", "svac_x")
|
||||
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline.jwt")
|
||||
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
with patch(
|
||||
"litellm.llms.anthropic.common_utils.exchange_anthropic_workload_identity_federation_token",
|
||||
return_value="anth_at_exchanged",
|
||||
):
|
||||
config = AnthropicModelInfo()
|
||||
headers = config.validate_environment(
|
||||
headers={},
|
||||
model="claude-3-5-sonnet-20241022",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert headers["authorization"] == "Bearer anth_at_exchanged"
|
||||
assert "x-api-key" not in headers
|
||||
|
||||
def test_validate_environment_raises_when_no_creds(self, monkeypatch):
|
||||
_clear_anthropic_env(monkeypatch)
|
||||
|
||||
import litellm
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
config = AnthropicModelInfo()
|
||||
with pytest.raises(litellm.AuthenticationError) as exc:
|
||||
config.validate_environment(
|
||||
headers={},
|
||||
model="claude-3-5-sonnet-20241022",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=None,
|
||||
)
|
||||
assert "ANTHROPIC_API_KEY" in str(exc.value)
|
||||
Loading…
Add table
Reference in a new issue