mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
refactor(mcp): model api_key placement as (header_name, value_prefix)
The scheme enum assumed every static credential rides on the Authorization header with a prefix, which is true for bearer/basic/token/raw but wrong for api_key: v1 writes that one to the X-API-Key header (client.py:443-444), a different header entirely. Parity-testing the graft surfaced it. API-key headers are not standardized (OpenAPI models the header name as configurable; real upstreams vary - Atlassian uses Authorization: Bearer, MS Logic Apps and Spring AI use X-API-Key, Azure APIM uses Ocp-Apim-Subscription-Key), so model the placement as data: header_name (default Authorization) + value_prefix (default Bearer), like OpenAPI's apiKey scheme. This expresses every v1 auth_type and any custom upstream header, with no per-scheme enum and no leak. Replaces ApiKeyConfig.scheme/header_for with header_name/value_prefix/header(); the arm builds StaticHeaderAuth(value, header_name=name) via a small helper. The capability is not yet exposed to admins (that is the deferred server-configuration phase); the v1->v2 adapter will fill the pair from v1's auth_type. Test now covers bearer/basic/token/raw on Authorization plus X-API-Key and a custom header. 56 tests, gates green.
This commit is contained in:
parent
dcd97c6c58
commit
4b55de7f45
3 changed files with 33 additions and 33 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue