feat(mcp): graft v2 resolver onto _create_mcp_client (none + api_key static family) (#31058)

* 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.

* feat(mcp): graft v2 resolver onto _create_mcp_client for migrated modes

Wire the none + api_key static-family resolver arms from PR4a onto v1's
live request path. In _create_mcp_client's HTTP/SSE branch, to_server_spec
decides per mode: a migrated mode resolves through the injected
UpstreamCredentialProvider and feeds the resulting httpx.Auth into the new
resolved_auth slot; every other mode returns None and falls through to the
unchanged v1 construction. resolve_mcp_auth now runs only when the mode
defers, so a migrated server skips the v1 token-exchange / M2M I/O.

stdio is untouched: auth_type/auth_value never reach the upstream on the
stdio path (_get_auth_headers is HTTP/SSE only), so there is nothing to
graft there. No v1 code is deleted yet; resolve_mcp_auth's static return
still backs stdio and the not-yet-migrated modes until later PRs retire it.

* test(mcp): cover the v2-resolver graft in _create_mcp_client

Regression tests for the PR4 graft. Migrated HTTP modes resolve through the
provider into resolved_auth: none -> NoOpAuth, and the static api_key family
emits the right header per scheme (X-API-Key, Bearer, token, raw authorization,
base64 basic). Deferred modes (oauth2) and a missing static token fall back to
v1's auth_value. A stdio server with a migrated auth_type still defers to v1,
since httpx.Auth never reaches the subprocess. A resolver Error is mapped to the
public HTTP contract (401) via an injected provider, exercising the DI seam.

* fix(mcp): defer to v1 when an inbound credential would be overridden

The graft attaches the resolved static credential as an httpx.Auth, whose auth
flow writes its header after extra_headers. That silently overrode an inbound
Authorization: a per-request mcp_auth_header override, or a header supplied via a
guardrail hook / static_headers / forwarded caller header. v1 lets those win, so
the graft had inverted the credential precedence for the migrated static modes.

Mirror the v2 egress credential-isolation invariant: defer the request to v1 when
mcp_auth_header is set, or when the header the resolved credential would write is
already present in extra_headers. none writes no header, so it never defers.

* test(mcp): cover the credential-isolation defer guard

Regression tests for the precedence fix. A per-request mcp_auth_header override and an
Authorization already present in extra_headers (guardrail hook like the JWT signer,
static_headers, or a forwarded caller header) both defer a migrated static server to v1
so the inbound credential wins; none stays on v2 and does not clobber an inbound
Authorization since NoOpAuth writes nothing. The deferred cases assert resolved_auth is
None, which fails if the guard is removed.

* refactor(mcp): resolve inbound-header conflict on v2 instead of deferring

For an Authorization already supplied via extra_headers (a guardrail hook such as the
JWT signer, static_headers, or a forwarded caller header), keep the request on the v2
path and skip resolved_auth rather than deferring to v1. The inbound header still wins
since nothing overwrites it, but hooks no longer pin a v1 fallback, which is what lets
resolve_mcp_auth be retired once the remaining modes migrate.

The mcp_auth_header per-request override still defers to v1, since that value becomes
the upstream credential rather than sitting in extra_headers; that defer falls away
once the per-user modes stop writing mcp_auth_header.

* fix(mcp): clear UP037 lint gate and fix allowed-servers test under the graft

adapter.py uses `from __future__ import annotations`, so the quoted "UserAPIKeyAuth" /
"MCPServer" annotations in to_subject/to_server_spec/_shared_key_spec were unnecessary
and pushed UP037 over the strict-rule budget; drop the quotes.

test_list_tools_only_returns_allowed_servers passed a MagicMock as user_api_key_auth.
The graft now builds a Subject from the principal, and the MagicMock's non-string
org_id/user_id fail Subject validation, so the listing came back empty. Use a real
UserAPIKeyAuth instead (MagicMock for an injected dependency was the anti-pattern here).

* test(mcp): assert config token via resolved_auth, not the headers dict

test_mcp_server_config_auth_value_header_used inspected _get_auth_headers(), but the
graft now carries the static credential on the client's httpx.Auth (resolved_auth) and
writes the header at send time, so that dict is empty. Assert the header the
StaticHeaderAuth emits onto the request instead. Both config keys (authentication_token,
auth_value) stay covered.

* chore(typecheck): set reportMatchNotExhaustive slack to 0

The previous slack of 3 put the ceiling at baseline + slack = 4, so a newly
non-exhaustive match (for instance dropping an Error arm off a Result match)
could land without tripping the gate. Setting slack to 0 pins the ceiling at
the current baseline of 1, so any added non-exhaustive match now fails CI while
the one pre-existing violation in router.py stays within budget
This commit is contained in:
tin-berri 2026-06-24 14:53:33 -07:00 • committed by GitHub
parent 6003187165
commit bbef1b84ab
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 688 additions and 52 deletions

View file

@ -69,7 +69,7 @@
},
"reportMatchNotExhaustive": {
"baseline": 1,
"slack": 3
"slack": 0
},
"reportMissingParameterType": {
"baseline": 3933,

View file

@ -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,

View file

@ -56,6 +56,16 @@ from litellm.proxy._experimental.mcp_server.sampling_handler import (
MCP_SAMPLING_AVAILABLE,
)
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
Error,
Ok,
UpstreamCredentialProvider,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
raise_public,
to_server_spec,
to_subject,
)
from litellm.proxy._experimental.mcp_server.utils import (
MCP_TOOL_PREFIX_SEPARATOR,
MCPMissingUserEnvVarsError,
@ -511,7 +521,8 @@ class MCPServerManager:
return "client_credentials"
return None
def __init__(self):
def __init__(self, cred_provider: Optional[UpstreamCredentialProvider] = None):
self._cred_provider = cred_provider or UpstreamCredentialProvider()
self.registry: Dict[str, MCPServer] = {}
self.config_mcp_servers: Dict[str, MCPServer] = {}
"""
@ -1942,11 +1953,19 @@ class MCPServerManager:
Returns:
Configured MCP client instance.
"""
auth_value = await resolve_mcp_auth(
server, mcp_auth_header, subject_token=subject_token
)
transport = server.transport or MCPTransport.sse
spec = None if transport == MCPTransport.stdio else to_server_spec(server)
# A per-request override is the caller-supplied credential v1 turns into the upstream
# auth, so it must win; defer those to v1 (this defer falls away once the per-user modes
# stop writing mcp_auth_header). An inbound header already in extra_headers is handled on
# the v2 path below, not here.
if spec is not None and mcp_auth_header:
spec = None
auth_value = (
await resolve_mcp_auth(server, mcp_auth_header, subject_token=subject_token)
if spec is None
else None
)
# Create sampling and elicitation callbacks for this client
sampling_cb = (
@ -2017,6 +2036,43 @@ class MCPServerManager:
# For HTTP/SSE transports
server_url = server.url or ""
if spec is not None:
match await self._cred_provider.resolve_credentials(
to_subject(user_api_key_auth, subject_token), spec
):
case Ok(auth):
resolved_auth = auth
# Do not override an Authorization already supplied via extra_headers
# (a guardrail hook such as the JWT signer, static_headers, or a
# forwarded caller header): v1 applies those last, so they win. NoOpAuth
# has no header_name and so never skips.
header_name = getattr(resolved_auth, "header_name", None)
if (
header_name
and extra_headers
and any(
key.lower() == header_name.lower()
for key in extra_headers
)
):
resolved_auth = None
case Error(err):
raise_public(err)
return MCPClient(
server_url=server_url,
transport_type=transport,
auth_type=server.auth_type,
timeout=(
server.timeout
if server.timeout is not None
else MCP_CLIENT_TIMEOUT
),
extra_headers=extra_headers,
resolved_auth=resolved_auth,
sampling_callback=sampling_cb,
elicitation_callback=elicitation_cb,
)
# Create SigV4 auth if configured
aws_auth = None
if server.auth_type == MCPAuth.aws_sigv4:

View file

@ -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)

View file

@ -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(

View file

@ -45,7 +45,18 @@ async def test_mcp_server_works_without_config_auth_value():
@pytest.mark.parametrize("token_key", ["authentication_token", "auth_value"])
async def test_mcp_server_config_auth_value_header_used(token_key):
"""Ensure auth header is sent when auth token configured in config"""
"""Ensure the configured auth token is emitted as the upstream Authorization header.
The token is resolved through the v2 credential resolver and rides on the client's
httpx.Auth, so assert the header it writes onto the request rather than the (now
credential-free) _get_auth_headers() dict.
"""
import httpx
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
config = {
"test_server": {
"url": "https://api.example.com/mcp",
@ -60,7 +71,8 @@ async def test_mcp_server_config_auth_value_header_used(token_key):
server = next(iter(manager.config_mcp_servers.values()))
client = await manager._create_mcp_client(server)
headers = client._get_auth_headers()
assert headers["Authorization"] == "Bearer example_token"
assert isinstance(client._resolved_auth, StaticHeaderAuth)
emitted = next(client._resolved_auth.auth_flow(httpx.Request("POST", server.url)))
assert emitted.headers["Authorization"] == "Bearer example_token"
assert client.auth_type == MCPAuth.bearer_token

View file

@ -1089,7 +1089,9 @@ async def test_list_tools_only_returns_allowed_servers(monkeypatch):
mock_client_constructor,
):
# Call list_tools
tools = await test_manager.list_tools(user_api_key_auth=MagicMock())
from litellm.proxy._types import UserAPIKeyAuth
tools = await test_manager.list_tools(user_api_key_auth=UserAPIKeyAuth())
# Should only return tools from server_a
assert len(tools) == 1
# The server should use the server_name as prefix since no alias is provided

View file

@ -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__])

View 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

View file

@ -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)

View file

@ -4568,5 +4568,173 @@ class TestGetPublicMCPServersLegacyMode:
assert sorted(s.server_id for s in result) == ["a", "b"]
class TestCreateMcpClientV2Graft:
"""The PR4 v2-resolver graft in ``_create_mcp_client``.
Migrated HTTP/SSE modes (``none`` plus the static ``api_key`` family) resolve through the
injected ``UpstreamCredentialProvider`` into the ``resolved_auth`` slot; every other mode,
and every stdio server, defers to v1's ``auth_value`` path unchanged.
"""
def _http_server(self, **overrides: Any) -> MCPServer:
base: Dict[str, Any] = dict(
server_id="http-graft",
name="graft_server",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
)
base.update(overrides)
return MCPServer(**base)
async def test_none_mode_resolves_to_noop_auth(self):
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
NoOpAuth,
)
client = await MCPServerManager()._create_mcp_client(
self._http_server(auth_type=None)
)
assert isinstance(client._resolved_auth, NoOpAuth)
assert client._mcp_auth_value is None
@pytest.mark.parametrize(
"auth_type, token, expected_name, expected_value",
[
(MCPAuth.api_key, "k-123", "X-API-Key", "k-123"),
(MCPAuth.bearer_token, "b-123", "Authorization", "Bearer b-123"),
(MCPAuth.token, "t-123", "Authorization", "token t-123"),
(MCPAuth.authorization, "raw-123", "Authorization", "raw-123"),
],
)
async def test_static_family_emits_expected_header(
self, auth_type, token, expected_name, expected_value
):
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
client = await MCPServerManager()._create_mcp_client(
self._http_server(auth_type=auth_type, authentication_token=token)
)
assert isinstance(client._resolved_auth, StaticHeaderAuth)
assert client._resolved_auth.header_name == expected_name
assert client._resolved_auth._header_value.get_secret_value() == expected_value
assert client._mcp_auth_value is None
async def test_basic_mode_base64_encodes(self):
import base64
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
client = await MCPServerManager()._create_mcp_client(
self._http_server(auth_type=MCPAuth.basic, authentication_token="user:pass")
)
encoded = base64.b64encode(b"user:pass").decode()
assert isinstance(client._resolved_auth, StaticHeaderAuth)
assert client._resolved_auth.header_name == "Authorization"
assert (
client._resolved_auth._header_value.get_secret_value() == f"Basic {encoded}"
)
async def test_deferred_mode_uses_v1_auth_value(self):
client = await MCPServerManager()._create_mcp_client(
self._http_server(
auth_type=MCPAuth.oauth2, authentication_token="legacy-token"
)
)
assert client._resolved_auth is None
assert client._mcp_auth_value == "legacy-token"
async def test_static_token_missing_defers_to_v1(self):
client = await MCPServerManager()._create_mcp_client(
self._http_server(auth_type=MCPAuth.api_key, authentication_token=None)
)
assert client._resolved_auth is None
async def test_stdio_migrated_auth_type_still_defers_to_v1(self):
client = await MCPServerManager()._create_mcp_client(
MCPServer(
server_id="stdio-graft",
name="stdio_graft",
transport=MCPTransport.stdio,
command="node",
args=["server.js"],
auth_type=MCPAuth.api_key,
authentication_token="k-stdio",
)
)
assert client.transport_type == MCPTransport.stdio
assert client._resolved_auth is None
assert client._mcp_auth_value == "k-stdio"
async def test_resolver_error_maps_to_http_exception(self):
from litellm.proxy._experimental.mcp_server.outbound_credentials import Error
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
CredError,
)
class _UnauthorizedProvider:
async def resolve_credentials(self, subject, server):
return Error(CredError.of_unauthorized("denied"))
manager = MCPServerManager(cred_provider=_UnauthorizedProvider())
with pytest.raises(HTTPException) as exc:
await manager._create_mcp_client(self._http_server(auth_type=None))
assert exc.value.status_code == 401
async def test_per_request_override_defers_to_v1(self):
# A per-request override (mcp_auth_header) must win over the shared static token,
# exactly as v1 did, so a migrated static server defers to v1 when one is present.
client = await MCPServerManager()._create_mcp_client(
self._http_server(
auth_type=MCPAuth.bearer_token, authentication_token="shared-tok"
),
mcp_auth_header="caller-override",
)
assert client._resolved_auth is None
assert client._mcp_auth_value == "caller-override"
async def test_conflicting_extra_header_skips_resolved_auth_on_v2(self):
# An Authorization already supplied via extra_headers (guardrail hook like the JWT
# signer, static_headers, or a forwarded caller header) must win. The server stays on
# the v2 path but skips resolved_auth, so nothing overwrites the inbound header.
client = await MCPServerManager()._create_mcp_client(
self._http_server(
auth_type=MCPAuth.bearer_token, authentication_token="shared-tok"
),
extra_headers={"Authorization": "Bearer hook-jwt"},
)
assert client._resolved_auth is None
assert client._mcp_auth_value is None
assert client._get_auth_headers()["Authorization"] == "Bearer hook-jwt"
async def test_none_with_extra_header_stays_v2_without_clobbering(self):
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
NoOpAuth,
)
# none resolves to NoOpAuth, which writes no header, so it cannot clobber an inbound
# Authorization; it stays on the v2 path and the inbound header is preserved verbatim.
client = await MCPServerManager()._create_mcp_client(
self._http_server(auth_type=None),
extra_headers={"Authorization": "Bearer hook-jwt"},
)
assert isinstance(client._resolved_auth, NoOpAuth)
assert client._get_auth_headers()["Authorization"] == "Bearer hook-jwt"
if __name__ == "__main__":
pytest.main([__file__])