From ccc45044458fe53e4e4df52e7440301afdf0b46c Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 25 Feb 2026 12:08:25 +0530 Subject: [PATCH] Add support for forwarding provider's auth headers --- litellm/proxy/litellm_pre_call_utils.py | 72 +++++++++++------ tests/proxy_unit_tests/test_proxy_server.py | 75 +++++++++++++++++ .../anthropic/test_anthropic_common_utils.py | 81 +++++++++++++++++++ 3 files changed, 203 insertions(+), 25 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 52f0b1d46e9..b128482daf2 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -10,10 +10,15 @@ import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging from litellm.litellm_core_utils.safe_json_loads import safe_json_loads -from litellm.proxy._types import (AddTeamCallback, CommonProxyErrors, - LitellmDataForBackendLLMCall, - LitellmUserRoles, SpecialHeaders, - TeamCallbackMetadata, UserAPIKeyAuth) +from litellm.proxy._types import ( + AddTeamCallback, + CommonProxyErrors, + LitellmDataForBackendLLMCall, + LitellmUserRoles, + SpecialHeaders, + TeamCallbackMetadata, + UserAPIKeyAuth, +) from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers # Cache special headers as a frozenset for O(1) lookup performance @@ -23,9 +28,12 @@ _SPECIAL_HEADERS_CACHE = frozenset( from litellm.router import Router from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS from litellm.types.services import ServiceTypes -from litellm.types.utils import (LlmProviders, ProviderSpecificHeader, - StandardLoggingUserAPIKeyMetadata, - SupportedCacheControls) +from litellm.types.utils import ( + LlmProviders, + ProviderSpecificHeader, + StandardLoggingUserAPIKeyMetadata, + SupportedCacheControls, +) service_logger_obj = ServiceLogging() # used for tracking latency on OTEL @@ -228,7 +236,9 @@ def _get_dynamic_logging_metadata( def clean_headers( - headers: Headers, litellm_key_header_name: Optional[str] = None + headers: Headers, + litellm_key_header_name: Optional[str] = None, + forward_llm_provider_auth_headers: bool = False, ) -> dict: """ Removes litellm api key from headers @@ -238,19 +248,25 @@ def clean_headers( clean_headers = {} litellm_key_lower = ( litellm_key_header_name.lower() if litellm_key_header_name is not None else None - ) - + ) for header, value in headers.items(): header_lower = header.lower() - # Preserve Authorization header if it contains Anthropic OAuth token (sk-ant-oat*) - # This allows OAuth tokens to be forwarded to Anthropic-compatible providers - # via add_provider_specific_headers_to_request() + verbose_proxy_logger.debug(f"header: {header}") + if header_lower == "authorization" and is_anthropic_oauth_key(value): + verbose_proxy_logger.debug(f"Adding Anthropic OAuth header: {header}") + clean_headers[header] = value + elif forward_llm_provider_auth_headers and header_lower in _SPECIAL_HEADERS_CACHE: + if litellm_key_lower and header_lower == litellm_key_lower: + continue + if header_lower == "authorization": + continue clean_headers[header] = value # Check if header should be excluded: either in special headers cache or matches custom litellm key elif header_lower not in _SPECIAL_HEADERS_CACHE and ( litellm_key_lower is None or header_lower != litellm_key_lower ): + verbose_proxy_logger.debug(f"Adding header and value: {header} {value}") clean_headers[header] = value return clean_headers @@ -654,7 +670,8 @@ class LiteLLMProxyRequestSetup: return data from litellm.proxy._types import ( LiteLLM_ManagementEndpoint_MetadataFields, - LiteLLM_ManagementEndpoint_MetadataFields_Premium) + LiteLLM_ManagementEndpoint_MetadataFields_Premium, + ) # ignore any special fields added_metadata = {} @@ -826,6 +843,11 @@ async def add_litellm_data_to_request( # noqa: PLR0915 from litellm.types.proxy.litellm_pre_call_utils import SecretFields _raw_headers: Dict[str, str] = _safe_get_request_headers(request) + + forward_llm_auth = False + if general_settings: + forward_llm_auth = general_settings.get("forward_llm_provider_auth_headers", False) + _headers: Dict[str, str] = clean_headers( request.headers, litellm_key_header_name=( @@ -833,7 +855,10 @@ async def add_litellm_data_to_request( # noqa: PLR0915 if general_settings is not None else None ), + forward_llm_provider_auth_headers=forward_llm_auth, ) + verbose_proxy_logger.debug(f"Request Headers: {_headers}") + verbose_proxy_logger.debug(f"Raw Headers: {_raw_headers}") ########################################################## # Init - Proxy Server Request @@ -1479,8 +1504,7 @@ async def move_guardrails_to_metadata( # Only check policy engine if no local config (avoid import + registry lookup) if not (has_key_config or has_team_config or has_request_config): - from litellm.proxy.policy_engine.policy_registry import \ - get_policy_registry + from litellm.proxy.policy_engine.policy_registry import get_policy_registry if not get_policy_registry().is_initialized(): # Nothing configured anywhere - clean up request body fields and return @@ -1544,16 +1568,14 @@ async def move_guardrails_to_metadata( def _is_policy_version_id(s: str) -> bool: """Return True if string is a policy version ID (starts with policy_ prefix).""" - from litellm.proxy.policy_engine.policy_registry import \ - POLICY_VERSION_ID_PREFIX + from litellm.proxy.policy_engine.policy_registry import POLICY_VERSION_ID_PREFIX return isinstance(s, str) and s.startswith(POLICY_VERSION_ID_PREFIX) def _extract_policy_id(s: str) -> Optional[str]: """Extract raw UUID from policy_ string, or None if not a valid version ID.""" - from litellm.proxy.policy_engine.policy_registry import \ - POLICY_VERSION_ID_PREFIX + from litellm.proxy.policy_engine.policy_registry import POLICY_VERSION_ID_PREFIX if not _is_policy_version_id(s): return None @@ -1574,9 +1596,10 @@ def _match_and_track_policies( """ from litellm._logging import verbose_proxy_logger from litellm.proxy.common_utils.callback_utils import ( - add_policy_sources_to_metadata, add_policy_to_applied_policies_header) - from litellm.proxy.policy_engine.attachment_registry import \ - get_attachment_registry + add_policy_sources_to_metadata, + add_policy_to_applied_policies_header, + ) + from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher # Get matching policies via attachments (with match reasons for attribution) @@ -1721,8 +1744,7 @@ async def add_guardrails_from_policy_engine( user_api_key_dict: The user's API key authentication info """ from litellm._logging import verbose_proxy_logger - from litellm.proxy.common_utils.http_parsing_utils import \ - get_tags_from_request_body + from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body from litellm.proxy.policy_engine.policy_registry import get_policy_registry from litellm.types.proxy.policy_engine import PolicyMatchContext diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index 0cbba7b5cc3..d39edef793c 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -334,6 +334,81 @@ def test_chat_completion_forward_headers( pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}") +@pytest.mark.parametrize("forward_llm_auth_headers", [True, False]) +@mock_patch_acompletion() +def test_chat_completion_forward_llm_provider_auth_headers( + mock_acompletion, client_no_auth, forward_llm_auth_headers +): + """ + Test that LLM provider auth headers (x-api-key, x-goog-api-key) are forwarded + when forward_llm_provider_auth_headers=True. + + This allows clients to send their own LLM provider API keys through the proxy. + """ + try: + # Configure general settings + gs = getattr(litellm.proxy.proxy_server, "general_settings") + gs["forward_client_headers_to_llm_api"] = True + gs["forward_llm_provider_auth_headers"] = forward_llm_auth_headers + setattr(litellm.proxy.proxy_server, "general_settings", gs) + + # Test data + test_data = { + "model": "gpt-3.5-turbo", + "messages": [ + {"role": "user", "content": "hello"}, + ], + "max_tokens": 10, + } + + # Headers including LLM provider auth + request_headers = { + "Authorization": "Bearer sk-proxy-auth-123", # Proxy auth (should be stripped) + "x-api-key": "sk-ant-api03-test-anthropic-key", # Anthropic API key + "x-goog-api-key": "google-api-key-123", # Google API key + "X-Custom-Header": "custom-value", # Custom header (should be forwarded) + } + + # Make request + response = client_no_auth.post( + "/v1/chat/completions", json=test_data, headers=request_headers + ) + + assert response.status_code == 200 + + # Check forwarded headers + forwarded_headers = mock_acompletion.call_args.kwargs.get("headers", {}) + + if forward_llm_auth_headers: + # LLM provider auth headers should be forwarded + assert "x-api-key" in forwarded_headers + assert forwarded_headers["x-api-key"] == "sk-ant-api03-test-anthropic-key" + assert "x-goog-api-key" in forwarded_headers + assert forwarded_headers["x-goog-api-key"] == "google-api-key-123" + else: + # LLM provider auth headers should be stripped + assert "x-api-key" not in forwarded_headers + assert "x-goog-api-key" not in forwarded_headers + + # Custom headers should always be forwarded (when forward_client_headers_to_llm_api=True) + assert "x-custom-header" in forwarded_headers + assert forwarded_headers["x-custom-header"] == "custom-value" + + # Proxy Authorization should never be forwarded + assert "authorization" not in forwarded_headers + + print(f"✓ Test passed with forward_llm_provider_auth_headers={forward_llm_auth_headers}") + print(f" Forwarded headers: {list(forwarded_headers.keys())}") + + except Exception as e: + pytest.fail(f"Test failed with forward_llm_auth_headers={forward_llm_auth_headers}: {str(e)}") + finally: + # Clean up + gs = getattr(litellm.proxy.proxy_server, "general_settings") + gs.pop("forward_llm_provider_auth_headers", None) + setattr(litellm.proxy.proxy_server, "general_settings", gs) + + @mock_patch_acompletion() @pytest.mark.asyncio async def test_team_disable_guardrails(mock_acompletion, client_no_auth): 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 ebffb56446e..3729e67f0da 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -363,6 +363,87 @@ class TestProxyOAuthHeaderForwarding: assert "authorization" not in cleaned assert cleaned["content-type"] == "application/json" + def test_clean_headers_forwards_anthropic_api_key_when_enabled(self): + """clean_headers should preserve x-api-key when forward_llm_provider_auth_headers=True.""" + from starlette.datastructures import Headers + + from litellm.proxy.litellm_pre_call_utils import clean_headers + + raw_headers = Headers( + raw=[ + (b"authorization", b"Bearer sk-proxy-auth"), + (b"x-api-key", b"sk-ant-api03-test-key"), + (b"content-type", b"application/json"), + ] + ) + cleaned = clean_headers(raw_headers, forward_llm_provider_auth_headers=True) + + # x-api-key should be preserved when flag is True + assert "x-api-key" in cleaned + assert cleaned["x-api-key"] == "sk-ant-api03-test-key" + # Authorization (proxy auth) should still be stripped + assert "authorization" not in cleaned + assert cleaned["content-type"] == "application/json" + + def test_clean_headers_strips_anthropic_api_key_when_disabled(self): + """clean_headers should strip x-api-key when forward_llm_provider_auth_headers=False (default).""" + from starlette.datastructures import Headers + + from litellm.proxy.litellm_pre_call_utils import clean_headers + + raw_headers = Headers( + raw=[ + (b"x-api-key", b"sk-ant-api03-test-key"), + (b"content-type", b"application/json"), + ] + ) + cleaned = clean_headers(raw_headers, forward_llm_provider_auth_headers=False) + + # x-api-key should be stripped by default + assert "x-api-key" not in cleaned + assert cleaned["content-type"] == "application/json" + + def test_clean_headers_forwards_google_api_key_when_enabled(self): + """clean_headers should preserve x-goog-api-key when forward_llm_provider_auth_headers=True.""" + from starlette.datastructures import Headers + + from litellm.proxy.litellm_pre_call_utils import clean_headers + + raw_headers = Headers( + raw=[ + (b"x-goog-api-key", b"google-api-key-123"), + (b"content-type", b"application/json"), + ] + ) + cleaned = clean_headers(raw_headers, forward_llm_provider_auth_headers=True) + + assert "x-goog-api-key" in cleaned + assert cleaned["x-goog-api-key"] == "google-api-key-123" + assert cleaned["content-type"] == "application/json" + + def test_clean_headers_preserves_oauth_regardless_of_forward_flag(self): + """clean_headers should always preserve OAuth tokens regardless of forward_llm_provider_auth_headers.""" + from starlette.datastructures import Headers + + from litellm.proxy.litellm_pre_call_utils import clean_headers + + raw_headers = Headers( + raw=[ + (b"authorization", f"Bearer {FAKE_OAUTH_TOKEN}".encode()), + (b"content-type", b"application/json"), + ] + ) + + # Should preserve OAuth even with flag=False + cleaned_without_flag = clean_headers(raw_headers, forward_llm_provider_auth_headers=False) + assert "authorization" in cleaned_without_flag + assert cleaned_without_flag["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" + + # Should also preserve OAuth with flag=True + cleaned_with_flag = clean_headers(raw_headers, forward_llm_provider_auth_headers=True) + assert "authorization" in cleaned_with_flag + assert cleaned_with_flag["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" + def test_add_provider_specific_headers_forwards_oauth(self): """add_provider_specific_headers_to_request should forward OAuth Authorization as a ProviderSpecificHeader scoped to Anthropic-compatible providers."""