mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(mcp): bind keyed OAuth connections to the initiating key
This commit is contained in:
parent
cc1a3157d3
commit
641f50afe2
11 changed files with 1723 additions and 13 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 <untrusted>" 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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue