diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 31131d722ab..8909f68a0a0 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -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 diff --git a/litellm/llms/anthropic/workload_identity_federation.py b/litellm/llms/anthropic/workload_identity_federation.py new file mode 100644 index 00000000000..251576c565f --- /dev/null +++ b/litellm/llms/anthropic/workload_identity_federation.py @@ -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() diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_workload_identity_federation.py b/tests/test_litellm/llms/anthropic/test_anthropic_workload_identity_federation.py new file mode 100644 index 00000000000..234fb635653 --- /dev/null +++ b/tests/test_litellm/llms/anthropic/test_anthropic_workload_identity_federation.py @@ -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)