diff --git a/litellm/llms/anthropic/wif.py b/litellm/llms/anthropic/wif.py index 2444a8f833e..b47c4f28eb2 100644 --- a/litellm/llms/anthropic/wif.py +++ b/litellm/llms/anthropic/wif.py @@ -65,6 +65,7 @@ _IDENTITY_SOURCE_PARAM: Final = "anthropic_identity_source" _IDENTITY_SOURCE_ENV: Final = "ANTHROPIC_IDENTITY_SOURCE" _IDENTITY_TOKEN_FILE_PARAM: Final = "anthropic_identity_token_file" _IDENTITY_TOKEN_PARAM: Final = "anthropic_identity_token" +_LEGACY_REF_PARAMS: Final = (_IDENTITY_TOKEN_FILE_PARAM, _IDENTITY_TOKEN_PARAM) # litellm_params key -> InternalIssuerSource/KeycloakSource field name. Every key here must # also be listed in ANTHROPIC_WIF_KWARGS_KEYS (types/workload_identity.py), which is what makes it @@ -139,7 +140,7 @@ def resolve_anthropic_wif_params(litellm_params: Mapping[str, object] | None) -> ) organization_id: Final = _config_value(litellm_params, "anthropic_organization_id", "ANTHROPIC_ORGANIZATION_ID") if federation_rule_id is None or organization_id is None: - _raise_if_identity_source_configured(litellm_params, federation_rule_id, organization_id) + _raise_if_federation_requested(litellm_params, federation_rule_id, organization_id) return None identity_source: Final = _resolve_identity_source(litellm_params) if identity_source is None: @@ -204,17 +205,15 @@ def _raise_unknown_source_kind(source_kind: str) -> NoReturn: ) -def _raise_if_identity_source_configured( +def _raise_if_federation_requested( litellm_params: Mapping[str, object] | None, federation_rule_id: str | None, organization_id: str | None ) -> None: - """A configured identity source is an explicit request to federate, so a missing rule or - organization id fails closed with the ids named, rather than silently skipping federation - and surfacing later as a missing API key.""" - source_kind: Final = _resolve_source_kind(litellm_params) - if source_kind is None: + """A configured identity source, or a token file or inline token set on the deployment, is an + explicit request to federate, so a missing rule or organization id fails closed with the ids + named, rather than silently skipping federation and surfacing later as a missing API key.""" + request: Final = _explicit_federation_request(litellm_params) + if request is None: return - if source_kind not in {kind.value for kind in AnthropicIdentitySourceKind}: - _raise_unknown_source_kind(source_kind) missing: Final = tuple( param for param, value in ( @@ -225,7 +224,7 @@ def _raise_if_identity_source_configured( ) raise litellm.AuthenticationError( message=( - f"{_IDENTITY_SOURCE_PARAM} is {source_kind!r}, but {' and '.join(missing)} " + f"{request}, but {' and '.join(missing)} " f"{'is' if len(missing) == 1 else 'are'} not set. {_MISSING_IDS_HINT}" ), llm_provider="anthropic", @@ -233,14 +232,25 @@ def _raise_if_identity_source_configured( ) +def _explicit_federation_request(litellm_params: Mapping[str, object] | None) -> str | None: + source_kind: Final = _resolve_source_kind(litellm_params) + if source_kind is None: + legacy_ref_param: Final = _legacy_ref_param(litellm_params) + return None if legacy_ref_param is None else f"{legacy_ref_param} is set" + if source_kind not in {kind.value for kind in AnthropicIdentitySourceKind}: + _raise_unknown_source_kind(source_kind) + return f"{_IDENTITY_SOURCE_PARAM} is {source_kind!r}" + + def _resolve_source_kind(litellm_params: Mapping[str, object] | None) -> str | None: param_kind: Final = _param_str(litellm_params, _IDENTITY_SOURCE_PARAM) if param_kind is not None: return param_kind - has_param_legacy_ref: Final = any( - _param_str(litellm_params, key) is not None for key in (_IDENTITY_TOKEN_FILE_PARAM, _IDENTITY_TOKEN_PARAM) - ) - return None if has_param_legacy_ref else _env_str(_IDENTITY_SOURCE_ENV) + return None if _legacy_ref_param(litellm_params) is not None else _env_str(_IDENTITY_SOURCE_ENV) + + +def _legacy_ref_param(litellm_params: Mapping[str, object] | None) -> str | None: + return next((key for key in _LEGACY_REF_PARAMS if _param_str(litellm_params, key) is not None), None) def _reject_foreign_variant_fields( diff --git a/tests/integration/management/test_credential_federation_values.py b/tests/integration/management/test_credential_federation_values.py index 3f812c87c59..e17fb1154d8 100644 --- a/tests/integration/management/test_credential_federation_values.py +++ b/tests/integration/management/test_credential_federation_values.py @@ -1,11 +1,13 @@ from __future__ import annotations +import asyncio import json import uuid from collections.abc import Callable, Iterator, Mapping, Sequence from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from functools import partial +from itertools import product from pathlib import Path from types import MappingProxyType from typing import Final, TypeVar @@ -43,7 +45,7 @@ from tests.integration._support.client import ( string_value, ) from tests.integration._support.database import read_rows -from tests.integration._support.process import OwnedProxy, owned_proxy_process +from tests.integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process from tests.integration._support.wire import Reply, Request, Wire, wire_server T = TypeVar("T") @@ -562,7 +564,13 @@ class _Secrets: ) -_REMOVED_ENVIRONMENT: Final = ("ANTHROPIC_IDENTITY_TOKEN_FILE", "ANTHROPIC_API_KEY", "ANTHROPIC_API_BASE") +_REMOVED_ENVIRONMENT: Final = ( + "ANTHROPIC_IDENTITY_TOKEN_FILE", + "ANTHROPIC_API_KEY", + "ANTHROPIC_API_BASE", + "ANTHROPIC_FEDERATION_RULE_ID", + "ANTHROPIC_ORGANIZATION_ID", +) def _secrets(directory: Path) -> _Secrets: @@ -817,3 +825,388 @@ def test_worker_kill_mid_credential_burst(gateway: Gateway, tmp_path: Path) -> N for name in landed: _converged(survivor, name, _shape("token_file", f"fdrl-{name}")) _stable(partial(_chat_outcome, survivor, model), lambda outcome: outcome == (200, True)) + + +_FAIL_CLOSED_HINT: Final = "Settings > Workload identity" +_REFUSAL_CLIENTS: Final = ( + "chat", + "chat_stream", + "chat_async", + "messages", + "messages_stream", + "responses", + "responses_stream", +) +_OWNED_PROXY_CELL_SECONDS: Final = 2 * graceful_stop_seconds() + 120 + + +@dataclass(frozen=True, slots=True) +class _Misconfiguration: + reference: str + missing: tuple[str, ...] + blank: bool + expected: str + + def values(self, rig: FederationRig, tag: str) -> dict[str, JsonValue]: + token: Final = ( + str(rig.secrets.token_file) + if self.reference == "anthropic_identity_token_file" + else f"oidc/env/{_IDENTITY_TOKEN_VARIABLE}" + ) + ids: Final = {"anthropic_federation_rule_id": f"fdrl-{tag}", "anthropic_organization_id": f"org-{tag}"} + kept: Final = {key: value for key, value in ids.items() if key not in self.missing} + blanked: Final = {key: "" for key in self.missing} if self.blank else {} + return {self.reference: token, **kept, **blanked} + + +_MISCONFIGURED: Final[Mapping[str, _Misconfiguration]] = MappingProxyType( + { + "token_file_without_org": _Misconfiguration( + "anthropic_identity_token_file", + ("anthropic_organization_id",), + False, + "anthropic_identity_token_file is set, but anthropic_organization_id is not set", + ), + "token_file_without_rule": _Misconfiguration( + "anthropic_identity_token_file", + ("anthropic_federation_rule_id",), + False, + "anthropic_identity_token_file is set, but anthropic_federation_rule_id is not set", + ), + "inline_token_without_org": _Misconfiguration( + "anthropic_identity_token", + ("anthropic_organization_id",), + False, + "anthropic_identity_token is set, but anthropic_organization_id is not set", + ), + "inline_token_without_rule": _Misconfiguration( + "anthropic_identity_token", + ("anthropic_federation_rule_id",), + False, + "anthropic_identity_token is set, but anthropic_federation_rule_id is not set", + ), + "token_file_without_ids": _Misconfiguration( + "anthropic_identity_token_file", + ("anthropic_federation_rule_id", "anthropic_organization_id"), + False, + "anthropic_identity_token_file is set, but anthropic_federation_rule_id and anthropic_organization_id" + " are not set", + ), + "token_file_blank_org": _Misconfiguration( + "anthropic_identity_token_file", + ("anthropic_organization_id",), + True, + "anthropic_identity_token_file is set, but anthropic_organization_id is not set", + ), + } +) + + +@dataclass(frozen=True, slots=True) +class _Misconfigured: + tag: str + credential: str + model: str + shape: _Misconfiguration + + +def _misconfigured(rig: FederationRig, scenario: Scenario, shape: _Misconfiguration) -> _Misconfigured: + tag: Final = uuid.uuid4().hex + credential: Final = _create(rig.owned.gateway, scenario, shape.values(rig, tag)) + model: Final = _federated_deployment(rig.owned.gateway, scenario, credential, rig.peer.wire.url) + return _Misconfigured(tag=tag, credential=credential, model=model, shape=shape) + + +async def _async_chat(base_url: str, key: str, model: str, marker: str) -> None: + await openai.AsyncOpenAI(base_url=base_url + "/v1", api_key=key, max_retries=0).chat.completions.create( + model=model, messages=[{"role": "user", "content": prompt(marker)}] + ) + + +def _attempt(rig: FederationRig, client: str, model: str, marker: str) -> None: + base_url: Final = str(rig.owned.gateway.client.base_url) + key: Final = rig.owned.gateway.key + chat: Final = openai.OpenAI(base_url=base_url + "/v1", api_key=key, max_retries=0) + messages: Final = anthropic.Anthropic(base_url=base_url, api_key=key, max_retries=0) + match client: + case "chat": + chat.chat.completions.create(model=model, messages=[{"role": "user", "content": prompt(marker)}]) + case "chat_stream": + for _ in chat.chat.completions.create( + model=model, messages=[{"role": "user", "content": prompt(marker)}], stream=True + ): + pass + case "chat_async": + asyncio.run(_async_chat(base_url, key, model, marker)) + case "messages": + messages.messages.create(model=model, max_tokens=64, messages=[{"role": "user", "content": prompt(marker)}]) + case "messages_stream": + with messages.messages.stream( + model=model, max_tokens=64, messages=[{"role": "user", "content": prompt(marker)}] + ) as stream: + stream.get_final_message() + case "responses": + chat.responses.create(model=model, input=prompt(marker)) + case "responses_stream": + for _ in chat.responses.create(model=model, input=prompt(marker), stream=True): + pass + case _: + pytest.fail(f"unknown client {client!r}") + + +def _refusal(rig: FederationRig, client: str, model: str, marker: str) -> tuple[int, str]: + try: + _attempt(rig, client, model, marker) + except (openai.APIStatusError, anthropic.APIStatusError) as error: + return error.status_code, str(error) + pytest.fail(f"{client} succeeded against the misconfigured deployment {model}") + + +def _assert_names_missing_id(message: str, shape: _Misconfiguration) -> None: + assert shape.expected in message, message + assert _FAIL_CLOSED_HINT in message, message + + +def _assert_fails_closed(outcome: tuple[int, str], shape: _Misconfiguration) -> None: + status, message = outcome + assert status == 401, message + _assert_names_missing_id(message, shape) + + +def _assert_peer_untouched(rig: FederationRig, *fragments: str) -> None: + bodies: Final = tuple(request.body.decode() for request in rig.peer.requests()) + touched: Final = tuple(body for body in bodies if any(fragment in body for fragment in fragments)) + assert touched == (), touched + + +@pytest.mark.timeout(240) +@pytest.mark.parametrize("client", _REFUSAL_CLIENTS) +@pytest.mark.parametrize("shape", tuple(_MISCONFIGURED)) +def test_legacy_reference_without_ids_fails_closed_before_any_exchange( + federation: FederationRig, shape: str, client: str +) -> None: + marker: Final = uuid.uuid4().hex + with federation.owned.gateway.scenario() as scenario: + misconfigured: Final = _misconfigured(federation, scenario, _MISCONFIGURED[shape]) + _deployments_visible(federation.owned.gateway, (misconfigured.model,)) + _assert_fails_closed(_refusal(federation, client, misconfigured.model, marker), misconfigured.shape) + _assert_peer_untouched(federation, misconfigured.tag, marker) + + +@pytest.mark.timeout(240) +def test_file_upload_fails_closed_naming_the_missing_id(federation: FederationRig) -> None: + note: Final = uuid.uuid4().hex + with federation.owned.gateway.scenario() as scenario: + misconfigured: Final = _misconfigured(federation, scenario, _MISCONFIGURED["token_file_without_org"]) + _deployments_visible(federation.owned.gateway, (misconfigured.model,)) + response: Final = federation.owned.gateway.request_multipart( + "/v1/files", + {"purpose": "user_data", "model": misconfigured.model}, + {"file": ("notes.jsonl", json.dumps({"note": note}).encode(), "application/jsonl")}, + ) + _assert_fails_closed((response.status_code, response.text), misconfigured.shape) + _assert_peer_untouched(federation, misconfigured.tag, note) + + +@pytest.mark.timeout(240) +def test_skills_listing_fails_closed_naming_the_missing_id(federation: FederationRig) -> None: + with federation.owned.gateway.scenario() as scenario: + misconfigured: Final = _misconfigured(federation, scenario, _MISCONFIGURED["inline_token_without_rule"]) + _deployments_visible(federation.owned.gateway, (misconfigured.model,)) + response: Final = federation.owned.gateway.request( + "GET", "/v1/skills", params={"beta": "true", "model": misconfigured.model}, headers=_CLOSE + ) + _assert_fails_closed((response.status_code, response.text), misconfigured.shape) + upstream: Final = federation.peer.requests() + listed: Final = tuple(request for request in upstream if request.target.startswith("/v1/skills")) + assert listed == (), listed + _assert_peer_untouched(federation, misconfigured.tag) + + +@pytest.mark.timeout(240) +def test_health_report_names_the_missing_id(federation: FederationRig) -> None: + with federation.owned.gateway.scenario() as scenario: + misconfigured: Final = _misconfigured(federation, scenario, _MISCONFIGURED["token_file_without_rule"]) + _deployments_visible(federation.owned.gateway, (misconfigured.model,)) + response: Final = federation.owned.gateway.request( + "GET", "/health", params={"model": misconfigured.model}, headers=_CLOSE + ) + assert response.status_code == 503, response.text + unhealthy: Final = JSON_OBJECT.validate_json(response.content)["unhealthy_endpoints"] + assert isinstance(unhealthy, list) and len(unhealthy) == 1, response.text + _assert_names_missing_id(string_value(object_value(unhealthy[0])["error"]), misconfigured.shape) + _assert_peer_untouched(federation, misconfigured.tag) + + +@pytest.mark.timeout(240) +def test_connection_probe_names_the_missing_id(federation: FederationRig) -> None: + with federation.owned.gateway.scenario() as scenario: + misconfigured: Final = _misconfigured(federation, scenario, _MISCONFIGURED["inline_token_without_org"]) + _deployments_visible(federation.owned.gateway, (misconfigured.model,)) + response: Final = federation.owned.gateway.request( + "POST", + "/health/test_connection", + { + "litellm_params": { + "model": _MODEL, + "api_base": federation.peer.wire.url, + "litellm_credential_name": misconfigured.credential, + }, + "mode": "chat", + }, + headers=_CLOSE, + ) + assert response.status_code == 200, response.text + assert JSON_OBJECT.validate_json(response.content)["status"] == "error", response.text + _assert_names_missing_id(response.text, misconfigured.shape) + _assert_peer_untouched(federation, misconfigured.tag) + + +@pytest.mark.timeout(240) +def test_static_key_beside_a_stray_token_reference_still_wins(federation: FederationRig) -> None: + tag: Final = uuid.uuid4().hex + marker: Final = uuid.uuid4().hex + api_key: Final = f"sk-ant-api03-{tag}" + owned: Final = federation.owned + with owned.gateway.scenario() as scenario: + credential: Final = _create( + owned.gateway, + scenario, + { + "api_key": api_key, + "anthropic_identity_token_file": str(federation.secrets.token_file), + "anthropic_federation_rule_id": f"fdrl-{tag}", + }, + ) + model: Final = _federated_deployment(owned.gateway, scenario, credential, federation.peer.wire.url) + _deployments_visible(owned.gateway, (model,)) + _call(federation, "chat", model, marker) + upstream: Final = federation.peer.requests() + sent: Final = tuple( + request for request in upstream if request.target == "/v1/messages" and marker in request.body.decode() + ) + assert len(sent) == 1, upstream + assert sent[0].headers.get("x-api-key") == api_key, sent[0].headers + assert "authorization" not in sent[0].headers, sent[0].headers + _assert_peer_untouched(federation, tag) + + +@pytest.mark.timeout(240) +def test_blank_token_file_reference_is_unset_and_the_ambient_token_federates(federation: FederationRig) -> None: + rule_id: Final = f"fdrl-blank-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + owned: Final = federation.owned + with owned.gateway.scenario() as scenario: + credential: Final = _create(owned.gateway, scenario, {**_ids(rule_id), "anthropic_identity_token_file": ""}) + model: Final = _federated_deployment(owned.gateway, scenario, credential, federation.peer.wire.url) + _deployments_visible(owned.gateway, (model,)) + _call(federation, "chat", model, marker) + grants: Final = tuple( + JSON_OBJECT.validate_json(request.body) + for request in federation.peer.requests() + if request.target == "/v1/oauth/token" + ) + mine: Final = tuple(grant for grant in grants if grant["federation_rule_id"] == rule_id) + assert mine, grants + assert all(grant["assertion"] == federation.secrets.environment_token for grant in mine), mine + + +def test_unauthenticated_call_is_refused_before_the_credential_is_read(federation: FederationRig) -> None: + owned: Final = federation.owned + marker: Final = uuid.uuid4().hex + with owned.gateway.scenario() as scenario: + entry: Final = _misconfigured(federation, scenario, _MISCONFIGURED["token_file_without_org"]) + _deployments_visible(owned.gateway, (entry.model,)) + refused: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": entry.model, "messages": [{"role": "user", "content": prompt(marker)}]}, + key=f"sk-not-a-key-{marker}", + headers=_CLOSE, + ) + assert refused.status_code == 401, refused.text + assert _FAIL_CLOSED_HINT not in refused.text, refused.text + _assert_fails_closed(_refusal(federation, "chat", entry.model, uuid.uuid4().hex), entry.shape) + _assert_peer_untouched(federation, entry.tag, marker) + + +@pytest.mark.timeout(_OWNED_PROXY_CELL_SECONDS) +def test_environment_organization_id_completes_a_legacy_reference(gateway: Gateway, tmp_path: Path) -> None: + directory: Final = tmp_path.resolve() + secrets: Final = _secrets(directory) + tag: Final = uuid.uuid4().hex + organization: Final = f"org-env-{tag}" + rule_id: Final = f"fdrl-{tag}" + marker: Final = uuid.uuid4().hex + with wire_server(_federation_peer) as wire: + with owned_proxy_process( + gateway, + directory, + {**secrets.overrides(), "ANTHROPIC_ORGANIZATION_ID": organization}, + remove_environment=_REMOVED_ENVIRONMENT, + workers=2, + ) as owned: + with owned.gateway.scenario() as scenario: + credential: Final = _create( + owned.gateway, + scenario, + {"anthropic_identity_token_file": str(secrets.token_file), "anthropic_federation_rule_id": rule_id}, + ) + model: Final = _federated_deployment(owned.gateway, scenario, credential, wire.url) + _deployments_visible(owned.gateway, (model,)) + response: Final = _chat(owned.gateway, model, marker) + assert response.status_code == 200, response.text + upstream: Final = wire.drain() + grants: Final = tuple( + JSON_OBJECT.validate_json(request.body) + for request in upstream + if request.target == "/v1/oauth/token" + ) + mine: Final = tuple(grant for grant in grants if grant["federation_rule_id"] == rule_id) + assert mine, upstream + assert all(grant["organization_id"] == organization for grant in mine), mine + assert all(grant["assertion"] == secrets.token_file.read_text() for grant in mine), mine + + +@pytest.mark.timeout(240) +def test_mixed_burst_fails_closed_without_disturbing_healthy_federation(federation: FederationRig) -> None: + owned: Final = federation.owned + shapes: Final = tuple(_MISCONFIGURED.values()) + healthy: Final = tuple(product(SOURCES, CLIENTS)) + with owned.gateway.scenario() as scenario: + misconfigured: Final = tuple( + _misconfigured(federation, scenario, shapes[index % len(shapes)]) for index in range(12) + ) + _deployments_visible(owned.gateway, tuple(entry.model for entry in misconfigured)) + control: Final = scenario.model() + serial_markers: Final = tuple(uuid.uuid4().hex for _ in _REFUSAL_CLIENTS) + for client, marker in zip(_REFUSAL_CLIENTS, serial_markers, strict=True): + _assert_fails_closed(_refusal(federation, client, misconfigured[0].model, marker), misconfigured[0].shape) + refused_markers: Final = tuple(uuid.uuid4().hex for _ in misconfigured) + healthy_markers: Final = tuple(uuid.uuid4().hex for _ in healthy) + + def refuse(index: int) -> tuple[int, str]: + client: Final = _REFUSAL_CLIENTS[index % len(_REFUSAL_CLIENTS)] + return _refusal(federation, client, misconfigured[index].model, refused_markers[index]) + + def federate(index: int) -> None: + source, client = healthy[index] + _call(federation, client, federation.deployments[source], healthy_markers[index]) + + with ThreadPoolExecutor(max_workers=16) as pool: + refusals: Final = tuple(pool.submit(refuse, index) for index in range(len(misconfigured))) + federations: Final = tuple(pool.submit(federate, index) for index in range(len(healthy))) + controls: Final = tuple(pool.submit(_chat_outcome, owned.gateway, control) for _ in range(4)) + outcomes: Final = tuple(future.result() for future in refusals) + for future in federations: + future.result() + assert tuple(future.result()[0] for future in controls) == (200,) * 4 + for entry, outcome in zip(misconfigured, outcomes, strict=True): + _assert_fails_closed(outcome, entry.shape) + sent: Final = tuple( + request.body.decode() for request in federation.peer.requests() if request.target == "/v1/messages" + ) + for marker in healthy_markers: + assert sum(marker in body for body in sent) == 1, marker + _assert_peer_untouched(federation, *serial_markers, *refused_markers, *(entry.tag for entry in misconfigured)) + assert owned.gateway.request("GET", "/health/liveliness").status_code == 200 diff --git a/tests/unit/llms/anthropic/test_anthropic_common_utils.py b/tests/unit/llms/anthropic/test_anthropic_common_utils.py index 089538882eb..17840a829c1 100644 --- a/tests/unit/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/unit/llms/anthropic/test_anthropic_common_utils.py @@ -14,6 +14,7 @@ import json import os import sys import threading +from pathlib import Path from types import SimpleNamespace from typing import Final from unittest.mock import patch @@ -2518,7 +2519,6 @@ class TestWifTierPrecedence: assert [record for record in caplog.records if "takes precedence" in record.getMessage()] == [] - class TestWifZeroBehaviorChange: def test_unconfigured_raises_same_authentication_error(self, clean_anthropic_env): """No WIF config and no keys: same AuthenticationError as today (message @@ -2536,6 +2536,29 @@ class TestWifZeroBehaviorChange: assert "ANTHROPIC_SERVICE_ACCOUNT_ID" in exc_info.value.message assert "ANTHROPIC_IDENTITY_TOKEN_FILE" in exc_info.value.message + def test_a_token_file_credential_missing_an_id_names_the_id_not_the_key(self, clean_anthropic_env: None, tmp_path: Path): + import litellm + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + with pytest.raises(litellm.AuthenticationError) as exc_info: + AnthropicModelInfo().validate_environment( + headers={}, + model="claude-haiku-5-5", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={ + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_identity_token_file": str(tmp_path / "token"), + }, + api_key=None, + api_base=None, + ) + + assert ( + "anthropic_identity_token_file is set, but anthropic_organization_id is not set" in exc_info.value.message + ) + assert "Missing Anthropic API Key" not in exc_info.value.message + class TestWifHeaderContract: def test_minted_token_headers(self, monkeypatch, wif_engine): diff --git a/tests/unit/llms/anthropic/test_anthropic_wif.py b/tests/unit/llms/anthropic/test_anthropic_wif.py index a054b4d130c..733e1472a8a 100644 --- a/tests/unit/llms/anthropic/test_anthropic_wif.py +++ b/tests/unit/llms/anthropic/test_anthropic_wif.py @@ -277,12 +277,18 @@ class TestExchangeHostTrust: def test_a_gateway_listed_with_its_port_is_trusted(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "gateway.internal:8443") - assert self._mint("https://gateway.internal:8443", monkeypatch) == "https://gateway.internal:8443/v1/oauth/token" + assert ( + self._mint("https://gateway.internal:8443", monkeypatch) == "https://gateway.internal:8443/v1/oauth/token" + ) def test_allowlist_matching_ignores_hostname_case(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "Gateway.Internal:8443") - assert self._mint("https://gateway.internal:8443", monkeypatch) == "https://gateway.internal:8443/v1/oauth/token" - assert self._mint("https://GATEWAY.internal:8443", monkeypatch) == "https://GATEWAY.internal:8443/v1/oauth/token" + assert ( + self._mint("https://gateway.internal:8443", monkeypatch) == "https://gateway.internal:8443/v1/oauth/token" + ) + assert ( + self._mint("https://GATEWAY.internal:8443", monkeypatch) == "https://GATEWAY.internal:8443/v1/oauth/token" + ) def test_a_gateway_listed_with_a_port_is_not_trusted_on_another_port(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "gateway.internal:8443") @@ -302,7 +308,9 @@ class TestExchangeHostTrust: def test_a_gateway_listed_without_a_port_is_trusted_on_every_port(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "gateway.internal") - assert self._mint("https://gateway.internal:9443", monkeypatch) == "https://gateway.internal:9443/v1/oauth/token" + assert ( + self._mint("https://gateway.internal:9443", monkeypatch) == "https://gateway.internal:9443/v1/oauth/token" + ) def test_an_entry_spelling_the_scheme_default_port_matches_a_base_that_omits_it( self, monkeypatch: pytest.MonkeyPatch @@ -311,7 +319,6 @@ class TestExchangeHostTrust: assert self._mint("https://gateway.internal", monkeypatch) == "https://gateway.internal/v1/oauth/token" - class TestBaseUrlDerivation: def _mint(self, api_base: str | None, monkeypatch: pytest.MonkeyPatch) -> str: monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt") @@ -548,8 +555,6 @@ class TestResolutionMatrix: {"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): @@ -834,7 +839,11 @@ class TestDenialHints: def test_only_console_pointer_when_both_set(self, monkeypatch: pytest.MonkeyPatch): message = self._raise( - {**self.BASE_PARAMS, "anthropic_federation_workspace_id": "wrkspc_1", "anthropic_service_account_id": "svac_1"}, + { + **self.BASE_PARAMS, + "anthropic_federation_workspace_id": "wrkspc_1", + "anthropic_service_account_id": "svac_1", + }, 401, monkeypatch, ) @@ -1064,7 +1073,6 @@ class TestKeycloakIdentitySourceDispatch: assert first.assertion_ref != second.assertion_ref - @pytest.mark.parametrize( "sparse_params", [TestInternalIssuerIdentitySourceDispatch.LITELLM_PARAMS, TestKeycloakIdentitySourceDispatch.LITELLM_PARAMS], @@ -1229,9 +1237,82 @@ class TestMissingIdsFailClosedWhenIdentitySourceConfigured: with pytest.raises(litellm.AuthenticationError, match="must be one of internal_issuer, keycloak"): resolve_anthropic_wif_params({"anthropic_identity_source": "bogus"}) - def test_legacy_token_params_without_ids_still_return_none(self, monkeypatch: pytest.MonkeyPatch): + +class TestLegacyRefsFailClosedWithoutIds: + """A token file or inline token on the deployment asks to federate as explicitly as a named + identity source does, so a missing rule or organization id is reported by name instead of + letting the request die later as a missing API key.""" + + def test_token_file_without_organization_id_names_the_file_param(self, tmp_path: Path): + with pytest.raises(litellm.AuthenticationError) as exc_info: + resolve_anthropic_wif_params( + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_identity_token_file": str(tmp_path / "token")} + ) + + message: Final = exc_info.value.message + assert "anthropic_identity_token_file is set, but anthropic_organization_id is not set. Copy" in message + assert "Settings > Workload identity" in message + assert "ANTHROPIC_FEDERATION_RULE_ID" in message + assert not message.endswith(".") + + def test_inline_token_without_rule_id_names_the_token_param(self): + with pytest.raises(litellm.AuthenticationError) as exc_info: + resolve_anthropic_wif_params( + {"anthropic_organization_id": "org-1", "anthropic_identity_token": "oidc/env/TOK"} + ) + + assert ( + "anthropic_identity_token is set, but anthropic_federation_rule_id is not set. Copy" + in exc_info.value.message + ) + + def test_token_file_with_both_ids_missing_names_both(self, tmp_path: Path): + with pytest.raises( + litellm.AuthenticationError, match="anthropic_federation_rule_id and anthropic_organization_id are not set" + ): + resolve_anthropic_wif_params({"anthropic_identity_token_file": str(tmp_path / "token")}) + + def test_fleet_wide_env_source_does_not_relabel_a_legacy_param(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("ANTHROPIC_IDENTITY_SOURCE", "internal_issuer") - assert resolve_anthropic_wif_params({"anthropic_identity_token": "oidc/env/TOK"}) is None + with pytest.raises(litellm.AuthenticationError) as exc_info: + resolve_anthropic_wif_params({"anthropic_identity_token": "oidc/env/TOK"}) + + assert "anthropic_identity_token is set, but" in exc_info.value.message + assert "internal_issuer" not in exc_info.value.message + + def test_env_token_file_without_ids_still_returns_none(self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path): + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN_FILE", str(tmp_path / "token")) + assert resolve_anthropic_wif_params({"anthropic_federation_rule_id": "fdrl_1"}) is None + + def test_environment_ids_complete_a_token_file_param(self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path): + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_env") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-env") + token_file: Final = tmp_path / "token" + + params: Final = resolve_anthropic_wif_params({"anthropic_identity_token_file": str(token_file)}) + + assert params is not None + assert (params.federation_rule_id, params.organization_id) == ("fdrl_env", "org-env") + assert params.assertion_ref == f"oidc/file/{token_file}" + + def test_disabling_federation_wins_over_the_gate(self, tmp_path: Path): + litellm_params: Final = { + "anthropic_disable_workload_identity_federation": True, + "anthropic_identity_token_file": str(tmp_path / "token"), + } + assert resolve_anthropic_wif_params(litellm_params) is None + + def test_facade_raises_without_an_engine_call(self): + poster: Final = ScriptedPoster([token_response()]) + engine: Final = make_engine(poster) + with pytest.raises(litellm.AuthenticationError, match="anthropic_identity_token is set, but"): + get_anthropic_wif_token( + {"anthropic_organization_id": "org-1", "anthropic_identity_token": "oidc/env/TOK"}, + None, + "claude-haiku-5-5", + engine, + ) + assert poster.requests == [] class TestConfigYamlShapedIdentitySources: