diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index b0640e4f0dd..0b2e42e8e45 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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( diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 64bab0a7832..f596f9af914 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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") diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index e66504af47a..2e53e8c76b3 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -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( + '' + '
{escape(client.client_name or 'MCP client')} wants access to " + f"{escape(server.name)}.
Return address: {escape(redirect_uri)}
Requested permissions: {escape(scopes or 'provider defaults')}
" + "Only approve if you started this connection. You will continue to the provider to sign in.
" + f'", + 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, + ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_context.py b/litellm/proxy/_experimental/mcp_server/mcp_context.py index 11325a9f127..c280bef337e 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_context.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_context.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index b293ab5a206..f041c677a1d 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/envelope.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/envelope.py index df883d5a208..1f87d87164e 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/envelope.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/envelope.py @@ -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() diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 3a9bca926b0..8f8a9fe9380 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 087c5a03498..92be8d411ed 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index b0cda30dfe5..e25932b1a8c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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