diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index f122a75bbcc..735a001e343 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -23,6 +23,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( from litellm.llms.anthropic.wif import aget_anthropic_wif_token, get_anthropic_wif_token from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.proxy._types import SpecialHeaders from litellm.types.llms.anthropic import ( ANTHROPIC_HOSTED_TOOLS, ANTHROPIC_OAUTH_BETA_HEADER, @@ -57,7 +58,7 @@ def _strip_bedrock_id_suffixes(model: str) -> str: ) -_SERVER_OWNED_AUTH_HEADERS: Final = frozenset({"x-api-key", "authorization"}) +_SERVER_OWNED_AUTH_HEADERS: Final = SpecialHeaders.litellm_credential_header_names() _WIF_ELIGIBILITY_ATTR: Final = "_workload_identity_eligible" diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index b52c5d82657..bdecc1469e0 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -25,7 +25,11 @@ from litellm.constants import ( ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS, BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES, ) -from litellm.llms.anthropic.common_utils import AnthropicModelInfo, merge_anthropic_beta_headers +from litellm.llms.anthropic.common_utils import ( + _SERVER_OWNED_AUTH_HEADERS, # pyright: ignore[reportPrivateUsage] # canonical set, must not be duplicated here + AnthropicModelInfo, + merge_anthropic_beta_headers, +) from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.proxy._types import * from litellm.proxy.auth.route_checks import RouteChecks @@ -579,6 +583,46 @@ def _anthropic_passthrough_headers(auth_header: Mapping[str, str] | None, client return MappingProxyType({**auth_header, "anthropic-beta": merge_anthropic_beta_headers(client_beta, auth_beta)}) +def _configured_litellm_key_header_name() -> str | None: + from litellm.proxy.proxy_server import ( + general_settings, # pyright: ignore[reportUnknownVariableType] # proxy_server declares it as a bare dict + ) + + configured: Final = general_settings.get( # pyright: ignore[reportUnknownMemberType] # proxy_server general_settings is a bare dict + "litellm_key_header_name" + ) + return configured if isinstance(configured, str) else None + + +def _anthropic_passthrough_header_plan( + request: Request, auth_header: Mapping[str, str] | None, litellm_key_header_name: str | None +) -> tuple[Mapping[str, str], bool]: + """Returns the headers to send upstream plus whether the relay should still forward the + caller's own headers. Once the server owns the Anthropic credential, none of the headers + the proxy accepts a LiteLLM key in (``SpecialHeaders`` plus the configured custom name) + may ride upstream beside it, so the forward merge runs here with those stripped and the + relay is told not to merge again. With no server credential the caller's key is the only + one there is, so forwarding stays on (BYOK).""" + server_headers: Final = _anthropic_passthrough_headers(auth_header, request.headers.get("anthropic-beta")) + if auth_header is None: + return server_headers, True + caller_owned: Final = _SERVER_OWNED_AUTH_HEADERS | frozenset( + (litellm_key_header_name.lower(),) if litellm_key_header_name else () + ) + caller_headers: Final[dict[str, str]] = { # mutable-ok: forward_headers_from_request takes a concrete dict + name: value for name, value in request.headers.items() if name.lower() not in caller_owned + } + merged: Final = cast( # cast-ok: forward_headers_from_request is untyped upstream, its result is a header dict + "dict[str, str]", + HttpPassThroughEndpointHelpers.forward_headers_from_request( # pyright: ignore[reportUnknownMemberType] # untyped upstream + request_headers=caller_headers, + headers=dict(server_headers), # mutable-ok: forward_headers_from_request takes a concrete dict + forward_headers=True, + ), + ) + return MappingProxyType(merged), False + + @router.api_route( "/anthropic/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], @@ -620,11 +664,14 @@ async def anthropic_proxy_route( auth_header: Final = await AnthropicModelInfo.aget_auth_header( anthropic_api_key or None, allow_workload_identity=True ) + upstream_headers, forward_caller_headers = _anthropic_passthrough_header_plan( + request, auth_header, _configured_litellm_key_header_name() + ) endpoint_func: Final = create_pass_through_route( endpoint=endpoint, target=str(updated_url), - custom_headers=_anthropic_passthrough_headers(auth_header, request.headers.get("anthropic-beta")), - _forward_headers=True, + custom_headers=upstream_headers, + _forward_headers=forward_caller_headers, is_streaming_request=is_streaming_request, ) # dynamically construct pass-through endpoint based on incoming path received_value: Final = await endpoint_func( diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index d3647c50902..73738d97259 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -19,9 +19,9 @@ from unittest.mock import patch import pytest -sys.path.insert( - 0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../..")) -) +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../.."))) + +from litellm.proxy._types import SpecialHeaders # noqa: E402 # sys.path must be patched before importing litellm # Fake tokens for testing (not real secrets) FAKE_OAUTH_TOKEN = "sk-ant-oat01-fake-token-for-testing-123456789abcdef" @@ -2069,6 +2069,8 @@ WIF_ENV = { "ANTHROPIC_IDENTITY_TOKEN": "inline-wire-jwt", } +PROXY_CREDENTIAL_HEADER_NAMES = sorted(SpecialHeaders.litellm_credential_header_names()) + class RecordingPoster: def __init__(self, response): @@ -2281,6 +2283,39 @@ class TestWifServerOwnedAuthHeaderStrip: assert caller_key not in headers.values() assert len(poster.requests) == 1 + def test_server_owned_set_is_every_proxy_credential_header(self): + """The strip list must track the proxy's own key-header list, not a hand-rolled + pair: every header user_api_key_auth accepts a LiteLLM key in must be here.""" + from litellm.llms.anthropic.common_utils import _SERVER_OWNED_AUTH_HEADERS + + assert _SERVER_OWNED_AUTH_HEADERS == SpecialHeaders.litellm_credential_header_names() + assert {"x-litellm-api-key", "api-key", "x-goog-api-key"} < _SERVER_OWNED_AUTH_HEADERS + + @pytest.mark.parametrize("header_name", PROXY_CREDENTIAL_HEADER_NAMES) + def test_mint_strips_every_proxy_credential_header(self, monkeypatch, wif_engine, header_name): + """A LiteLLM virtual key arrives in any of the proxy's accepted key headers; once + a mint happened none of them may reach Anthropic in any header slot.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + caller_key = "sk-litellm-CALLER-VIRTUAL-KEY" + + headers = AnthropicModelInfo().validate_environment( + headers={header_name.title(): caller_key, "user-agent": "caller/1.0"}, + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + api_key=None, + api_base=None, + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert header_name == "authorization" or header_name not in {name.lower() for name in headers} + assert all(caller_key not in value for value in headers.values()) + assert headers["user-agent"] == "caller/1.0" + def test_no_mint_preserves_caller_supplied_authorization(self, monkeypatch, clean_anthropic_env): """No-regression: LiteLLM deliberately lets a caller-forwarded credential header ride alongside a statically configured ANTHROPIC_API_KEY, because the diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 6f7ea0298f6..7e48a1563e5 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -35,7 +35,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( vertex_proxy_route, vllm_proxy_route, ) -from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy._types import LitellmUserRoles, SpecialHeaders, UserAPIKeyAuth from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials @@ -2345,7 +2345,6 @@ class TestMilvusProxyRoute: """ from fastapi import HTTPException - mock_request = MagicMock(spec=Request) mock_response = MagicMock(spec=Response) mock_user_api_key_dict = MagicMock() @@ -2379,7 +2378,6 @@ class TestMilvusProxyRoute: """ from fastapi import HTTPException - mock_request = MagicMock(spec=Request) mock_response = MagicMock(spec=Response) mock_user_api_key_dict = MagicMock() @@ -2406,7 +2404,6 @@ class TestMilvusProxyRoute: """ from fastapi import HTTPException - collection_name = "test-collection" mock_request = MagicMock(spec=Request) @@ -2443,7 +2440,6 @@ class TestMilvusProxyRoute: """ from fastapi import HTTPException - collection_name = "unmanaged-collection" mock_request = MagicMock(spec=Request) @@ -3946,3 +3942,229 @@ class TestAnthropicProxyRoute: assert custom_headers["anthropic-beta"] == "oauth-2025-04-20" assert sync_calls == [] assert thread_ids and thread_ids[0] != threading.get_ident() + + +class TestAnthropicProxyRouteCallerAuthHeaders: + """Regression for a caller credential riding upstream next to a server-owned one. + + /anthropic forwards the caller's headers, so a caller-supplied ``x-api-key`` used to reach + Anthropic alongside the server-minted ``Authorization: Bearer``. These drive the real relay + (only the httpx client is stubbed) and assert on the bytes actually handed to the upstream. + """ + + _MINTED: Final = "sk-ant-oat01-plan-minted" + + def _clear_anthropic_env(self, monkeypatch) -> None: + for name in ( + "ANTHROPIC_API_KEY", + "ANTHROPIC_AUTH_TOKEN", + "ANTHROPIC_API_BASE", + "ANTHROPIC_BASE_URL", + "ANTHROPIC_FEDERATION_RULE_ID", + "ANTHROPIC_ORGANIZATION_ID", + "ANTHROPIC_IDENTITY_TOKEN_FILE", + "ANTHROPIC_IDENTITY_TOKEN", + ): + monkeypatch.delenv(name, raising=False) + + def _enable_wif(self, monkeypatch) -> None: + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.wif import aget_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_plan") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-plan") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "plan-inline-jwt") + + minted: Final = self._MINTED + + class StubPoster: + def post(self, url, *, content, headers, timeout): + return httpx.Response( + 200, + json={"access_token": minted, "token_type": "Bearer", "expires_in": 3600}, + ) + + engine: Final = JwtBearerTokenExchangeEngine(poster=StubPoster()) + + async def async_shim(litellm_params, api_base, model): + return await aget_anthropic_wif_token(litellm_params, api_base, model, engine) + + monkeypatch.setattr(anthropic_common_utils, "aget_anthropic_wif_token", async_shim) + + def _request(self, headers: Mapping[str, str]) -> Request: + body: Final = b'{"model":"claude-sonnet-4-5","messages":[]}' + scope: Final = { + "type": "http", + "http_version": "1.1", + "method": "POST", + "scheme": "https", + "path": "/anthropic/v1/messages", + "raw_path": b"/anthropic/v1/messages", + "root_path": "", + "query_string": b"", + "headers": [(name.lower().encode(), value.encode()) for name, value in headers.items()], + "client": ("127.0.0.1", 51234), + "server": ("proxy.local", 4000), + "state": {}, + } + + async def receive() -> dict: + return {"type": "http.request", "body": body, "more_body": False} + + return Request(scope, receive) + + async def _upstream_headers(self, request: Request) -> dict: + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + anthropic_proxy_route, + ) + + upstream_response: Final = MagicMock() + upstream_response.status_code = 200 + upstream_response.headers = {"content-type": "application/json"} + upstream_response.aread = AsyncMock(return_value=b'{"ok": true}') + upstream_response.aiter_bytes = AsyncMock(return_value=[b'{"ok": true}']) + + httpx_client: Final = MagicMock() + httpx_client.build_request = MagicMock(return_value=MagicMock()) + httpx_client.send = AsyncMock(return_value=upstream_response) + client_wrapper: Final = MagicMock() + client_wrapper.client = httpx_client + + with ( + patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client", + return_value=client_wrapper, + ), + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging_obj, + ): + mock_logging_obj.pre_call_hook = AsyncMock(return_value={"model": "claude-sonnet-4-5", "messages": []}) + mock_logging_obj.post_call_success_hook = AsyncMock() + mock_logging_obj.post_call_failure_hook = AsyncMock() + mock_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={}) + + await anthropic_proxy_route( + endpoint="v1/messages", + request=request, + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-caller-virtual-key"), + ) + + assert httpx_client.send.called + return {name.lower(): value for name, value in dict(httpx_client.build_request.call_args[1]["headers"]).items()} + + @pytest.mark.asyncio + async def test_wif_credential_drops_caller_supplied_api_key(self, monkeypatch): + self._clear_anthropic_env(monkeypatch) + self._enable_wif(monkeypatch) + + sent: Final = await self._upstream_headers( + self._request( + { + "content-type": "application/json", + "x-api-key": "sk-caller-virtual-key", + "user-agent": "caller/1.0", + } + ) + ) + + assert sent["authorization"] == f"Bearer {self._MINTED}" + assert "x-api-key" not in sent + assert sent["user-agent"] == "caller/1.0" + + @pytest.mark.asyncio + async def test_wif_credential_drops_caller_supplied_authorization(self, monkeypatch): + self._clear_anthropic_env(monkeypatch) + self._enable_wif(monkeypatch) + + sent: Final = await self._upstream_headers( + self._request( + { + "content-type": "application/json", + "authorization": "Bearer sk-caller-virtual-key", + } + ) + ) + + assert sent["authorization"] == f"Bearer {self._MINTED}" + assert all("sk-caller-virtual-key" not in value for value in sent.values()) + + @pytest.mark.asyncio + @pytest.mark.parametrize("header_name", sorted(SpecialHeaders.litellm_credential_header_names())) + async def test_wif_credential_drops_every_proxy_key_header(self, monkeypatch, header_name: str): + """The proxy accepts a LiteLLM key in any SpecialHeaders slot, so the caller's virtual + key must not reach Anthropic from any of them once the server owns the credential.""" + self._clear_anthropic_env(monkeypatch) + self._enable_wif(monkeypatch) + + sent: Final = await self._upstream_headers( + self._request( + { + "content-type": "application/json", + header_name: "sk-caller-virtual-key", + "user-agent": "caller/1.0", + } + ) + ) + + assert sent["authorization"] == f"Bearer {self._MINTED}" + assert header_name == "authorization" or header_name not in sent + assert all("sk-caller-virtual-key" not in value for value in sent.values()) + assert sent["user-agent"] == "caller/1.0" + + @pytest.mark.asyncio + async def test_wif_credential_drops_configured_custom_key_header(self, monkeypatch): + self._clear_anthropic_env(monkeypatch) + self._enable_wif(monkeypatch) + + with patch.dict("litellm.proxy.proxy_server.general_settings", {"litellm_key_header_name": "X-Tenant-Key"}): + sent: Final = await self._upstream_headers( + self._request( + { + "content-type": "application/json", + "x-tenant-key": "sk-caller-virtual-key", + "x-tenant-region": "eu", + } + ) + ) + + assert sent["authorization"] == f"Bearer {self._MINTED}" + assert "x-tenant-key" not in sent + assert all("sk-caller-virtual-key" not in value for value in sent.values()) + assert sent["x-tenant-region"] == "eu" + + @pytest.mark.asyncio + async def test_server_api_key_drops_caller_supplied_authorization(self, monkeypatch): + self._clear_anthropic_env(monkeypatch) + monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-server-owned") + + sent: Final = await self._upstream_headers( + self._request( + { + "content-type": "application/json", + "authorization": "Bearer sk-caller-virtual-key", + "x-api-key": "sk-caller-virtual-key", + } + ) + ) + + assert sent["x-api-key"] == "sk-ant-server-owned" + assert "authorization" not in sent + + @pytest.mark.asyncio + async def test_byok_caller_key_still_reaches_upstream(self, monkeypatch): + self._clear_anthropic_env(monkeypatch) + + sent: Final = await self._upstream_headers( + self._request( + { + "content-type": "application/json", + "x-api-key": "sk-ant-caller-owned", + "anthropic-version": "2023-06-01", + } + ) + ) + + assert sent["x-api-key"] == "sk-ant-caller-owned" + assert sent["anthropic-version"] == "2023-06-01" + assert "authorization" not in sent