fix(mcp): bind keyed OAuth connections to the initiating key

This commit is contained in:
Joshua Valluru 2026-09-21 11:56:01 -07:00
parent cc1a3157d3
commit 641f50afe2
11 changed files with 1723 additions and 13 deletions

View file

@ -13,6 +13,12 @@ from typing_extensions import assert_never
import litellm
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import (
CONNECTION_SCOPE_KEY,
connection_challenge,
is_connection_credential,
open_connection_credential,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
get_passthrough_resource_metadata_url,
get_passthrough_www_authenticate,
@ -28,6 +34,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credenti
resolve_bridge_envelope,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
ConnectionBinding,
EnvelopeIdentity,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import (
@ -42,6 +49,7 @@ from litellm.proxy._types import (
SpecialMCPServerName,
SpecialMCPServerNames,
UserAPIKeyAuth,
hash_token,
user_api_key_has_admin_view,
)
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
@ -541,6 +549,52 @@ class MCPRequestHandler:
bearer_presented=False,
)
scope.pop(CONNECTION_SCOPE_KEY, None)
connection_header: Final = headers.get("authorization")
if is_connection_credential(connection_header):
if not has_explicit_litellm_key or request_route != "/mcp":
raise HTTPException(status_code=401, detail="A connection credential requires the original MCP key")
targets: Final = MCPRequestHandler._resolve_target_server_names(request_route, mcp_servers)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
target: Final = (
global_mcp_server_manager.get_mcp_server_by_name(
targets[0], client_ip=IPAddressUtils.get_mcp_client_ip(request)
)
if len(targets) == 1
else None
)
allowed: Final = await MCPRequestHandler.get_allowed_mcp_servers(validated_user_api_key_auth)
if (
target is None
or target.server_id not in allowed
or not target.needs_user_oauth_token
or target.oauth_identity_binding is not None
):
raise HTTPException(status_code=403, detail="Connection credential does not authorize this MCP server")
expected_binding: Final = ConnectionBinding(
key_hash=hash_token(_get_bearer_token_or_received_api_key(litellm_api_key)),
server_id=target.server_id,
resource=f"{get_request_base_url(request)}/mcp",
)
connection: Final = open_connection_credential(connection_header or "")
if connection is None:
raise HTTPException(
status_code=401,
detail="Invalid or expired MCP connection credential",
headers=MappingProxyType(
{
"www-authenticate": connection_challenge(request, expected_binding),
"Cache-Control": "no-store",
}
),
)
if connection.binding != expected_binding:
raise HTTPException(
status_code=401, detail="Connection credential belongs to a different key or resource"
)
scope[CONNECTION_SCOPE_KEY] = connection
# Leak-defense (single chokepoint): a gateway admission credential (session bearer or bridge
# envelope) is NEVER a valid upstream token. Scrub it from EVERY egress context so no
# client-forwarded, OBO, or passthrough path can send it upstream for replay. Anchored to the
@ -573,7 +627,9 @@ class MCPRequestHandler:
"""True when a header value is a gateway admission credential — a session bearer or bridge
envelope. It proves who signed in to the GATEWAY, never a valid UPSTREAM token, so it must never
be forwarded (a hostile upstream could capture and replay it against the aggregate ``/mcp`` scope)."""
return value is not None and (is_session_bearer_shaped(value) or is_bridge_envelope_shaped(value))
return value is not None and (
is_session_bearer_shaped(value) or is_bridge_envelope_shaped(value) or is_connection_credential(value)
)
@staticmethod
def _scrub_gateway_admission_credentials(

View file

@ -5,6 +5,7 @@ import secrets
import time
from collections.abc import Callable, Mapping
from datetime import datetime, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
@ -46,20 +47,38 @@ from litellm.proxy._experimental.mcp_server.faults import (
render_token_fault,
)
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import (
CONNECTION_FLOW_PREFIX,
VendorCredentialState,
_cookie_path_and_secure, # pyright: ignore[reportPrivateUsage] # reuse the gateway cookie scope policy
_oauth_error, # pyright: ignore[reportPrivateUsage] # reuse the gateway OAuth error contract
_pkce_verifier_matches, # pyright: ignore[reportPrivateUsage] # reuse the gateway S256 verification
_seal, # pyright: ignore[reportPrivateUsage] # reuse authenticated gateway flow sealing
aggregate_authorize,
aggregate_token,
authorize_connection,
claim_connection_once,
complete_connect_flow,
describe_connect_flow,
introspect_gateway_token,
is_gateway_dcr_client_id,
is_proxy_api_resource,
mint_connection_tokens,
native_client_auth_contract,
native_client_authorize,
open_connection_bootstrap,
open_connection_code,
open_connection_credential,
open_connection_flow,
open_gateway_dcr_client,
register_aggregate_client,
relative_request_url,
revoke_refresh_token,
seal_connection_code,
supported_grant_types,
validate_connection_binding,
)
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import (
_append_query_params as append_connection_query, # pyright: ignore[reportPrivateUsage] # reuse gateway redirect encoding
)
from litellm.proxy._experimental.mcp_server.idp_token_exchange import (
exchange_idp_subject_token,
@ -78,6 +97,10 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import (
validate_trusted_redirect_uri,
well_known_root_suffix,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
ConnectionAuthorization,
ConnectionBinding,
)
from litellm.proxy._experimental.mcp_server.proxy_api_credentials import (
lookup_consent_teams,
mint_proxy_credential,
@ -153,6 +176,7 @@ def encode_state_with_base_url(
dcr_client_secret: str | None = None,
dcr_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None,
oauth_nonce: str | None = None,
connection_flow: str | None = None,
) -> str:
"""
Encode the base_url, original state, and PKCE parameters using encryption.
@ -182,6 +206,7 @@ def encode_state_with_base_url(
An encrypted string that encodes all values
"""
state_data: Final = {
"connection_flow": connection_flow,
"oauth_nonce": oauth_nonce,
"base_url": base_url,
"original_state": original_state,
@ -881,6 +906,7 @@ async def authorize_with_server(
response_type: str | None = None,
scope: str | None = None,
ephemeral_dcr_client: "EphemeralDcrClient | None" = None,
connection: ConnectionAuthorization | None = None,
):
_raise_if_not_oauth2(mcp_server)
resolved_server: Final = await _server_with_oauth_endpoints(mcp_server, _register_flow_needed_endpoint)
@ -944,6 +970,7 @@ async def authorize_with_server(
encoded_state: Final = encode_state_with_base_url(
base_url=base_url,
original_state=state,
connection_flow=_seal(CONNECTION_FLOW_PREFIX, connection) if connection is not None else None,
oauth_nonce=oauth_nonce,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
@ -989,6 +1016,8 @@ async def authorize_with_server(
final_url: Final = urlunparse(parsed_auth_url._replace(query=urlencode(existing_params)))
response: Final = RedirectResponse(final_url)
_set_oauth_state_cookie(response, request, relay_state, encoded_state)
if connection is not None and len(response.headers["set-cookie"]) > 4096:
return _oauth_error(400, "invalid_request", "Connection authorization metadata is too large")
return response
@ -1011,6 +1040,7 @@ async def exchange_token_with_server(
refresh_token: str | None = None,
scope: str | None = None,
client_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None,
connection_binding: ConnectionBinding | None = None,
):
_raise_if_not_oauth2(mcp_server)
if grant_type not in ("authorization_code", "refresh_token"):
@ -1222,6 +1252,11 @@ async def exchange_token_with_server(
else None
)
if connection_binding is not None:
return mint_connection_tokens(
connection_binding, client_id, token_response, fallback_refresh=refresh_token, fallback_scope=scope
)
# Store server-side when the server is configured for per-user OAuth and
# the calling client has provided a valid LiteLLM identity.
# Errors are non-fatal: the token is still returned to the client.
@ -1912,6 +1947,21 @@ async def authorize(
scope: str | None = None,
resource: str | None = None,
):
registered: Final = open_gateway_dcr_client(client_id) if client_id else None
if registered is not None and registered.connection is not None:
if mcp_server_name is not None:
return _oauth_error(400, "invalid_client", "Use the registered connection authorization endpoint")
return await authorize_connection(
request,
client_id or "",
redirect_uri,
state,
code_challenge,
code_challenge_method,
response_type,
resource,
scope,
)
# Redirect to real OAuth provider with PKCE support
if mcp_server_name is None and client_id and is_gateway_dcr_client_id(client_id):
if is_proxy_api_resource(request, resource):
@ -1999,6 +2049,22 @@ async def token_endpoint(
3. Return the token
4. Return a virtual key in this response
"""
registered: Final = open_gateway_dcr_client(client_id)
if registered is not None and registered.connection is not None:
if mcp_server_name is not None:
return _oauth_error(400, "invalid_client", "Use the registered connection token endpoint")
return await exchange_connection_token(
request,
registered.connection,
grant_type,
client_id,
code,
redirect_uri,
code_verifier,
refresh_token,
resource,
scope,
)
if mcp_server_name is None and is_gateway_dcr_client_id(client_id):
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # circular import at module load
master_key,
@ -2271,6 +2337,24 @@ async def callback(
mcp_server_id: Final = state_data.get("mcp_server_id")
dcr_client_id: Final = state_data.get("dcr_client_id")
dcr_client_secret: Final = state_data.get("dcr_client_secret")
encoded_connection: Final = state_data.get("connection_flow")
connection: Final = open_connection_flow(encoded_connection) if isinstance(encoded_connection, str) else None
if encoded_connection is not None:
if connection is None or connection.redirect_uri != redirect_uri:
raise HTTPException(status_code=400, detail="Invalid or expired connection authorization")
replay: Final = await claim_connection_once(f"callback:{connection.jti}", connection.exp)
if replay is not None:
return replay
destination: Final = append_connection_query(
redirect_uri,
(
("code", seal_connection_code(connection, code)),
("state", connection.state),
),
)
connection_response: Final = RedirectResponse(destination, status_code=302)
_clear_oauth_state_cookie(connection_response, request, state)
return connection_response
forwarded_code = code
if isinstance(litellm_user_id, str) and litellm_user_id and isinstance(mcp_server_id, str) and mcp_server_id:
forwarded_code = seal_bridge_authorization_code(
@ -2665,7 +2749,7 @@ def _build_aggregate_authorization_server_response(request: Request, token_excha
# in registration order, and /.well-known/oauth-authorization-server/{name}
# would otherwise capture the "/mcp" suffix as a server name.
@router.get(f"/.well-known/oauth-protected-resource{well_known_root_suffix()}/mcp")
async def oauth_protected_resource_aggregate(request: Request):
async def oauth_protected_resource_aggregate(request: Request, connection: str | None = None):
"""
OAuth protected resource discovery for the aggregate /mcp endpoint.
@ -2673,9 +2757,54 @@ async def oauth_protected_resource_aggregate(request: Request):
(those are two-segment: ``/mcp/{server}`` or ``/{server}/mcp``), so this unambiguously
describes the aggregate resource.
"""
bootstrap: Final = connection
if bootstrap is not None:
binding: Final = open_connection_bootstrap(bootstrap)
if binding is None:
raise HTTPException(status_code=400, detail="Invalid or expired connection discovery")
server: Final = await validate_connection_binding(request, binding)
from mcp.shared.auth import ProtectedResourceMetadata
metadata: Final = ProtectedResourceMetadata.model_validate(
MappingProxyType(
{
"resource": binding.resource,
"authorization_servers": (f"{get_request_base_url(request)}/mcp-connect/{bootstrap}",),
"scopes_supported": tuple(server.scopes or ()),
}
)
)
return JSONResponse(metadata.model_dump(mode="json", exclude_none=True), headers=TOKEN_NO_CACHE_HEADERS)
return _build_aggregate_protected_resource_response(request)
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/mcp-connect/{{bootstrap}}")
async def connection_authorization_metadata(request: Request, bootstrap: str) -> JSONResponse:
binding: Final = open_connection_bootstrap(bootstrap)
if binding is None:
raise HTTPException(status_code=400, detail="Invalid or expired connection discovery")
server: Final = await validate_connection_binding(request, binding)
base: Final = get_request_base_url(request)
from mcp.shared.auth import OAuthMetadata
metadata: Final = OAuthMetadata.model_validate(
MappingProxyType(
{
"issuer": f"{base}/mcp-connect/{bootstrap}",
"authorization_endpoint": f"{base}/authorize",
"token_endpoint": f"{base}/token",
"registration_endpoint": append_connection_query(f"{base}/register", (("connection", bootstrap),)),
"scopes_supported": tuple(server.scopes or ()),
"response_types_supported": ("code",),
"grant_types_supported": ("authorization_code", "refresh_token"),
"code_challenge_methods_supported": ("S256",),
"token_endpoint_auth_methods_supported": ("none",),
}
)
)
return JSONResponse(metadata.model_dump(mode="json", exclude_none=True), headers=TOKEN_NO_CACHE_HEADERS)
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/mcp")
async def oauth_authorization_server_aggregate(request: Request):
"""
@ -2895,7 +3024,15 @@ async def oauth_authorization_server_legacy(request: Request, mcp_server_name: s
@router.post("/{mcp_server_name}/register")
@router.post("/register")
async def register_client(request: Request, mcp_server_name: str | None = None):
async def register_client(request: Request, mcp_server_name: str | None = None, connection: str | None = None):
bootstrap: Final = connection
if bootstrap is not None:
binding: Final = open_connection_bootstrap(bootstrap)
if binding is None or mcp_server_name is not None:
return _oauth_error(400, "invalid_client", "Invalid or expired connection discovery")
await validate_connection_binding(request, binding)
body: Final = await _read_request_body(request=request)
return await register_aggregate_client(request, body, False, connection=binding)
# Get the correct base URL considering X-Forwarded-* headers
request_base_url: Final = get_request_base_url(request)
@ -2946,3 +3083,111 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
fallback_client_id=mcp_server_name,
client_redirect_uris=client_redirect_uris,
)
@router.post("/authorize/connection/complete")
async def complete_connection(request: Request, flow: str = Form(...), decision: str = Form(...)) -> Response:
cookie_name: Final = f"mcp_connection_{flow}"
opened: Final = open_connection_flow(request.cookies.get(cookie_name, ""))
if opened is None or opened.jti != flow or decision not in ("approve", "deny"):
return _oauth_error(400, "invalid_request", "Invalid or expired consent")
server: Final = await validate_connection_binding(request, opened.binding)
claim: Final = await claim_connection_once(f"consent:{opened.jti}", opened.exp)
if claim is not None:
return claim
if decision == "deny":
denied: Final = RedirectResponse(
append_connection_query(
opened.redirect_uri,
(
("error", "access_denied"),
("state", opened.state),
),
),
status_code=302,
)
cookie_path, _ = _cookie_path_and_secure(request)
denied.delete_cookie(cookie_name, path=cookie_path)
return denied
response: Final = await authorize_with_server(
request,
server,
opened.client_id,
opened.redirect_uri,
opened.state,
opened.code_challenge,
"S256",
"code",
opened.scope,
connection=opened,
)
cookie_path, _ = _cookie_path_and_secure(request)
response.delete_cookie(cookie_name, path=cookie_path)
return response
async def exchange_connection_token(
request: Request,
binding: ConnectionBinding,
grant_type: str,
client_id: str,
code: str | None,
redirect_uri: str | None,
code_verifier: str | None,
refresh_token: str | None,
resource: str | None,
scope: str | None,
) -> Response:
if resource is not None and resource != binding.resource:
return _oauth_error(400, "invalid_target", "The requested resource does not match this connection")
server: Final = await validate_connection_binding(request, binding)
if grant_type == "authorization_code":
opened: Final = open_connection_code(code or "")
if (
opened is None
or opened.authorization.binding != binding
or opened.authorization.client_id != client_id
or opened.authorization.redirect_uri != redirect_uri
or not code_verifier
or not 43 <= len(code_verifier) <= 128
or not _pkce_verifier_matches(code_verifier, opened.authorization.code_challenge)
):
return _oauth_error(400, "invalid_grant", "Invalid authorization code or PKCE verifier")
claim: Final = await claim_connection_once(f"code:{opened.jti}", opened.exp)
if claim is not None:
return claim
return await exchange_token_with_server(
request,
server,
grant_type,
opened.upstream_code.get_secret_value(),
redirect_uri,
client_id,
None,
code_verifier,
scope=opened.authorization.scope,
connection_binding=binding,
)
if grant_type == "refresh_token":
refreshed: Final = open_connection_credential(refresh_token or "", refresh=True)
if refreshed is None or refreshed.binding != binding or refreshed.client_id != client_id:
return _oauth_error(400, "invalid_grant", "Invalid refresh credential")
if scope and not frozenset(scope.split()).issubset((refreshed.scope or "").split()):
return _oauth_error(400, "invalid_scope", "Refresh cannot expand the granted scopes")
claimed: Final = await claim_connection_once(f"refresh:{refreshed.jti}", refreshed.exp)
if claimed is not None:
return claimed
return await exchange_token_with_server(
request,
server,
grant_type,
None,
None,
client_id,
None,
None,
refresh_token=refreshed.token.get_secret_value(),
scope=scope or refreshed.scope,
connection_binding=binding,
)
return _oauth_error(400, "unsupported_grant_type", "Unsupported connection grant type")

View file

@ -30,9 +30,9 @@ Nothing here stores state server-side except the single-use code guard (a TTL ca
entry). Every sealed value is authenticated encryption over the proxy salt/master key
family, opened totally (bad input maps to an OAuth error, never a raise), and every
identity is a stable reference re-validated live at mint, refresh, and (in the admission
PR) tool-call time. Upstream server credentials never appear anywhere in this flow; they
are vaulted per user by the existing ``/v1/mcp`` authorize endpoints and resolved at
egress by user id.
PR) tool-call time. The SSO flow vaults upstream credentials per user through the existing
``/v1/mcp`` authorize endpoints. Keyed connections instead seal a server-specific upstream
credential for the client, require the original key at admission, and never write the vault.
"""
from __future__ import annotations
@ -50,7 +50,7 @@ from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
from fastapi import HTTPException, Request
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response
from pydantic import BaseModel, ConfigDict, Field, ValidationError
from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError
from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never
from litellm._logging import verbose_logger
@ -62,6 +62,15 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import (
get_request_base_url,
is_loopback_redirect_host,
validate_redirect_uri_shape,
well_known_root_suffix,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
ConnectionAuthorization,
ConnectionBinding,
ConnectionBootstrap,
ConnectionCode,
ConnectionCredential,
RefreshCredential,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import (
SessionRefreshOpened,
@ -284,6 +293,8 @@ class GatewayDcrClient(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
redirect_uris: tuple[str, ...] = Field(min_length=1, max_length=MAX_REDIRECT_URIS)
iat: int
connection: ConnectionBinding | None = None
client_name: str | None = None
class _ConnectFlow(BaseModel):
@ -371,7 +382,10 @@ def open_gateway_dcr_client(client_id: str) -> GatewayDcrClient | None:
async def register_aggregate_client(
request: Request, request_body: Mapping[str, object], token_exchange_available: bool
request: Request,
request_body: Mapping[str, object],
token_exchange_available: bool,
connection: ConnectionBinding | None = None,
) -> Response:
"""RFC 7591 dynamic registration against the gateway itself, statelessly.
@ -425,7 +439,13 @@ async def register_aggregate_client(
)
now: Final = datetime.now(timezone.utc)
client_id: Final = _seal(
GATEWAY_DCR_CLIENT_ID_PREFIX, GatewayDcrClient(redirect_uris=tuple(raw_uris), iat=int(now.timestamp()))
GATEWAY_DCR_CLIENT_ID_PREFIX,
GatewayDcrClient(
redirect_uris=tuple(raw_uris),
iat=int(now.timestamp()),
connection=connection,
client_name=str(request_body.get("client_name", "MCP client"))[:100] if connection else None,
),
)
if len(client_id) > MAX_CLIENT_ID_LENGTH:
return _oauth_error(400, "invalid_client_metadata", "registered metadata is too large")
@ -533,6 +553,9 @@ def aggregate_authorize(
)
if rejected is not None:
return rejected
registered: Final = open_gateway_dcr_client(client_id)
if registered is not None and registered.connection is not None:
return _oauth_error(400, "invalid_client", "Use the registered connection authorization endpoint")
base_url: Final = get_request_base_url(request)
if session_user_id is None:
return _login_redirect(base_url, request)
@ -1547,3 +1570,246 @@ async def introspect_gateway_token(
if failure is not None:
return _inactive_introspection_response()
return _active_introspection_response(opened)
CONNECTION_BOOTSTRAP_PREFIX: Final = "llm_cboot_"
CONNECTION_FLOW_PREFIX: Final = "llm_cflow_"
CONNECTION_CODE_PREFIX: Final = "llm_ccode_"
CONNECTION_ACCESS_PREFIX: Final = "llm_caccess_"
CONNECTION_REFRESH_PREFIX: Final = "llm_crefresh_"
CONNECTION_SCOPE_KEY: Final = "litellm.mcp.connection_grant"
def mint_connection_bootstrap(binding: ConnectionBinding) -> str:
return _seal(
CONNECTION_BOOTSTRAP_PREFIX,
ConnectionBootstrap(
binding=binding,
exp=int(datetime.now(timezone.utc).timestamp()) + CONNECT_FLOW_TTL_SECONDS,
),
)
def connection_challenge(request: Request, binding: ConnectionBinding) -> str:
bootstrap: Final = mint_connection_bootstrap(binding)
metadata: Final = _append_query_params(
f"{get_request_base_url(request)}/.well-known/oauth-protected-resource{well_known_root_suffix()}/mcp",
(("connection", bootstrap),),
)
return f'Bearer resource_metadata="{metadata}"'
def open_connection_bootstrap(value: str) -> ConnectionBinding | None:
opened: Final = _open_sealed(value, CONNECTION_BOOTSTRAP_PREFIX, ConnectionBootstrap, "connection_bootstrap")
if opened is None or opened.exp <= int(datetime.now(timezone.utc).timestamp()):
return None
return opened.binding
def is_connection_credential(value: str | None) -> bool:
bearer: Final = value[7:] if value and value[:7].lower() == "bearer " else value
return bool(bearer and bearer.startswith((CONNECTION_ACCESS_PREFIX, CONNECTION_REFRESH_PREFIX)))
def open_connection_credential(value: str, *, refresh: bool = False) -> ConnectionCredential | None:
prefix: Final = CONNECTION_REFRESH_PREFIX if refresh else CONNECTION_ACCESS_PREFIX
expected: Final = "connection_refresh" if refresh else "connection_access"
bearer: Final = value[7:] if value[:7].lower() == "bearer " else value
if len(bearer) > 12288:
return None
opened: Final = _open_sealed(bearer, prefix, ConnectionCredential, "connection_credential")
if opened is None or opened.kind != expected or opened.exp <= int(datetime.now(timezone.utc).timestamp()):
return None
return opened
async def validate_connection_binding(request: Request, binding: ConnectionBinding) -> MCPServer:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
if binding.resource != f"{get_request_base_url(request)}/mcp":
raise HTTPException(status_code=400, detail="Invalid connection resource")
key: Final = await MCPRequestHandler._reload_admitted_key(binding.key_hash) # pyright: ignore[reportPrivateUsage] # reuse key revocation and SCIM checks
await MCPRequestHandler._enforce_admitted_live_policy(key.model_copy(), request, "/mcp") # pyright: ignore[reportPrivateUsage] # enforce the same MCP route and budget policy
allowed: Final = await MCPRequestHandler.get_allowed_mcp_servers(key)
server: Final = global_mcp_server_manager.get_mcp_server_by_id(
binding.server_id, client_ip=IPAddressUtils.get_mcp_client_ip(request)
)
if server is None or server.server_id not in allowed:
raise HTTPException(status_code=403, detail="Key is not allowed to access the selected MCP server")
if not server.needs_user_oauth_token or server.oauth_identity_binding is not None:
raise HTTPException(status_code=400, detail="Server does not support a keyed connection grant")
return server
async def claim_connection_once(jti: str, expires: int) -> Response | None:
from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache
if redis_usage_cache is None and getattr(user_api_key_cache, "redis_cache", None) is None:
return _oauth_error(503, "temporarily_unavailable", "Keyed MCP OAuth requires a shared Redis cache")
ttl: Final = max(1, expires - int(datetime.now(timezone.utc).timestamp()) + _CLAIM_TTL_BUFFER_SECONDS)
outcome: Final = await _SingleUseGuard(user_api_key_cache).claim(f"mcp_connection_used:{jti}", ttl)
return _claim_refusal(outcome, _oauth_error(400, "invalid_grant", "This authorization has already been used"))
async def authorize_connection(
request: Request,
client_id: str,
redirect_uri: str,
state: str,
code_challenge: str | None,
code_challenge_method: str | None,
response_type: str | None,
resource: str | None,
scope: str | None,
) -> Response:
from html import escape
rejected: Final = _rejected_authorize_request(
client_id, redirect_uri, state, code_challenge, code_challenge_method, response_type
)
if rejected is not None:
return rejected
client: Final = open_gateway_dcr_client(client_id)
if client is None or client.connection is None or resource != client.connection.resource:
return _oauth_error(400, "invalid_target", "The requested resource does not match this connection")
if (
code_challenge is None
or len(code_challenge) != 43
or not all(c.isascii() and (c.isalnum() or c in "-_") for c in code_challenge)
):
return _oauth_error(400, "invalid_request", "A valid S256 code challenge is required")
server: Final = await validate_connection_binding(request, client.connection)
scopes: Final = scope or " ".join(server.scopes or ())
if not frozenset(scopes.split()).issubset(server.scopes or ()):
return _oauth_error(400, "invalid_scope", "Requested scopes are not configured for this server")
handle: Final = secrets.token_urlsafe(24)
flow: Final = ConnectionAuthorization(
binding=client.connection,
client_id=client_id,
redirect_uri=redirect_uri,
state=state,
code_challenge=code_challenge,
scope=scopes,
jti=handle,
exp=int(datetime.now(timezone.utc).timestamp()) + CONNECT_FLOW_TTL_SECONDS,
)
action: Final = f"{get_request_base_url(request)}/authorize/connection/complete"
response: Final = HTMLResponse(
'<!doctype html><html lang="en"><meta charset="utf-8"><meta name="referrer" content="no-referrer">'
'<meta name="viewport" content="width=device-width, initial-scale=1"><title>Authorize MCP connection</title>'
"<h1>Authorize MCP connection</h1>"
f"<p><strong>{escape(client.client_name or 'MCP client')}</strong> wants access to "
f"<strong>{escape(server.name)}</strong>.</p><p>Return address: <code>{escape(redirect_uri)}</code></p>"
f"<p>Requested permissions: {escape(scopes or 'provider defaults')}</p>"
"<p>Only approve if you started this connection. You will continue to the provider to sign in.</p>"
f'<form method="post" action="{escape(action)}"><input type="hidden" name="flow" value="{escape(handle)}">'
'<button name="decision" value="deny">Deny</button> <button name="decision" value="approve">Continue</button>'
"</form></html>",
headers=MappingProxyType(
{
**TOKEN_NO_CACHE_HEADERS,
"Referrer-Policy": "no-referrer",
"X-Frame-Options": "DENY",
"Content-Security-Policy": "default-src 'none'; form-action 'self'; frame-ancestors 'none'",
}
),
)
cookie_path, secure = _cookie_path_and_secure(request)
response.set_cookie(
f"mcp_connection_{handle}",
_seal(CONNECTION_FLOW_PREFIX, flow),
max_age=CONNECT_FLOW_TTL_SECONDS,
httponly=True,
secure=secure,
samesite="lax",
path=cookie_path,
)
if len(response.headers["set-cookie"]) > 4096:
return _oauth_error(400, "invalid_request", "Connection authorization metadata is too large")
return response
def open_connection_flow(value: str) -> ConnectionAuthorization | None:
flow: Final = _open_sealed(value, CONNECTION_FLOW_PREFIX, ConnectionAuthorization, "connection_flow")
return flow if flow is not None and flow.exp > int(datetime.now(timezone.utc).timestamp()) else None
def seal_connection_code(flow: ConnectionAuthorization, upstream_code: str) -> str:
return _seal(
CONNECTION_CODE_PREFIX,
ConnectionCode(
authorization=flow,
upstream_code=SecretStr(upstream_code),
jti=secrets.token_urlsafe(24),
exp=int(datetime.now(timezone.utc).timestamp()) + GATEWAY_AUTH_CODE_TTL_SECONDS,
),
)
def open_connection_code(value: str) -> ConnectionCode | None:
code: Final = _open_sealed(value, CONNECTION_CODE_PREFIX, ConnectionCode, "connection_code")
return code if code is not None and code.exp > int(datetime.now(timezone.utc).timestamp()) else None
def mint_connection_tokens(
binding: ConnectionBinding,
client_id: str,
token_response: object,
fallback_refresh: str | None = None,
fallback_scope: str | None = None,
) -> Response:
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_bridge_grant_from_token_response, # pyright: ignore[reportPrivateUsage] # reuse provider token validation
_upstream_refresh_credential, # pyright: ignore[reportPrivateUsage] # reuse provider refresh lifetime parsing
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import UpstreamTokenGrant
grant: Final = _bridge_grant_from_token_response(token_response)
if not isinstance(grant, UpstreamTokenGrant) or grant.token_type.lower() != "bearer":
return _oauth_error(502, "server_error", "The upstream did not return a usable bearer token")
granted_scope: Final = grant.scope if grant.scope is not None else fallback_scope
now: Final = int(datetime.now(timezone.utc).timestamp())
ttl: Final = min(grant.expires_in or 3600, 3600)
access: Final = _seal(
CONNECTION_ACCESS_PREFIX,
ConnectionCredential(
kind="connection_access",
binding=binding,
client_id=client_id,
token=grant.access_token,
scope=granted_scope,
jti=secrets.token_urlsafe(24),
exp=now + ttl,
),
)
refresh: Final = _upstream_refresh_credential(token_response) or (
RefreshCredential(refresh_token=SecretStr(fallback_refresh), scope=granted_scope) if fallback_refresh else None
)
sealed_refresh: Final = (
_seal(
CONNECTION_REFRESH_PREFIX,
ConnectionCredential(
kind="connection_refresh",
binding=binding,
client_id=client_id,
token=refresh.refresh_token,
scope=refresh.scope if refresh.scope is not None else granted_scope,
jti=secrets.token_urlsafe(24),
exp=now + min(refresh.expires_in or 1209600, 1209600),
),
)
if refresh is not None
else None
)
if len(access) > 12288 or (sealed_refresh is not None and len(sealed_refresh) > 12288):
return _oauth_error(502, "server_error", "The upstream credential is too large")
from mcp.shared.auth import OAuthToken
return JSONResponse(
OAuthToken(
access_token=access, token_type="Bearer", expires_in=ttl, refresh_token=sealed_refresh, scope=granted_scope
).model_dump(exclude_none=True),
headers=TOKEN_NO_CACHE_HEADERS,
)

View file

@ -11,6 +11,8 @@ from typing import TYPE_CHECKING, Final
if TYPE_CHECKING:
from mcp.server.context import ServerRequestContext
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential
# The SDK 1.x ``mcp.server.lowlevel.server.request_ctx`` ContextVar was removed in
# SDK 2, which hands each request handler a ``ServerRequestContext`` argument
# instead. The handlers set this var so downstream helpers (session auth caching,
@ -40,3 +42,19 @@ _mcp_gateway_server_name: Final[ContextVar[str | None]] = ContextVar("_mcp_gatew
# Set server-side by the /mcp/proxy route. Never populated from client-supplied headers.
_mcp_proxy_mode: Final[ContextVar[bool]] = ContextVar("_mcp_proxy_mode", default=False)
def get_connection_credential(server_id: str) -> "ConnectionCredential | None":
from starlette.requests import Request
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential
context: Final = get_active_mcp_request_ctx()
request: Final = context.request if context is not None else None
if not isinstance(request, Request):
return None
value: Final = request.scope.get("litellm.mcp.connection_grant")
if not isinstance(value, ConnectionCredential) or value.binding.server_id != server_id:
return None
return value

View file

@ -4176,7 +4176,22 @@ class MCPServerManager:
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
transport: Final = resolved_server.transport or MCPTransport.sse
spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server)
provider: Final = cred_provider or self._cred_provider
from litellm.proxy._experimental.mcp_server.mcp_context import get_connection_credential
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken
from litellm.proxy._experimental.mcp_server.outbound_credentials.presented_token_store import (
PresentedOAuthTokenStore,
)
connection: Final = get_connection_credential(resolved_server.server_id)
if connection is not None and connection.exp <= int(datetime.datetime.now(datetime.timezone.utc).timestamp()):
raise HTTPException(status_code=401, detail="MCP connection credential expired; reconnect")
provider: Final = (
UpstreamCredentialProvider(
oauth_token_store=PresentedOAuthTokenStore(OAuthToken(access_token=connection.token.get_secret_value()))
)
if connection is not None and resolved_server.needs_user_oauth_token
else cred_provider or self._cred_provider
)
# A caller-supplied per-request override (mcp_auth_header / x-mcp-*) defers to the v1 path
# so it wins - except for the modes the v2 resolver owns per-caller (authorization_code's
# stored token, token_exchange's RFC 8693 minted token, id_jag's minted assertion, and the

View file

@ -38,7 +38,7 @@ from datetime import datetime, timedelta
from typing import Final, Literal, TypeAlias
import jwt
from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError
from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError, field_serializer
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value, encrypt_value
@ -580,3 +580,58 @@ def _decrypt_refresh(
return RefreshCredential.model_validate_json(plaintext)
except ValidationError:
return MalformedPayload()
class ConnectionBinding(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
key_hash: str = Field(min_length=1)
server_id: str = Field(min_length=1)
resource: str = Field(min_length=1)
class ConnectionBootstrap(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
kind: Literal["connection_bootstrap"] = "connection_bootstrap"
binding: ConnectionBinding
exp: int
class ConnectionAuthorization(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
kind: Literal["connection_authorization"] = "connection_authorization"
binding: ConnectionBinding
client_id: str = Field(min_length=1)
redirect_uri: str = Field(min_length=1)
state: str
code_challenge: str = Field(min_length=43, max_length=43)
scope: str
jti: str = Field(min_length=1)
exp: int
class ConnectionCode(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
kind: Literal["connection_code"] = "connection_code"
authorization: ConnectionAuthorization
upstream_code: SecretStr = Field(min_length=1)
jti: str = Field(min_length=1)
exp: int
@field_serializer("upstream_code", when_used="json")
def serialize_code(self, value: SecretStr) -> str:
return value.get_secret_value()
class ConnectionCredential(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
kind: Literal["connection_access", "connection_refresh"]
binding: ConnectionBinding
client_id: str = Field(min_length=1)
token: SecretStr = Field(min_length=1)
scope: str | None = None
jti: str = Field(min_length=1)
exp: int
@field_serializer("token", when_used="json")
def serialize_token(self, value: SecretStr) -> str:
return value.get_secret_value()

View file

@ -16,7 +16,8 @@ import types
import uuid
from collections import Counter
from collections.abc import AsyncIterator, Callable, Iterable, Mapping, Sequence
from datetime import datetime
from datetime import datetime, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol
import httpx
@ -59,6 +60,7 @@ from litellm.proxy._experimental.mcp_server.exceptions import (
MCPToolResultError,
MCPUpstreamAuthError,
)
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import CONNECTION_SCOPE_KEY, connection_challenge
from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_active_toolset_id,
_mcp_gateway_initialize_instructions,
@ -79,6 +81,7 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import (
get_route_relative_request_path,
well_known_root_suffix,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionBinding, ConnectionCredential
from litellm.proxy._experimental.mcp_server.ui_session_utils import is_ui_session_credential
from litellm.proxy._experimental.mcp_server.utils import (
LITELLM_MCP_SERVER_DESCRIPTION,
@ -97,6 +100,7 @@ from litellm.proxy._types import (
ProxyException,
SpecialMCPServerNames,
UserAPIKeyAuth,
hash_token,
)
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
@ -4262,7 +4266,22 @@ if MCP_AVAILABLE:
# authorization server is the gateway itself, vaulting via the
# authorize interlude); the per-server relay advertised below
# cannot vault without a litellm key on its token request.
if await global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth):
connection = scope.get(CONNECTION_SCOPE_KEY)
if (
isinstance(connection, ConnectionCredential)
and connection.binding.server_id == server.server_id
and connection.exp > int(datetime.now(timezone.utc).timestamp())
):
status, _ = await _probe_upstream_auth(
server.url or "", f"Bearer {connection.token.get_secret_value()}"
)
if status == 403:
raise HTTPException(status_code=403, detail="Upstream denied access")
if status != 401:
continue
if connection is None and await global_mcp_server_manager.has_user_oauth_token(
server, user_api_key_auth
):
continue
if _is_mcp_admitted_user_subject(user_api_key_auth):
@ -4279,6 +4298,39 @@ if MCP_AVAILABLE:
request = StarletteRequest(scope)
base_url = get_request_base_url(request)
if (
get_route_relative_request_path(scope) == "/mcp"
and len(mcp_servers or ()) == 1
and request.headers.get("x-litellm-api-key")
and user_api_key_auth is not None
and server.oauth_identity_binding is None
):
allowed = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
if server.server_id not in allowed:
raise HTTPException(
status_code=403, detail="Key is not allowed to access the selected MCP server"
)
from litellm.proxy.auth.user_api_key_auth import (
_get_bearer_token_or_received_api_key, # pyright: ignore[reportPrivateUsage] # reuse the admission header parser
)
binding = ConnectionBinding(
key_hash=hash_token(
_get_bearer_token_or_received_api_key(request.headers["x-litellm-api-key"])
),
server_id=server.server_id,
resource=f"{base_url}/mcp",
)
challenge_headers = MappingProxyType(
{
"www-authenticate": connection_challenge(request, binding),
"Cache-Control": "no-store",
}
)
raise HTTPException(
status_code=401, detail="Authorize the selected MCP server", headers=challenge_headers
)
_path = get_route_relative_request_path(scope)
# Pick the well-known AS-metadata form that matches the inbound route

View file

@ -9437,3 +9437,92 @@ class TestScopedSessionAdmission:
def test_scope_field_cannot_be_forged_through_construction(self):
forged = UserAPIKeyAuth(user_id="u1", mcp_session_resource_server_id="any-server")
assert forged.mcp_session_resource_server_id is None
@pytest.mark.asyncio
@pytest.mark.parametrize(
"key,selector,expected",
[
("sk-original", "connection-target", 200),
("sk-other", "connection-target", 401),
("sk-original", "other-target", 403),
(None, "connection-target", 401),
],
)
@pytest.mark.parametrize("grant_state", ["valid", "expired", "refresh", "oversized"])
async def test_connection_credential_requires_exact_key_and_server(monkeypatch, key, selector, expected, grant_state):
import json
from litellm.proxy._experimental.mcp_server import gateway_dcr_flow as flow
from litellm.proxy._experimental.mcp_server.auth import user_api_key_auth_mcp as admission
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionBinding
from litellm.proxy._types import hash_token
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
monkeypatch.setenv("LITELLM_SALT_KEY", "connection-admission-test-salt")
server = MCPServer(
server_id="connection-target",
name="connection-target",
server_name="connection-target",
alias="connection-target",
url="https://mcp.example/mcp",
transport="http",
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
)
monkeypatch.setitem(global_mcp_server_manager.registry, server.server_id, server)
auth = UserAPIKeyAuth(
api_key=hash_token(key or "sk-original"),
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="connection-test",
mcp_servers=[server.server_id],
),
)
monkeypatch.setattr(admission, "user_api_key_auth", AsyncMock(return_value=auth))
binding = ConnectionBinding(
key_hash=hash_token("sk-original"), server_id=server.server_id, resource="https://gateway.example/mcp"
)
issued = json.loads(
flow.mint_connection_tokens(
binding, "client", {"access_token": "provider-token", "refresh_token": "provider-refresh"}
).body
)
credential = flow.open_connection_credential(issued["access_token"])
assert credential is not None
token = (
flow._seal(flow.CONNECTION_ACCESS_PREFIX, credential.model_copy(update={"exp": 1}))
if grant_state == "expired"
else flow.CONNECTION_ACCESS_PREFIX + "x" * 12289
if grant_state == "oversized"
else issued["refresh_token"]
if grant_state == "refresh"
else issued["access_token"]
)
expected_status = 401 if expected == 200 and grant_state != "valid" else expected
scope = {
"type": "http",
"method": "POST",
"scheme": "https",
"path": "/mcp",
"headers": [
(b"host", b"gateway.example"),
*([(b"x-litellm-api-key", key.encode())] if key is not None else []),
(b"x-mcp-servers", selector.encode()),
(b"authorization", f"Bearer {token}".encode()),
],
}
if expected_status != 200:
with pytest.raises(HTTPException) as exc:
await MCPRequestHandler.process_mcp_request(scope)
assert exc.value.status_code == expected_status
if key == "sk-original" and selector == "connection-target" and grant_state != "valid":
assert "resource_metadata=" in exc.value.headers["www-authenticate"]
assert flow.CONNECTION_SCOPE_KEY not in scope
return
_, _, _, server_headers, oauth_headers, raw_headers = await MCPRequestHandler.process_mcp_request(scope)
assert scope[flow.CONNECTION_SCOPE_KEY].token.get_secret_value() == "provider-token"
assert not oauth_headers
assert "authorization" not in raw_headers
assert all(token not in value for value in raw_headers.values())
assert not server_headers

View file

@ -12345,3 +12345,690 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session(
proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called()
proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called()
proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called()
@pytest.fixture
def keyed_oauth_client(monkeypatch):
from types import SimpleNamespace
import httpx
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.caching.caching import DualCache
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
from litellm.proxy._experimental.mcp_server import gateway_dcr_flow as flow
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionBinding
from litellm.proxy._types import hash_token
from litellm.types.mcp_server.mcp_server_manager import MCPServer
monkeypatch.setenv("LITELLM_SALT_KEY", "connection-flow-regression-salt")
cache = DualCache()
redis = SimpleNamespace(async_increment=cache.async_increment_cache, async_get_cache=cache.async_get_cache)
monkeypatch.setattr(proxy_server, "redis_usage_cache", redis)
server = MCPServer(
server_id="github-test",
name="GitHub",
server_name="github-test",
alias="github-test",
url="https://mcp.example/mcp",
transport="http",
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
client_id="upstream-app",
client_secret="upstream-secret",
scopes=["read:user"],
authorization_url="https://provider.example/authorize",
token_url="https://provider.example/token",
)
real_validate = flow.validate_connection_binding
validate = AsyncMock(return_value=server)
monkeypatch.setattr(flow, "validate_connection_binding", validate)
monkeypatch.setattr(endpoints, "validate_connection_binding", validate)
response = httpx.Response(
200,
json={
"access_token": "provider-access",
"refresh_token": "provider-refresh",
"expires_in": 3600,
"scope": "read:user",
"token_type": "Bearer",
},
request=httpx.Request("POST", "https://provider.example/token"),
)
upstream = SimpleNamespace(post=AsyncMock(return_value=response))
monkeypatch.setattr(endpoints, "get_async_httpx_client", lambda **kwargs: upstream)
vault = AsyncMock(side_effect=AssertionError("a connection must not write the user vault"))
monkeypatch.setattr(endpoints, "_store_per_user_token_server_side", vault)
app = FastAPI()
app.include_router(endpoints.router)
binding = ConnectionBinding(
key_hash=hash_token("sk-original"), server_id=server.server_id, resource="https://gateway.example/mcp"
)
bootstrap = flow.mint_connection_bootstrap(binding)
with TestClient(app, base_url="https://gateway.example", follow_redirects=False) as client:
yield SimpleNamespace(
client=client,
binding=binding,
bootstrap=bootstrap,
upstream=upstream,
vault=vault,
validate=validate,
real_validate=real_validate,
server=server,
redis=redis,
)
def _start_keyed_oauth(harness):
from urllib.parse import parse_qs, urlparse
client = harness.client
metadata = client.get("/.well-known/oauth-protected-resource/mcp", params={"connection": harness.bootstrap})
assert metadata.status_code == 200, metadata.text
assert metadata.json()["resource"] == harness.binding.resource
issuer = metadata.json()["authorization_servers"][0]
discovery = client.get("/.well-known/oauth-authorization-server" + urlparse(issuer).path)
assert discovery.status_code == 200, discovery.text
assert discovery.json()["issuer"] == issuer
registered = client.post(
discovery.json()["registration_endpoint"],
json={
"redirect_uris": ["http://localhost:33418/callback"],
"client_name": "Cursor <untrusted>",
},
)
assert registered.status_code == 201, registered.text
client_id = registered.json()["client_id"]
verifier = "v" * 43
challenge = urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode()
consent = client.get(
discovery.json()["authorization_endpoint"],
params={
"client_id": client_id,
"redirect_uri": "http://localhost:33418/callback",
"response_type": "code",
"state": "client-state",
"code_challenge": challenge,
"code_challenge_method": "S256",
"resource": harness.binding.resource,
"scope": "read:user",
},
)
assert consent.status_code == 200, consent.text
assert "Cursor &lt;untrusted&gt;" in consent.text
assert "http://localhost:33418/callback" in consent.text
assert "read:user" in consent.text
assert "/sso/" not in consent.text
handle = next(
cookie.name.removeprefix("mcp_connection_")
for cookie in client.cookies.jar
if cookie.name.startswith("mcp_connection_")
)
return client_id, verifier, handle
def _complete_keyed_oauth(harness):
from urllib.parse import parse_qs, urlparse
client_id, verifier, handle = _start_keyed_oauth(harness)
consent = harness.client.post("/authorize/connection/complete", data={"flow": handle, "decision": "approve"})
assert consent.status_code == 307, consent.text
upstream = urlparse(consent.headers["location"])
assert upstream.netloc == "provider.example"
params = parse_qs(upstream.query)
assert params["client_id"] == ["upstream-app"]
assert params["redirect_uri"] == ["https://gateway.example/callback"]
callback = harness.client.get("/callback", params={"state": params["state"][0], "code": "provider-code"})
assert callback.status_code == 302, callback.text
result = parse_qs(urlparse(callback.headers["location"]).query)
assert result["state"] == ["client-state"]
assert result["code"] != ["provider-code"]
return {
"grant_type": "authorization_code",
"client_id": client_id,
"code": result["code"][0],
"code_verifier": verifier,
"redirect_uri": "http://localhost:33418/callback",
"resource": harness.binding.resource,
}
def test_keyed_connection_headerless_exchange_and_rotating_refresh(keyed_oauth_client):
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import open_connection_credential
harness = keyed_oauth_client
payload = _complete_keyed_oauth(harness)
token = harness.client.post("/token", data=payload)
assert token.status_code == 200, token.text
grant = open_connection_credential(token.json()["access_token"])
assert grant is not None
assert grant.binding == harness.binding
assert grant.token.get_secret_value() == "provider-access"
assert grant.client_id == payload["client_id"]
posted = harness.upstream.post.call_args.kwargs["data"]
assert posted["code"] == "provider-code"
assert posted["client_id"] == "upstream-app"
assert posted["code_verifier"] == payload["code_verifier"]
replay = harness.client.post("/token", data=payload)
assert replay.status_code == 400
assert replay.json()["error"] == "invalid_grant"
assert harness.upstream.post.call_count == 1
refresh_payload = {
"grant_type": "refresh_token",
"client_id": payload["client_id"],
"refresh_token": token.json()["refresh_token"],
"resource": harness.binding.resource,
}
refreshed = harness.client.post("/token", data=refresh_payload)
assert refreshed.status_code == 200, refreshed.text
assert refreshed.json()["refresh_token"] != token.json()["refresh_token"]
assert harness.upstream.post.call_args.kwargs["data"]["refresh_token"] == "provider-refresh"
replayed_refresh = harness.client.post("/token", data=refresh_payload)
assert replayed_refresh.status_code == 400
assert harness.upstream.post.call_count == 2
harness.vault.assert_not_called()
@pytest.mark.parametrize("change", ["verifier", "redirect", "resource", "tamper"])
def test_keyed_connection_rejects_bad_exchange_before_upstream(keyed_oauth_client, change):
harness = keyed_oauth_client
payload = _complete_keyed_oauth(harness)
field, value = {
"verifier": ("code_verifier", "w" * 43),
"redirect": ("redirect_uri", "http://localhost:33419/callback"),
"resource": ("resource", "https://gateway.example/mcp/other"),
"tamper": ("code", "llm_ccode_invalid"),
}[change]
response = harness.client.post("/token", data={**payload, field: value})
assert response.status_code == 400, response.text
harness.upstream.post.assert_not_called()
harness.vault.assert_not_called()
def test_keyed_connection_denial_and_replayed_consent_never_exchange(keyed_oauth_client):
harness = keyed_oauth_client
_, _, handle = _start_keyed_oauth(harness)
denied = harness.client.post("/authorize/connection/complete", data={"flow": handle, "decision": "deny"})
assert denied.status_code == 302
assert "error=access_denied" in denied.headers["location"]
assert "state=client-state" in denied.headers["location"]
repeated = harness.client.post("/authorize/connection/complete", data={"flow": handle, "decision": "approve"})
assert repeated.status_code == 400
harness.upstream.post.assert_not_called()
harness.vault.assert_not_called()
def test_keyed_connection_redis_outage_prevents_exchange(keyed_oauth_client):
harness = keyed_oauth_client
payload = _complete_keyed_oauth(harness)
harness.redis.async_increment = AsyncMock(side_effect=ConnectionError("redis unavailable"))
refused = harness.client.post("/token", data=payload)
assert refused.status_code == 503
assert refused.json()["error"] == "temporarily_unavailable"
harness.upstream.post.assert_not_called()
harness.vault.assert_not_called()
@pytest.mark.parametrize("state_length", [1, 1024])
def test_keyed_connection_cookie_sizes_are_browser_safe(keyed_oauth_client, state_length):
harness = keyed_oauth_client
client_id, verifier, _ = _start_keyed_oauth(harness)
challenge = urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode()
response = harness.client.get(
"/authorize",
params={
"client_id": client_id,
"redirect_uri": "http://localhost:33418/callback",
"response_type": "code",
"state": "s" * state_length,
"code_challenge": challenge,
"code_challenge_method": "S256",
"resource": harness.binding.resource,
},
)
if response.status_code == 400:
assert response.json()["error"] == "invalid_request"
assert state_length == 1024
harness.upstream.post.assert_not_called()
return
assert response.status_code == 200
assert all(len(value) <= 4096 for value in response.headers.get_list("set-cookie"))
handle = response.text.split('name="flow" value="')[1].split('"')[0]
approved = harness.client.post("/authorize/connection/complete", data={"flow": handle, "decision": "approve"})
if approved.status_code == 400:
assert approved.json()["error"] == "invalid_request"
assert state_length == 1024
else:
assert approved.status_code == 307
assert all(len(value) <= 4096 for value in approved.headers.get_list("set-cookie"))
harness.upstream.post.assert_not_called()
def test_keyed_connection_client_cannot_enter_session_login_flow(keyed_oauth_client):
harness = keyed_oauth_client
client_id, verifier, _ = _start_keyed_oauth(harness)
challenge = urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode()
response = harness.client.get(
"/authorize/mcp-session",
params={
"client_id": client_id,
"redirect_uri": "http://localhost:33418/callback",
"response_type": "code",
"state": "state",
"code_challenge": challenge,
"code_challenge_method": "S256",
"resource": harness.binding.resource,
},
)
assert response.status_code == 400
assert response.json()["error"] == "invalid_client"
assert "location" not in response.headers
harness.upstream.post.assert_not_called()
@pytest.mark.parametrize("stage", ["authorize", "exchange", "refresh"])
@pytest.mark.parametrize("policy", ["blocked", "expired", "denied_server", "denied_route", "allowed"])
def test_keyed_connection_reloads_live_key_before_provider(keyed_oauth_client, monkeypatch, stage, policy):
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import gateway_dcr_flow as flow
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
from litellm.proxy._experimental.mcp_server.auth import user_api_key_auth_mcp as admission
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
from litellm.proxy.auth import auth_checks
harness = keyed_oauth_client
payload = _complete_keyed_oauth(harness)
token = harness.client.post("/token", data=payload) if stage == "refresh" else None
harness.upstream.post.reset_mock()
real_validate = harness.real_validate
monkeypatch.setattr(flow, "validate_connection_binding", real_validate)
monkeypatch.setattr(endpoints, "validate_connection_binding", real_validate)
key = UserAPIKeyAuth(
api_key=harness.binding.key_hash,
blocked=policy == "blocked",
expires=datetime.now(timezone.utc) - timedelta(seconds=1) if policy == "expired" else None,
allowed_routes=["/chat/completions"] if policy == "denied_route" else None,
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="connection-policy-test",
mcp_servers=["other-server"] if policy == "denied_server" else [harness.server.server_id],
),
)
lookup = AsyncMock(return_value=key)
monkeypatch.setattr(auth_checks, "get_key_object", lookup)
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
monkeypatch.setattr(admission, "_run_centralized_common_checks", AsyncMock())
monkeypatch.setitem(global_mcp_server_manager.registry, harness.server.server_id, harness.server)
if stage == "authorize":
challenge = urlsafe_b64encode(hashlib.sha256(payload["code_verifier"].encode()).digest()).rstrip(b"=").decode()
response = harness.client.get(
"/authorize",
params={
"client_id": payload["client_id"],
"redirect_uri": payload["redirect_uri"],
"response_type": "code",
"code_challenge": challenge,
"code_challenge_method": "S256",
"resource": harness.binding.resource,
},
)
elif stage == "exchange":
response = harness.client.post("/token", data=payload)
else:
assert token is not None and token.status_code == 200
response = harness.client.post(
"/token",
data={
"grant_type": "refresh_token",
"client_id": payload["client_id"],
"refresh_token": token.json()["refresh_token"],
"resource": harness.binding.resource,
},
)
assert response.status_code == (200 if policy == "allowed" else 401 if policy in ("blocked", "expired") else 403), (
response.text
)
lookup.assert_awaited_once()
assert lookup.call_args.kwargs["hashed_token"] == harness.binding.key_hash
assert harness.upstream.post.call_count == int(policy == "allowed" and stage != "authorize")
harness.vault.assert_not_called()
@pytest.mark.parametrize(
"case",
[
"absent_redis",
"wrong_client",
"expired_code",
"refresh_as_code",
"access_as_refresh",
"expanded_scope",
"missing_pkce",
"named_route",
],
)
def test_keyed_connection_rejects_invalid_grants_without_provider_calls(keyed_oauth_client, monkeypatch, case):
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import gateway_dcr_flow as flow
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCode
harness = keyed_oauth_client
payload = _complete_keyed_oauth(harness)
issued = (
harness.client.post("/token", data=payload)
if case in ("refresh_as_code", "access_as_refresh", "expanded_scope")
else None
)
harness.upstream.post.reset_mock()
if case == "absent_redis":
monkeypatch.setattr(proxy_server, "redis_usage_cache", None)
monkeypatch.setattr(proxy_server.user_api_key_cache, "redis_cache", None)
response = harness.client.post("/token", data=payload)
elif case == "wrong_client":
other = harness.client.post(
"/register",
params={"connection": harness.bootstrap},
json={
"redirect_uris": [payload["redirect_uri"]],
"client_name": "Different application",
},
)
assert other.status_code == 201
response = harness.client.post("/token", data={**payload, "client_id": other.json()["client_id"]})
elif case == "expired_code":
opened = flow.open_connection_code(payload["code"])
assert opened is not None
expired = ConnectionCode(
authorization=opened.authorization, upstream_code=opened.upstream_code, jti=opened.jti, exp=1
)
response = harness.client.post(
"/token", data={**payload, "code": flow._seal(flow.CONNECTION_CODE_PREFIX, expired)}
)
elif case == "missing_pkce":
response = harness.client.post("/token", data={k: v for k, v in payload.items() if k != "code_verifier"})
elif case == "named_route":
response = harness.client.post("/github-test/token", data=payload)
else:
assert issued is not None and issued.status_code == 200
if case == "refresh_as_code":
response = harness.client.post("/token", data={**payload, "code": issued.json()["refresh_token"]})
else:
response = harness.client.post(
"/token",
data={
"client_id": payload["client_id"],
"grant_type": "refresh_token",
"refresh_token": issued.json()["access_token" if case == "access_as_refresh" else "refresh_token"],
"scope": "admin" if case == "expanded_scope" else "read:user",
},
)
assert response.status_code == (503 if case == "absent_redis" else 400), response.text
harness.upstream.post.assert_not_called()
harness.vault.assert_not_called()
@pytest.mark.parametrize(
"body,status",
[
({}, 502),
({"access_token": "bad", "token_type": "MAC"}, 502),
({"access_token": "old", "expires_in": 0}, 502),
({"access_token": "x" * 13000}, 502),
({"access_token": "short", "expires_in": 30}, 200),
],
)
def test_keyed_connection_validates_provider_token_response(keyed_oauth_client, body, status):
import httpx
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import open_connection_credential
harness = keyed_oauth_client
payload = _complete_keyed_oauth(harness)
harness.upstream.post.return_value = httpx.Response(
200, json=body, request=httpx.Request("POST", harness.server.token_url)
)
response = harness.client.post("/token", data=payload)
assert response.status_code == status, response.text
harness.upstream.post.assert_awaited_once()
if status == 200:
grant = open_connection_credential(response.json()["access_token"])
assert grant is not None and grant.token.get_secret_value() == "short"
assert response.json()["expires_in"] == 30
assert "refresh_token" not in response.json()
else:
assert "access_token" not in response.json()
harness.vault.assert_not_called()
def test_keyed_connection_refresh_preserves_nonrotating_provider_token(keyed_oauth_client):
import httpx
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import open_connection_credential
harness = keyed_oauth_client
payload = _complete_keyed_oauth(harness)
issued = harness.client.post("/token", data=payload)
assert issued.status_code == 200
harness.upstream.post.return_value = httpx.Response(
200,
json={"access_token": "renewed", "expires_in": 1800},
request=httpx.Request("POST", harness.server.token_url),
)
response = harness.client.post(
"/token",
data={
"client_id": payload["client_id"],
"grant_type": "refresh_token",
"refresh_token": issued.json()["refresh_token"],
},
)
assert response.status_code == 200, response.text
refresh = open_connection_credential(response.json()["refresh_token"], refresh=True)
assert refresh is not None and refresh.token.get_secret_value() == "provider-refresh"
assert refresh.scope == "read:user"
assert response.json()["refresh_token"] != issued.json()["refresh_token"]
assert harness.upstream.post.call_args.kwargs["data"]["refresh_token"] == "provider-refresh"
harness.vault.assert_not_called()
@pytest.mark.parametrize(
"route",
[
"/.well-known/oauth-protected-resource/mcp",
"/.well-known/oauth-authorization-server/mcp-connect/llm_cboot_invalid",
"/register",
],
)
def test_keyed_connection_discovery_rejects_tampered_bootstrap(keyed_oauth_client, route):
harness = keyed_oauth_client
response = (
harness.client.post(
route, params={"connection": "llm_cboot_invalid"}, json={"redirect_uris": ["http://localhost/callback"]}
)
if route == "/register"
else harness.client.get(route, params={"connection": "llm_cboot_invalid"})
)
assert response.status_code == 400
harness.validate.assert_not_awaited()
harness.upstream.post.assert_not_called()
@pytest.mark.parametrize(
"field,value,error",
[
("resource", "https://wrong.example/mcp", "invalid_target"),
("code_challenge", "bad", "invalid_request"),
("code_challenge_method", "plain", "invalid_request"),
("scope", "admin", "invalid_scope"),
("redirect_uri", "http://localhost:4001/wrong", "invalid_request"),
],
)
def test_keyed_connection_authorize_rejects_invalid_parameters(keyed_oauth_client, field, value, error):
harness = keyed_oauth_client
client_id, verifier, _ = _start_keyed_oauth(harness)
challenge = urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode()
response = harness.client.get(
"/authorize",
params={
"client_id": client_id,
"redirect_uri": "http://localhost:33418/callback",
"response_type": "code",
"state": "state",
"code_challenge": challenge,
"code_challenge_method": "S256",
"resource": harness.binding.resource,
field: value,
},
)
assert response.status_code == 400
assert response.json()["error"] == error
assert "location" not in response.headers
harness.upstream.post.assert_not_called()
@pytest.mark.parametrize("cancel", [True, False])
def test_keyed_connection_callback_cancellation_and_replay(keyed_oauth_client, cancel):
from urllib.parse import parse_qs, urlparse
harness = keyed_oauth_client
_, _, handle = _start_keyed_oauth(harness)
approved = harness.client.post("/authorize/connection/complete", data={"flow": handle, "decision": "approve"})
assert approved.status_code == 307
state = parse_qs(urlparse(approved.headers["location"]).query)["state"][0]
cookies = dict(harness.client.cookies.items())
response = harness.client.get(
"/callback", params={"state": state, **({"error": "access_denied"} if cancel else {"code": "provider-code"})}
)
assert response.status_code == 302
returned = parse_qs(urlparse(response.headers["location"]).query)
assert returned["state"] == ["client-state"]
if cancel:
assert returned["error"] == ["access_denied"]
assert "code" not in returned
else:
harness.client.cookies.update(cookies)
replay = harness.client.get("/callback", params={"state": state, "code": "provider-code"})
assert replay.status_code == 400
assert "location" not in replay.headers
harness.upstream.post.assert_not_called()
harness.vault.assert_not_called()
def test_keyed_connection_replayed_approved_consent_does_not_redirect_twice(keyed_oauth_client):
harness = keyed_oauth_client
_, _, handle = _start_keyed_oauth(harness)
cookies = dict(harness.client.cookies.items())
approved = harness.client.post("/authorize/connection/complete", data={"flow": handle, "decision": "approve"})
assert approved.status_code == 307
harness.client.cookies.update(cookies)
repeated = harness.client.post("/authorize/connection/complete", data={"flow": handle, "decision": "approve"})
assert repeated.status_code == 400
assert repeated.json()["error"] == "invalid_grant"
assert "location" not in repeated.headers
harness.upstream.post.assert_not_called()
def test_keyed_connection_expired_callback_never_issues_code(keyed_oauth_client, monkeypatch):
from urllib.parse import parse_qs, urlparse
from litellm.proxy._experimental.mcp_server import gateway_dcr_flow as flow
harness = keyed_oauth_client
_, _, handle = _start_keyed_oauth(harness)
approved = harness.client.post("/authorize/connection/complete", data={"flow": handle, "decision": "approve"})
assert approved.status_code == 307
state = parse_qs(urlparse(approved.headers["location"]).query)["state"][0]
clock = MagicMock()
clock.now.return_value = datetime.now(timezone.utc) + timedelta(days=1)
monkeypatch.setattr(flow, "datetime", clock)
refused = harness.client.get("/callback", params={"state": state, "code": "provider-code"})
assert refused.status_code == 400
assert "location" not in refused.headers
assert "expired" in refused.json()["detail"]
harness.upstream.post.assert_not_called()
def test_keyed_connection_consent_rejects_metadata_over_cookie_limit(keyed_oauth_client):
harness = keyed_oauth_client
redirect = "http://localhost:33418/" + "c" * 230
registered = harness.client.post(
"/register",
params={"connection": harness.bootstrap},
json={"redirect_uris": [redirect, redirect + "1", redirect + "2"]},
)
assert registered.status_code == 201, registered.text
response = harness.client.get(
"/authorize",
params={
"client_id": registered.json()["client_id"],
"redirect_uri": redirect,
"response_type": "code",
"state": "s" * 1024,
"code_challenge": "c" * 43,
"code_challenge_method": "S256",
"resource": harness.binding.resource,
},
)
assert response.status_code == 400
assert response.json()["error"] == "invalid_request"
assert "set-cookie" not in response.headers
harness.upstream.post.assert_not_called()
@pytest.mark.parametrize("endpoint", ["/github-test/authorize", "/token"])
def test_keyed_connection_rejects_wrong_endpoint_and_grant_type(keyed_oauth_client, endpoint):
harness = keyed_oauth_client
client_id, _, _ = _start_keyed_oauth(harness)
response = (
harness.client.get(endpoint, params={"client_id": client_id, "redirect_uri": "http://localhost:33418/callback"})
if endpoint.endswith("authorize")
else harness.client.post(endpoint, data={"client_id": client_id, "grant_type": "client_credentials"})
)
assert response.status_code == 400
assert response.json()["error"] == (
"invalid_client" if endpoint.endswith("authorize") else "unsupported_grant_type"
)
harness.upstream.post.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("mode", ["wrong_resource", "m2m"])
async def test_keyed_connection_binding_rejects_resource_or_server_mode(keyed_oauth_client, monkeypatch, mode):
from starlette.requests import Request
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.auth import user_api_key_auth_mcp as admission
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
from litellm.proxy.auth import auth_checks
harness = keyed_oauth_client
lookup = AsyncMock(
return_value=UserAPIKeyAuth(
api_key=harness.binding.key_hash,
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="connection-mode-test",
mcp_servers=[harness.server.server_id],
),
)
)
monkeypatch.setattr(auth_checks, "get_key_object", lookup)
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
monkeypatch.setattr(admission, "_run_centralized_common_checks", AsyncMock())
monkeypatch.setitem(
global_mcp_server_manager.registry,
harness.server.server_id,
harness.server.model_copy(update={"oauth2_flow": "client_credentials"}),
)
request = Request(
{
"type": "http",
"scheme": "https",
"method": "GET",
"path": "/authorize",
"headers": [(b"host", b"wrong.example" if mode == "wrong_resource" else b"gateway.example")],
}
)
with pytest.raises(HTTPException) as exc:
await harness.real_validate(request, harness.binding)
assert exc.value.status_code == 400
assert lookup.call_count == (0 if mode == "wrong_resource" else 1)
harness.upstream.post.assert_not_called()

View file

@ -9469,6 +9469,47 @@ class TestPreemptive401ModeAware:
discovery.assert_not_awaited()
@pytest.mark.asyncio
async def test_keyed_aggregate_managed_oauth_uses_resource_metadata(self):
from litellm.proxy._experimental.mcp_server import server as server_module
upstream = _make_oauth2_server("interactive")
scope = {
"type": "http",
"method": "POST",
"scheme": "https",
"path": "/mcp",
"headers": [(b"host", b"gateway.example"), (b"x-litellm-api-key", b"Bearer sk-test")],
}
with (
patch.object(server_module.global_mcp_server_manager, "get_mcp_server_by_name", return_value=upstream),
patch.object(
server_module.global_mcp_server_manager, "has_user_oauth_token", new=AsyncMock(return_value=False)
),
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers",
new=AsyncMock(return_value=[upstream.server_id]),
),
patch.dict(os.environ, {"LITELLM_SALT_KEY": "regression-salt-not-a-real-secret"}),
pytest.raises(HTTPException) as exc,
):
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope=scope,
mcp_servers=["interactive"],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-test"),
client_ip=None,
allowed_server_ids={upstream.server_id},
)
assert exc.value.status_code == 401
assert (
'resource_metadata="https://gateway.example/.well-known/oauth-protected-resource/mcp?'
in exc.value.headers["www-authenticate"]
)
assert "sk-test" not in exc.value.headers["www-authenticate"]
@pytest.mark.asyncio
async def test_gateway_managed_interactive_no_token_challenges_with_x_litellm_api_key(self):
"""No stored token, key in x-litellm-api-key (oauth2_headers empty): 401."""
@ -10208,3 +10249,81 @@ async def test_streamable_http_rejects_modern_protocol_version(header_value: str
assert header_value in body["error"]["message"]
for version in body["error"]["message"].split("supported: ")[1].split(", "):
assert version in HANDSHAKE_PROTOCOL_VERSIONS
@pytest.mark.asyncio
@pytest.mark.parametrize(
"upstream_status,expected,key_allowed", [(200, None, True), (401, 401, True), (403, 403, True), (401, 403, False)]
)
async def test_keyed_connection_preflight_uses_presented_token_without_vault_fallback(
monkeypatch, upstream_status, expected, key_allowed
):
from datetime import timezone
from pydantic import SecretStr
from litellm.proxy._experimental.mcp_server import server as server_module
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import CONNECTION_SCOPE_KEY
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
ConnectionBinding,
ConnectionCredential,
)
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
upstream = _make_oauth2_server("connection-preflight")
credential = ConnectionCredential(
kind="connection_access",
binding=ConnectionBinding(
key_hash="test-hash", server_id=upstream.server_id, resource="https://gateway.example/mcp"
),
client_id="client",
token=SecretStr("connection-provider-token"),
jti="connection-preflight",
exp=int(datetime.now(timezone.utc).timestamp()) + 300,
)
scope = {
"type": "http",
"method": "POST",
"scheme": "https",
"path": "/mcp",
CONNECTION_SCOPE_KEY: credential,
"headers": [(b"host", b"gateway.example"), (b"x-litellm-api-key", b"Bearer sk-test")],
}
manager = server_module.global_mcp_server_manager
probe = AsyncMock(return_value=(upstream_status, {}))
vault = AsyncMock(side_effect=AssertionError("a rejected connection must not use stored credentials"))
with (
patch.object(manager, "get_mcp_server_by_name", return_value=upstream),
patch.object(manager, "has_user_oauth_token", vault),
patch.object(
MCPRequestHandler,
"get_allowed_mcp_servers",
AsyncMock(return_value=[upstream.server_id] if key_allowed else []),
),
patch.object(server_module, "_probe_upstream_auth", probe),
patch.dict(os.environ, {"LITELLM_SALT_KEY": "preflight-regression-salt"}),
):
if expected is None:
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope,
[upstream.server_name],
None,
None,
UserAPIKeyAuth(api_key="sk-test"),
None,
allowed_server_ids={upstream.server_id},
)
else:
with pytest.raises(HTTPException) as exc:
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope,
[upstream.server_name],
None,
None,
UserAPIKeyAuth(api_key="sk-test"),
None,
allowed_server_ids={upstream.server_id},
)
assert exc.value.status_code == expected
if expected == 401:
assert "resource_metadata=" in exc.value.headers["www-authenticate"]
probe.assert_awaited_once_with(upstream.url, "Bearer connection-provider-token")
vault.assert_not_awaited()

View file

@ -14075,3 +14075,111 @@ async def test_request_selected_during_guardrail_runs_concurrently_with_tool(mon
assert guardrail_started.is_set() is selected
assert result.is_error is False
assert result.content[0].text == "executed"
@pytest.mark.asyncio
async def test_connection_grants_follow_current_message_and_never_leak_to_another_server():
from types import SimpleNamespace
from datetime import timezone
from pydantic import SecretStr
from starlette.requests import Request
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import CONNECTION_SCOPE_KEY
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
from litellm.proxy._experimental.mcp_server.outbound_credentials import UpstreamCredentialProvider
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
ConnectionBinding,
ConnectionCredential,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken
store = SimpleNamespace(fetch=AsyncMock(return_value=OAuthToken(access_token="saved-vault-token")))
manager = MCPServerManager(cred_provider=UpstreamCredentialProvider(oauth_token_store=store))
server = MCPServer(
server_id="connection-target",
name="connection-target",
url="https://mcp.example/mcp",
transport="http",
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
)
other = server.model_copy(update={"server_id": "other-target"})
binding = ConnectionBinding(key_hash="same-key", server_id=server.server_id, resource="https://gateway.example/mcp")
for token in ("connection-one", "connection-two"):
credential = ConnectionCredential(
kind="connection_access",
binding=binding,
client_id="client",
token=SecretStr(token),
jti=token,
exp=int(datetime.now(timezone.utc).timestamp()) + 300,
)
request = Request(
{"type": "http", "method": "POST", "path": "/mcp", "headers": [], CONNECTION_SCOPE_KEY: credential}
)
reset = active_mcp_request_ctx_var.set(SimpleNamespace(request=request))
try:
client = await manager._create_mcp_client(server)
sent = await client.prepare_request_auth()
assert sent.headers["authorization"] == f"Bearer {token}"
store.fetch.assert_not_awaited()
finally:
active_mcp_request_ctx_var.reset(reset)
reset = active_mcp_request_ctx_var.set(SimpleNamespace(request=request))
try:
other_client = await manager._create_mcp_client(other)
other_sent = await other_client.prepare_request_auth()
assert other_sent.headers["authorization"] == "Bearer saved-vault-token"
assert store.fetch.call_args.args[1] == "other-target"
finally:
active_mcp_request_ctx_var.reset(reset)
saved_client = await manager._create_mcp_client(server)
assert (await saved_client.prepare_request_auth()).headers["authorization"] == "Bearer saved-vault-token"
@pytest.mark.asyncio
async def test_connection_expiry_between_admission_and_egress_never_uses_vault():
from types import SimpleNamespace
from pydantic import SecretStr
from starlette.requests import Request
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import CONNECTION_SCOPE_KEY
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
from litellm.proxy._experimental.mcp_server.outbound_credentials import UpstreamCredentialProvider
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
ConnectionBinding,
ConnectionCredential,
)
store = SimpleNamespace(fetch=AsyncMock(side_effect=AssertionError("must not fall back to another credential")))
manager = MCPServerManager(cred_provider=UpstreamCredentialProvider(oauth_token_store=store))
server = MCPServer(
server_id="expired-connection",
name="expired-connection",
url="https://mcp.example/mcp",
transport="http",
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
)
credential = ConnectionCredential(
kind="connection_access",
binding=ConnectionBinding(key_hash="key", server_id=server.server_id, resource="https://gateway.example/mcp"),
client_id="client",
token=SecretStr("expired-provider-token"),
jti="expired",
exp=1,
)
request = Request(
{"type": "http", "method": "POST", "path": "/mcp", "headers": [], CONNECTION_SCOPE_KEY: credential}
)
reset = active_mcp_request_ctx_var.set(SimpleNamespace(request=request))
try:
with pytest.raises(HTTPException) as exc:
await manager._create_mcp_client(server)
assert exc.value.status_code == 401
assert "expired" in exc.value.detail
store.fetch.assert_not_awaited()
finally:
active_mcp_request_ctx_var.reset(reset)