diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py b/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py index 2e752fc9ff6..5231f952f2c 100644 --- a/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py +++ b/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py @@ -111,9 +111,7 @@ class UpstreamCredentialProvider: case SharedKey() as source: # The shared key is read straight from ServerSpec.config, not the per-user # store; it is the same credential for every caller. - return Ok( - StaticHeaderAuth(config.header_for(source.value.get_secret_value())) - ) + return Ok(_api_key_auth(config, source.value.get_secret_value())) case PerUserEnvVar(): fetched = await self._per_user_value(subject, server) if isinstance(fetched, Error): @@ -125,7 +123,7 @@ class UpstreamCredentialProvider: "api_key: per-user env var not set for this subject" ) ) - return Ok(StaticHeaderAuth(config.header_for(value))) + return Ok(_api_key_auth(config, value)) case Byok(): fetched = await self._per_user_value(subject, server) if isinstance(fetched, Error): @@ -137,7 +135,7 @@ class UpstreamCredentialProvider: "api_key: no BYOK credential for this subject" ) ) - return Ok(StaticHeaderAuth(config.header_for(value))) + return Ok(_api_key_auth(config, value)) assert_never(config.key_source) async def _per_user_value( @@ -276,3 +274,8 @@ class UpstreamCredentialProvider: def _bearer(token: StoredToken) -> StaticHeaderAuth: return StaticHeaderAuth(f"Bearer {token.access_token.get_secret_value()}") + + +def _api_key_auth(config: ApiKeyConfig, value: str) -> StaticHeaderAuth: + header_name, header_value = config.header(value) + return StaticHeaderAuth(header_value, header_name=header_name) diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/types.py b/litellm/proxy/gateway/mcp/outbound_credentials/types.py index 36d0f0008ab..6e45f7162a8 100644 --- a/litellm/proxy/gateway/mcp/outbound_credentials/types.py +++ b/litellm/proxy/gateway/mcp/outbound_credentials/types.py @@ -134,9 +134,6 @@ class CredError: assert_never(self.tag) -ApiKeyScheme = Literal["bearer", "apikey", "basic", "token", "raw"] - - class AuthorizationCodeConfig(BaseModel): """Per-user 3LO; the gateway is the OAuth client and stores the user's token. @@ -212,27 +209,20 @@ ApiKeySource = Annotated[ class ApiKeyConfig(BaseModel): """A fixed credential injected as a header. The value is shared (in config) or seeded - per-user (pulled from the store); `scheme` is how it is written into the header.""" + per-user (pulled from the store); `header_name` and `value_prefix` say where and how it is + written, modeled like OpenAPI's apiKey scheme so any upstream convention is expressible + (Authorization + Bearer, a raw value on X-API-Key, Ocp-Apim-Subscription-Key, etc.). + """ model_config = ConfigDict(frozen=True) kind: Literal[AuthSpecKind.api_key] = AuthSpecKind.api_key - scheme: ApiKeyScheme = "bearer" + header_name: str = "Authorization" + value_prefix: str = "Bearer" key_source: ApiKeySource - def header_for(self, value: str) -> str: - # Reconstructs v1's per-scheme Authorization prefixes (mcp_server_manager.py:877-884). - match self.scheme: - case "bearer": - return f"Bearer {value}" - case "apikey": - return f"ApiKey {value}" - case "basic": - return f"Basic {value}" - case "token": - return f"token {value}" - case "raw": - return value - assert_never(self.scheme) + def header(self, value: str) -> tuple[str, str]: + formatted = f"{self.value_prefix} {value}" if self.value_prefix else value + return self.header_name, formatted class PassthroughConfig(BaseModel): diff --git a/tests/mcp_tests/gateway/test_resolver.py b/tests/mcp_tests/gateway/test_resolver.py index 590fc8d9098..e599396257d 100644 --- a/tests/mcp_tests/gateway/test_resolver.py +++ b/tests/mcp_tests/gateway/test_resolver.py @@ -262,21 +262,28 @@ async def test_none_attaches_no_credential(): @pytest.mark.parametrize( - "scheme,expected", + "header_name,value_prefix,expected_header,expected_value", [ - ("bearer", "Bearer k"), - ("apikey", "ApiKey k"), - ("basic", "Basic k"), - ("token", "token k"), - ("raw", "k"), + ("Authorization", "Bearer", "Authorization", "Bearer k"), # v1 bearer_token + ("Authorization", "Basic", "Authorization", "Basic k"), # v1 basic + ("Authorization", "token", "Authorization", "token k"), # v1 token + ("Authorization", "", "Authorization", "k"), # v1 authorization (raw) + ("X-API-Key", "", "X-API-Key", "k"), # v1 api_key (its own header, raw) + ("Ocp-Apim-Subscription-Key", "", "Ocp-Apim-Subscription-Key", "k"), # custom ], ) -async def test_api_key_emits_the_right_scheme(scheme: str, expected: str): - config = ApiKeyConfig(scheme=scheme, key_source=SharedKey(value="k")) # type: ignore[arg-type] +async def test_api_key_writes_value_to_the_configured_header( + header_name: str, value_prefix: str, expected_header: str, expected_value: str +): + config = ApiKeyConfig( + header_name=header_name, + value_prefix=value_prefix, + key_source=SharedKey(value="k"), + ) result = await PROVIDER.resolve(SUBJECT, _spec(config)) assert isinstance(result, Ok) assert isinstance(result.ok, StaticHeaderAuth) - assert _applied_headers(result.ok)["Authorization"] == expected + assert _applied_headers(result.ok)[expected_header] == expected_value def test_secret_fields_are_masked_in_serialization():