mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(mcp): add v1 bridge + none/api_key resolver arms (unwired)
PR4a of the MCP v2 outbound-credential migration, stacked on the resolver skeleton. Builds the bridge for the first live modes without wiring it onto the request path: - resolver.py: the none arm (NoOpAuth) and the api_key shared-key arm (StaticHeaderAuth from the config); the BYOK source and the other five arms stay not_implemented. - adapter.py: the v1 <-> v2 edge (to_subject, to_server_spec, raise_public, should_defer). to_server_spec maps only none + the static-header family and returns None to defer every other mode to v1. Imports v1, kept out of the package __init__ so the resolver core stays v1-free. - MCPClient gains an optional resolved_auth that feeds the factory's auth= slot, taking precedence over the SigV4 aws_auth; default None keeps current behavior. Nothing calls these from _create_mcp_client yet, so production behavior is unchanged; the graft lands in PR4b. Unit tests cover the two arms, the full mapping table, and the auth plumbing.
This commit is contained in:
parent
1d5ab42e14
commit
161d438791
6 changed files with 440 additions and 42 deletions
|
|
@ -224,6 +224,7 @@ class MCPClient:
|
|||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
ssl_verify: Optional[VerifyTypes] = None,
|
||||
aws_auth: Optional[httpx.Auth] = None,
|
||||
resolved_auth: Optional[httpx.Auth] = None,
|
||||
sampling_callback: Optional[Callable] = None,
|
||||
elicitation_callback: Optional[Callable] = None,
|
||||
logging_callback: Optional[Callable] = None,
|
||||
|
|
@ -237,6 +238,9 @@ class MCPClient:
|
|||
self.extra_headers: Optional[Dict[str, str]] = extra_headers
|
||||
self.ssl_verify: Optional[VerifyTypes] = ssl_verify
|
||||
self._aws_auth: Optional[httpx.Auth] = aws_auth
|
||||
# A pre-resolved httpx.Auth (e.g. from the v2 credential resolver) attached to the
|
||||
# upstream client's auth= slot, taking precedence over the SigV4 aws_auth.
|
||||
self._resolved_auth: Optional[httpx.Auth] = resolved_auth
|
||||
self._last_initialize_instructions: Optional[str] = None
|
||||
self._sampling_callback: Optional[Callable] = sampling_callback
|
||||
self._elicitation_callback: Optional[Callable] = elicitation_callback
|
||||
|
|
@ -482,11 +486,15 @@ class MCPClient:
|
|||
verbose_logger.debug(
|
||||
f"MCP client using SSL configuration: {type(ssl_config).__name__}"
|
||||
)
|
||||
# Use SigV4 auth if configured and no explicit auth provided.
|
||||
# The MCP SDK's sse_client and streamable_http_client call this
|
||||
# factory without passing auth=, so self._aws_auth is used.
|
||||
# For non-SigV4 clients, self._aws_auth is None — no behavior change.
|
||||
effective_auth = auth if auth is not None else self._aws_auth
|
||||
# The MCP SDK's sse_client and streamable_http_client call this factory without
|
||||
# passing auth=, so the fallback is used: a v2-resolved auth if present, else the
|
||||
# SigV4 aws_auth. Both are None for the common case — no behavior change.
|
||||
fallback_auth = (
|
||||
self._resolved_auth
|
||||
if self._resolved_auth is not None
|
||||
else self._aws_auth
|
||||
)
|
||||
effective_auth = auth if auth is not None else fallback_auth
|
||||
return httpx.AsyncClient(
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,140 @@
|
|||
"""The v1 <-> v2 bridge for the credential resolver.
|
||||
|
||||
These edge functions translate v1's request objects into the resolver's typed inputs and map
|
||||
its typed errors onto the proxy's public exception contract. They import v1 and live outside the
|
||||
package's public surface so the resolver core (``resolver.py`` / ``types.py``) stays v1-free.
|
||||
Nothing wires them into ``_create_mcp_client`` yet.
|
||||
|
||||
``to_server_spec`` maps only the modes the resolver has gone live for, returning ``None`` for
|
||||
every other mode so the caller defers to v1 (parity-safe); it grows one branch per migrated mode.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
from typing import TYPE_CHECKING, NoReturn, Optional
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import SecretStr
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
ApiKeyConfig,
|
||||
CredError,
|
||||
NoneConfig,
|
||||
ServerSpec,
|
||||
SharedKey,
|
||||
Subject,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
def to_subject(
|
||||
user_api_key_auth: Optional["UserAPIKeyAuth"], subject_token: Optional[str]
|
||||
) -> Subject:
|
||||
"""Map v1's authenticated principal onto the resolver's Subject.
|
||||
|
||||
tenant_id / subject_id are empty for an unauthenticated caller; the per-user arms must reject
|
||||
an empty subject rather than share one credential slot across callers.
|
||||
"""
|
||||
inbound = SecretStr(subject_token) if subject_token else None
|
||||
if user_api_key_auth is None:
|
||||
return Subject(tenant_id="", subject_id="", inbound_token=inbound)
|
||||
return Subject(
|
||||
tenant_id=user_api_key_auth.org_id or user_api_key_auth.team_id or "",
|
||||
subject_id=user_api_key_auth.user_id or "",
|
||||
inbound_token=inbound,
|
||||
)
|
||||
|
||||
|
||||
def to_server_spec(server: "MCPServer") -> Optional[ServerSpec]:
|
||||
"""Map a v1 server onto a ServerSpec for a migrated mode, or None to defer to v1.
|
||||
|
||||
BYOK is the per-user source of the ``api_key`` mode; its scheme rides on ``auth_type`` just
|
||||
like a shared key, but the value is per-user and not migrated yet, so a BYOK server defers
|
||||
to v1 regardless of ``auth_type`` (this guard is the seam the BYOK arm replaces later).
|
||||
|
||||
Dispatches on the declared ``auth_type``. The match is exhaustive over ``MCPAuthType`` with
|
||||
an ``assert_never`` tail, so a newly added auth mode fails the type gate here until it is
|
||||
explicitly mapped or explicitly deferred, rather than silently falling through to v1. Live
|
||||
modes: ``none`` and the static-header family (``api_key`` plus the Authorization schemes),
|
||||
all shared-key; every other mode returns None and stays on v1.
|
||||
"""
|
||||
if server.is_byok:
|
||||
return (
|
||||
None # per-user BYOK source not migrated yet -> defer to v1 (any auth_type)
|
||||
)
|
||||
resource = server.url or server.server_id
|
||||
auth_type = server.auth_type
|
||||
match auth_type:
|
||||
case None | MCPAuth.none:
|
||||
if server.is_oauth_passthrough:
|
||||
return None # passthrough is not migrated yet -> defer to v1
|
||||
return ServerSpec(
|
||||
server_id=server.server_id, resource=resource, config=NoneConfig()
|
||||
)
|
||||
case MCPAuth.api_key:
|
||||
return _shared_key_spec(server, resource, "X-API-Key", "")
|
||||
case MCPAuth.bearer_token:
|
||||
return _shared_key_spec(server, resource, "Authorization", "Bearer")
|
||||
case MCPAuth.token:
|
||||
return _shared_key_spec(server, resource, "Authorization", "token")
|
||||
case MCPAuth.authorization:
|
||||
return _shared_key_spec(server, resource, "Authorization", "")
|
||||
case MCPAuth.basic:
|
||||
return _shared_key_spec(
|
||||
server, resource, "Authorization", "Basic", encode=True
|
||||
)
|
||||
case MCPAuth.oauth2 | MCPAuth.oauth2_token_exchange | MCPAuth.aws_sigv4:
|
||||
return None # OAuth grants and SigV4 are not migrated yet -> defer to v1
|
||||
assert_never(auth_type)
|
||||
|
||||
|
||||
def _shared_key_spec(
|
||||
server: "MCPServer",
|
||||
resource: str,
|
||||
header_name: str,
|
||||
value_prefix: str,
|
||||
*,
|
||||
encode: bool = False,
|
||||
) -> Optional[ServerSpec]:
|
||||
"""Build an api_key spec from the server's static token, or defer (None) if it is absent.
|
||||
|
||||
Covers the whole shared-key static-header family: ``api_key`` on ``X-API-Key`` and the
|
||||
Authorization schemes (bearer / token / authorization sent verbatim, basic base64-encoded).
|
||||
"""
|
||||
token = server.authentication_token
|
||||
if not token:
|
||||
return None # no key configured -> defer to v1 (parity-safe)
|
||||
value = base64.b64encode(token.encode("utf-8")).decode() if encode else token
|
||||
return ServerSpec(
|
||||
server_id=server.server_id,
|
||||
resource=resource,
|
||||
config=ApiKeyConfig(
|
||||
header_name=header_name,
|
||||
value_prefix=value_prefix,
|
||||
key_source=SharedKey(value=SecretStr(value)),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def raise_public(error: CredError) -> NoReturn:
|
||||
"""Map a resolver CredError onto the proxy's public HTTP contract. The one edge that raises."""
|
||||
match error.tag:
|
||||
case "unauthorized":
|
||||
raise HTTPException(status_code=401, detail=error.summary)
|
||||
case "misconfigured":
|
||||
raise HTTPException(status_code=500, detail=error.summary)
|
||||
case "upstream_unavailable":
|
||||
raise HTTPException(status_code=503, detail=error.summary)
|
||||
case "unsupported_mode":
|
||||
raise HTTPException(status_code=500, detail=error.summary)
|
||||
case "precondition_required":
|
||||
raise HTTPException(status_code=412, detail=error.summary)
|
||||
case "not_implemented":
|
||||
raise HTTPException(status_code=501, detail=error.summary)
|
||||
assert_never(error.tag)
|
||||
|
|
@ -7,9 +7,9 @@ no precedence cascade. It is wildcard-free with an `assert_never` tail, so addin
|
|||
an arm fails the type gate (basedpyright `reportMatchNotExhaustive`); a bypassed gate fails loudly
|
||||
at runtime instead of returning `None`.
|
||||
|
||||
This skeleton ships every arm as a `not_implemented` stub. Each mode's real body, with its
|
||||
injected seam, lands in its own follow-up PR; until then the arm returns a typed error rather
|
||||
than silently producing no credential. Pure v2: no imports from v1.
|
||||
`none` and `api_key` (shared-key source) are live; the remaining arms are `not_implemented`
|
||||
stubs that each land in a follow-up PR with their injected seam. The self-contained arms read
|
||||
straight from the config and need no collaborator. Pure v2: no imports from v1.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -17,8 +17,13 @@ from __future__ import annotations
|
|||
import httpx
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
|
||||
NoOpAuth,
|
||||
StaticHeaderAuth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
|
||||
Error,
|
||||
Ok,
|
||||
Result,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
|
|
@ -26,11 +31,13 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
|||
AuthorizationCodeConfig,
|
||||
AuthSpecKind,
|
||||
AwsSigV4Config,
|
||||
Byok,
|
||||
ClientCredentialsConfig,
|
||||
CredError,
|
||||
NoneConfig,
|
||||
PassthroughConfig,
|
||||
ServerSpec,
|
||||
SharedKey,
|
||||
Subject,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
|
|
@ -40,7 +47,7 @@ class UpstreamCredentialProvider:
|
|||
"""Produces the one `httpx.Auth` for a `(subject, upstream)` pair, per declared mode.
|
||||
|
||||
Collaborators (the per-mode credential stores and token fetchers) are injected as each arm
|
||||
is built; the skeleton needs none, since every arm is a stub.
|
||||
is built; the live `none` and `api_key`-shared arms read from the config and need none.
|
||||
"""
|
||||
|
||||
async def resolve_credentials(
|
||||
|
|
@ -48,9 +55,9 @@ class UpstreamCredentialProvider:
|
|||
) -> Result[httpx.Auth, CredError]:
|
||||
match server.config:
|
||||
case NoneConfig():
|
||||
return _not_implemented(AuthSpecKind.none)
|
||||
case ApiKeyConfig():
|
||||
return _not_implemented(AuthSpecKind.api_key)
|
||||
return Ok(NoOpAuth())
|
||||
case ApiKeyConfig() as config:
|
||||
return self._api_key(config)
|
||||
case PassthroughConfig():
|
||||
return _not_implemented(AuthSpecKind.passthrough)
|
||||
case ClientCredentialsConfig():
|
||||
|
|
@ -63,6 +70,22 @@ class UpstreamCredentialProvider:
|
|||
return _not_implemented(AuthSpecKind.aws_sigv4)
|
||||
assert_never(server.config)
|
||||
|
||||
def _api_key(self, config: ApiKeyConfig) -> Result[httpx.Auth, CredError]:
|
||||
match config.key_source:
|
||||
case SharedKey() as source:
|
||||
header_name, header_value = config.header(
|
||||
source.value.get_secret_value()
|
||||
)
|
||||
return Ok(StaticHeaderAuth(header_value, header_name=header_name))
|
||||
case Byok():
|
||||
# Per-user key pulled from the credential store; lands with that seam.
|
||||
return Error(
|
||||
CredError.of_not_implemented(
|
||||
"api_key BYOK source not implemented yet"
|
||||
)
|
||||
)
|
||||
assert_never(config.key_source)
|
||||
|
||||
|
||||
def _not_implemented(kind: AuthSpecKind) -> Result[httpx.Auth, CredError]:
|
||||
return Error(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import asyncio
|
||||
import os
|
||||
import ssl
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -543,5 +542,45 @@ class TestExecuteSessionOperationSurfacesTransportError:
|
|||
assert result == "done"
|
||||
|
||||
|
||||
class TestMCPClientResolvedAuth:
|
||||
"""A pre-resolved httpx.Auth is attached to the upstream client's auth= slot."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolved_auth_feeds_the_auth_slot(self):
|
||||
resolved = httpx.Auth()
|
||||
client = MCPClient(
|
||||
server_url="https://upstream.example.com", resolved_auth=resolved
|
||||
)
|
||||
http_client = client._create_httpx_client_factory()()
|
||||
try:
|
||||
assert http_client.auth is resolved
|
||||
finally:
|
||||
await http_client.aclose()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolved_auth_takes_precedence_over_aws_auth(self):
|
||||
resolved = httpx.Auth()
|
||||
client = MCPClient(
|
||||
server_url="https://upstream.example.com",
|
||||
resolved_auth=resolved,
|
||||
aws_auth=httpx.Auth(),
|
||||
)
|
||||
http_client = client._create_httpx_client_factory()()
|
||||
try:
|
||||
assert http_client.auth is resolved
|
||||
finally:
|
||||
await http_client.aclose()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_without_resolved_auth_falls_back_to_aws_auth(self):
|
||||
aws = httpx.Auth()
|
||||
client = MCPClient(server_url="https://upstream.example.com", aws_auth=aws)
|
||||
http_client = client._create_httpx_client_factory()()
|
||||
try:
|
||||
assert http_client.auth is aws
|
||||
finally:
|
||||
await http_client.aclose()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
|
|
|
|||
|
|
@ -0,0 +1,141 @@
|
|||
"""Tests for the v1 -> v2 bridge.
|
||||
|
||||
`to_server_spec` maps the migrated modes (none + the static-header family, shared-key) and
|
||||
defers everything else to v1 by returning None; `to_subject` maps the principal; `raise_public`
|
||||
maps each CredError onto its HTTP status. These pin the parity-critical mapping before the graft.
|
||||
"""
|
||||
|
||||
import base64
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_public,
|
||||
to_server_spec,
|
||||
to_subject,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
ApiKeyConfig,
|
||||
CredError,
|
||||
NoneConfig,
|
||||
SharedKey,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
def _server(**kwargs) -> MCPServer:
|
||||
return MCPServer(server_id="s", name="n", transport=MCPTransport.http, **kwargs)
|
||||
|
||||
|
||||
def test_none_maps_to_none_config():
|
||||
spec = to_server_spec(_server(auth_type=None))
|
||||
assert spec is not None
|
||||
assert isinstance(spec.config, NoneConfig)
|
||||
|
||||
|
||||
def test_api_key_maps_to_x_api_key_shared():
|
||||
spec = to_server_spec(_server(auth_type=MCPAuth.api_key, authentication_token="k"))
|
||||
assert spec is not None and isinstance(spec.config, ApiKeyConfig)
|
||||
assert spec.config.header_name == "X-API-Key"
|
||||
assert spec.config.value_prefix == ""
|
||||
assert isinstance(spec.config.key_source, SharedKey)
|
||||
assert spec.config.key_source.value.get_secret_value() == "k"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"auth_type, prefix",
|
||||
[
|
||||
(MCPAuth.bearer_token, "Bearer"),
|
||||
(MCPAuth.token, "token"),
|
||||
(MCPAuth.authorization, ""),
|
||||
],
|
||||
)
|
||||
def test_authorization_schemes_map_with_their_prefix(auth_type, prefix):
|
||||
spec = to_server_spec(_server(auth_type=auth_type, authentication_token="t"))
|
||||
assert spec is not None and isinstance(spec.config, ApiKeyConfig)
|
||||
assert spec.config.header_name == "Authorization"
|
||||
assert spec.config.value_prefix == prefix
|
||||
assert spec.config.key_source.value.get_secret_value() == "t"
|
||||
|
||||
|
||||
def test_basic_scheme_base64_encodes_the_token():
|
||||
spec = to_server_spec(
|
||||
_server(auth_type=MCPAuth.basic, authentication_token="user:pass")
|
||||
)
|
||||
assert spec is not None and isinstance(spec.config, ApiKeyConfig)
|
||||
assert spec.config.value_prefix == "Basic"
|
||||
expected = base64.b64encode(b"user:pass").decode()
|
||||
assert spec.config.key_source.value.get_secret_value() == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"server",
|
||||
[
|
||||
_server(auth_type=MCPAuth.api_key), # no token configured
|
||||
_server(auth_type=MCPAuth.bearer_token), # no token configured
|
||||
_server(auth_type=MCPAuth.oauth2),
|
||||
_server(auth_type=MCPAuth.oauth2_token_exchange),
|
||||
_server(auth_type=MCPAuth.aws_sigv4),
|
||||
_server(
|
||||
auth_type=None, oauth_passthrough=True, extra_headers=["Authorization"]
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_unmigrated_modes_defer_to_v1(server):
|
||||
# A None spec is the defer signal; the caller falls back to v1.
|
||||
assert to_server_spec(server) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"server",
|
||||
[
|
||||
_server(auth_type=MCPAuth.api_key, is_byok=True),
|
||||
# BYOK rides on auth_type, so it must defer for every scheme, not just api_key. A stray
|
||||
# static token must not route a BYOK server to a v2 shared-key spec with the wrong value.
|
||||
_server(auth_type=MCPAuth.bearer_token, is_byok=True, authentication_token="x"),
|
||||
_server(auth_type=MCPAuth.basic, is_byok=True, authentication_token="x"),
|
||||
_server(
|
||||
auth_type=MCPAuth.authorization, is_byok=True, authentication_token="x"
|
||||
),
|
||||
_server(auth_type=MCPAuth.token, is_byok=True, authentication_token="x"),
|
||||
_server(auth_type=None, is_byok=True),
|
||||
],
|
||||
)
|
||||
def test_byok_defers_regardless_of_auth_type(server):
|
||||
assert to_server_spec(server) is None
|
||||
|
||||
|
||||
def test_to_subject_unauthenticated_is_empty_with_inbound_token():
|
||||
subject = to_subject(None, "inbound-jwt")
|
||||
assert subject.tenant_id == ""
|
||||
assert subject.subject_id == ""
|
||||
assert subject.inbound_token is not None
|
||||
assert subject.inbound_token.get_secret_value() == "inbound-jwt"
|
||||
|
||||
|
||||
def test_to_subject_maps_principal_fields():
|
||||
principal = SimpleNamespace(org_id="org1", team_id="team1", user_id="user1")
|
||||
subject = to_subject(principal, None)
|
||||
assert subject.tenant_id == "org1"
|
||||
assert subject.subject_id == "user1"
|
||||
assert subject.inbound_token is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"error, status",
|
||||
[
|
||||
(CredError.of_unauthorized("x"), 401),
|
||||
(CredError.of_misconfigured("x"), 500),
|
||||
(CredError.of_upstream_unavailable("x"), 503),
|
||||
(CredError.of_unsupported_mode("x"), 500),
|
||||
(CredError.of_precondition_required("x"), 412),
|
||||
(CredError.of_not_implemented("x"), 501),
|
||||
],
|
||||
)
|
||||
def test_raise_public_maps_each_error_to_its_status(error, status):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
raise_public(error)
|
||||
assert exc_info.value.status_code == status
|
||||
|
|
@ -1,57 +1,104 @@
|
|||
"""Tests for the resolver dispatch skeleton.
|
||||
"""Tests for the resolver dispatch: live arms produce auth, stubbed arms fail closed.
|
||||
|
||||
Every mode must reach its own arm and, until that arm is built, return a typed
|
||||
`not_implemented` CredError rather than silently producing no credential. Parametrizing over
|
||||
one config per mode also guards reachability: if a `case` were dropped, that mode would fall to
|
||||
the `assert_never` tail and raise here instead of returning the stub.
|
||||
`none` and `api_key` (shared-key source) are implemented; every other arm, plus the `api_key`
|
||||
BYOK source, returns a typed `not_implemented` error until its mode lands. Parametrizing the
|
||||
stubs over one config each also guards reachability: a dropped `case` would hit `assert_never`
|
||||
and raise instead of returning the stub.
|
||||
"""
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from pydantic import SecretStr
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
|
||||
ApiKeyConfig,
|
||||
AuthorizationCodeConfig,
|
||||
AuthSpecKind,
|
||||
AwsSigV4Config,
|
||||
Byok,
|
||||
ClientCredentialsConfig,
|
||||
Error,
|
||||
NoneConfig,
|
||||
NoOpAuth,
|
||||
Ok,
|
||||
PassthroughConfig,
|
||||
ServerSpec,
|
||||
SharedKey,
|
||||
StaticHeaderAuth,
|
||||
Subject,
|
||||
TokenExchangeConfig,
|
||||
UpstreamCredentialProvider,
|
||||
)
|
||||
|
||||
_ONE_CONFIG_PER_MODE = [
|
||||
(AuthSpecKind.none, NoneConfig()),
|
||||
(AuthSpecKind.api_key, ApiKeyConfig(key_source=SharedKey(value=SecretStr("k")))),
|
||||
(AuthSpecKind.passthrough, PassthroughConfig()),
|
||||
(AuthSpecKind.client_credentials, ClientCredentialsConfig()),
|
||||
(AuthSpecKind.token_exchange, TokenExchangeConfig()),
|
||||
(AuthSpecKind.authorization_code, AuthorizationCodeConfig()),
|
||||
(AuthSpecKind.aws_sigv4, AwsSigV4Config(region="us-east-1")),
|
||||
_SUBJECT = Subject(tenant_id="", subject_id="")
|
||||
|
||||
|
||||
def _spec(config):
|
||||
return ServerSpec(
|
||||
server_id="s", resource="https://upstream.example.com", config=config
|
||||
)
|
||||
|
||||
|
||||
def _emitted(auth: httpx.Auth) -> httpx.Headers:
|
||||
request = httpx.Request("GET", "https://upstream.example.com/mcp")
|
||||
flow = auth.auth_flow(request)
|
||||
next(flow)
|
||||
flow.close()
|
||||
return request.headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_none_mode_yields_a_no_op_auth():
|
||||
result = await UpstreamCredentialProvider().resolve_credentials(
|
||||
_SUBJECT, _spec(NoneConfig())
|
||||
)
|
||||
assert isinstance(result, Ok)
|
||||
assert isinstance(result.ok, NoOpAuth)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_key_shared_emits_the_configured_header():
|
||||
config = ApiKeyConfig(
|
||||
header_name="X-API-Key",
|
||||
value_prefix="",
|
||||
key_source=SharedKey(value=SecretStr("secret-key")),
|
||||
)
|
||||
result = await UpstreamCredentialProvider().resolve_credentials(
|
||||
_SUBJECT, _spec(config)
|
||||
)
|
||||
assert isinstance(result, Ok)
|
||||
assert isinstance(result.ok, StaticHeaderAuth)
|
||||
assert _emitted(result.ok)["X-API-Key"] == "secret-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_key_shared_honors_authorization_scheme():
|
||||
config = ApiKeyConfig(
|
||||
header_name="Authorization",
|
||||
value_prefix="Bearer",
|
||||
key_source=SharedKey(value=SecretStr("tok")),
|
||||
)
|
||||
result = await UpstreamCredentialProvider().resolve_credentials(
|
||||
_SUBJECT, _spec(config)
|
||||
)
|
||||
assert isinstance(result, Ok)
|
||||
assert _emitted(result.ok)["Authorization"] == "Bearer tok"
|
||||
|
||||
|
||||
_STUBBED = [
|
||||
("api_key_byok", ApiKeyConfig(key_source=Byok())),
|
||||
("passthrough", PassthroughConfig()),
|
||||
("client_credentials", ClientCredentialsConfig()),
|
||||
("token_exchange", TokenExchangeConfig()),
|
||||
("authorization_code", AuthorizationCodeConfig()),
|
||||
("aws_sigv4", AwsSigV4Config(region="us-east-1")),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("kind, config", _ONE_CONFIG_PER_MODE)
|
||||
async def test_every_mode_reaches_its_arm_and_returns_not_implemented(kind, config):
|
||||
spec = ServerSpec(
|
||||
server_id="s", resource="https://upstream.example.com", config=config
|
||||
@pytest.mark.parametrize("label, config", _STUBBED)
|
||||
async def test_unbuilt_arms_fail_closed_with_not_implemented(label, config):
|
||||
result = await UpstreamCredentialProvider().resolve_credentials(
|
||||
_SUBJECT, _spec(config)
|
||||
)
|
||||
subject = Subject(tenant_id="", subject_id="")
|
||||
|
||||
result = await UpstreamCredentialProvider().resolve_credentials(subject, spec)
|
||||
|
||||
assert isinstance(result, Error)
|
||||
assert result.error.tag == "not_implemented"
|
||||
assert kind.value in result.error.summary
|
||||
|
||||
|
||||
def test_all_seven_modes_are_covered():
|
||||
# Guards that the parametrization (and therefore the dispatch) spans every AuthSpecKind, so a
|
||||
# newly added mode without a test row is caught here rather than slipping through.
|
||||
assert {kind for kind, _ in _ONE_CONFIG_PER_MODE} == set(AuthSpecKind)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue