From 4eadf92adee832aa1ef3e52af6b66787614eda0d Mon Sep 17 00:00:00 2001 From: tin-berri Date: Sat, 15 Aug 2026 14:56:34 -0700 Subject: [PATCH] feat(mcp): scope gateway session bearers to the RFC 8707 resource (#35045) --- .../mcp_server/auth/user_api_key_auth_mcp.py | 11 +- .../mcp_server/discoverable_endpoints.py | 4 + .../mcp_server/gateway_dcr_flow.py | 92 ++++++- .../mcp_server/mcp_server_manager.py | 22 +- .../_experimental/mcp_server/oauth_utils.py | 4 +- .../outbound_credentials/session_token.py | 16 +- litellm/proxy/_types.py | 8 + .../auth/test_user_api_key_auth_mcp.py | 65 ++++- .../mcp_server/test_gateway_dcr_flow.py | 244 ++++++++++++++++++ .../test_mcp_oauth_passthrough_cold_start.py | 2 +- .../mcp_server/test_mcp_server_manager.py | 63 +++++ 11 files changed, 515 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 95a3806e8ad..d13b39661ad 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -83,7 +83,7 @@ class UnloadableEntitlementError(Exception): def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: list[str] | None = None) -> list[str] | None: """Resolve the single MCP server name a cold-start passthrough bypass may target. Delegates parsing to - :meth:`MCPRequestHandler._extract_target_server_names_from_path` so the + :meth:`MCPRequestHandler.extract_target_server_names_from_path` so the names used here always match the names downstream routing uses; returns ``None`` whenever the bypass must not activate (aggregate ``/mcp``, multi-server CSV paths, or any other unrecognized path). @@ -94,7 +94,7 @@ def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: list[str] | header/path mismatch here is a sign of a confused or hostile caller — refuse the cold-start bypass rather than admit anonymously based on the path while the header advertises a stricter, non-passthrough target.""" - servers: Final = MCPRequestHandler._extract_target_server_names_from_path(path) + servers: Final = MCPRequestHandler.extract_target_server_names_from_path(path) if len(servers) != 1: verbose_logger.debug( "MCP cold-start: path %r resolved to %r; passthrough 401 bypass " @@ -215,7 +215,7 @@ def _is_gateway_dcr_challenge_scope( return False if _has_client_supplied_mcp_auth(mcp_auth_header, mcp_server_auth_headers): return False - if len(MCPRequestHandler._extract_target_server_names_from_path(route)) == 0: + if len(MCPRequestHandler.extract_target_server_names_from_path(route)) == 0: return True return _gateway_dcr_challenge_target(route, mcp_servers, client_ip) is not None @@ -579,7 +579,7 @@ class MCPRequestHandler: return oauth2_headers, raw_headers, mcp_auth_header, mcp_server_auth_headers @staticmethod - def _extract_target_server_names_from_path(path: str) -> list[str]: + def extract_target_server_names_from_path(path: str) -> list[str]: """ Extract the target MCP server name(s) from the standard MCP transport URL patterns: ``/mcp/{server_name_or_csv}[/...]`` and @@ -836,6 +836,7 @@ class MCPRequestHandler: case SessionBearerAdmitted(): try: admitted: Final = await MCPRequestHandler._reload_admitted_user(result.principal.user_id) + admitted.mcp_session_resource_server_id = result.principal.resource_server_id await MCPRequestHandler._enforce_admitted_live_policy( admitted=admitted, request=request, route=route ) @@ -1168,7 +1169,7 @@ class MCPRequestHandler: (header/path TOCTOU). For non-``/mcp/...`` paths (where the path does not encode targets), fall back to the header. """ - path_targets: Final = MCPRequestHandler._extract_target_server_names_from_path(path) + path_targets: Final = MCPRequestHandler.extract_target_server_names_from_path(path) if path_targets: return path_targets # Path did not resolve to /mcp/... targets — trust the header diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 693e3f8e47d..86e97b55a8e 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1655,6 +1655,7 @@ async def authorize( code_challenge_method: str | None = None, response_type: str | None = None, scope: str | None = None, + resource: str | None = None, ): # Redirect to real OAuth provider with PKCE support from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -1671,6 +1672,7 @@ async def authorize( code_challenge_method=code_challenge_method, response_type=response_type, session_user_id=_session_cookie_user_id(request), + resource=resource, ) lookup_name: Final[str | None] = mcp_server_name or client_id @@ -1721,6 +1723,7 @@ async def token_endpoint( code_verifier: str = Form(None), refresh_token: str | None = Form(None), scope: str | None = Form(None), + resource: str | None = Form(None), mcp_server_name: str | None = None, ): """ @@ -1753,6 +1756,7 @@ async def token_endpoint( master_key=master_key, reload_user=_reload_active_user_by_id, cache=user_api_key_cache, + resource=resource, ) lookup_name: Final = mcp_server_name or client_id diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index 4c1b78c754a..85885fc75f5 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -56,6 +56,8 @@ from litellm._logging import verbose_logger from litellm.caching.caching import DualCache from litellm.proxy._experimental.mcp_server.oauth_utils import ( TOKEN_NO_CACHE_HEADERS, + canonical_resource_uri, + canonicalize_url_identity, get_request_base_url, is_loopback_redirect_host, validate_redirect_uri_shape, @@ -77,6 +79,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) +from litellm.types.mcp_server.mcp_server_manager import MCPServer GATEWAY_DCR_CLIENT_ID_PREFIX: Final = "llm_dcrc_" """Marker prefix on every gateway-issued DCR client_id so the root authorize/token @@ -169,6 +172,7 @@ class _ConnectFlow(BaseModel): code_challenge: str = Field(min_length=1) jti: str = Field(min_length=1) exp: int + resource_server_id: str | None = None class _GatewayAuthCode(BaseModel): @@ -185,6 +189,7 @@ class _GatewayAuthCode(BaseModel): jti: str = Field(min_length=1) iat: int exp: int + resource_server_id: str | None = None def is_gateway_dcr_client_id(client_id: str | None) -> bool: @@ -204,7 +209,13 @@ def _oauth_error(status_code: int, error: str, description: str) -> JSONResponse def _seal(prefix: str, payload: BaseModel) -> str: - return prefix + encrypt_value_helper(payload.model_dump_json()) + """Serialized ``exclude_none`` for the same reason session JWTs are minted that way: an + optional claim that is unset never reaches the wire, so during a rolling deploy a blob + sealed by a new pod without the new claim set stays byte-compatible with predating pods + whose strict models forbid unknown keys. This holds for every sealed artifact and every + future optional claim by construction; it requires each optional field to default to + ``None`` so reopening restores exactly what was sealed.""" + return prefix + encrypt_value_helper(payload.model_dump_json(exclude_none=True)) _SealedModelT = TypeVar("_SealedModelT", bound=BaseModel) @@ -320,6 +331,44 @@ def relative_request_url(request: Request) -> str: return f"{path}?{request.url.query}" if request.url.query else path +def resolve_scoped_resource_server(request: Request, resource: str | None) -> MCPServer | None: + """Resolve an RFC 8707 ``resource`` value to the single gateway-managed oauth2 server it + names, or ``None`` for every other shape: absent, the aggregate resource, a foreign + host, an unparseable value, a multi-server path, an unknown name, or any server mode the + keyless gateway flow does not serve (whose protected-resource metadata never directs a + client here). ``None`` means the flow stays unscoped and byte-identical to today, so a + hostile or confused ``resource`` can never widen anything; a resolved server only ever + NARROWS the session via the sealed scope. + + Resolution is an IDENTITY question, deliberately free of the per-IP visibility filter: + access is enforced where it belongs (grant intersection at admission, IP checks on the + MCP routes), while filtering here would mint an entitlement-wide UNSCOPED bearer exactly + when the caller asked to narrow, and would let authorize-time vs token-time IP drift + turn a matching redemption into a spurious ``invalid_target``.""" + if resource is None: + return None + canonical: Final = canonical_resource_uri(resource) + if canonical is None: + return None + base: Final = canonicalize_url_identity(get_request_base_url(request)) + if canonical == f"{base}/mcp" or not canonical.startswith(f"{base}/"): + return None + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( # noqa: PLC0415 # proxy import cycle + MCPRequestHandler, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # proxy import cycle + global_mcp_server_manager, + ) + + names: Final = MCPRequestHandler.extract_target_server_names_from_path(canonical[len(base) :]) + if len(names) != 1: + return None + server: Final = global_mcp_server_manager.get_mcp_server_by_name(names[0]) + if server is None or not server.is_gateway_managed_oauth2: + return None + return server + + def aggregate_authorize( request: Request, client_id: str, @@ -329,11 +378,16 @@ def aggregate_authorize( code_challenge_method: str | None, response_type: str | None, session_user_id: str | None, + resource: str | None = None, ) -> Response: """The aggregate authorize verb: validate the client, require S256 PKCE, interpose LiteLLM sign-in, and hand the browser to the connect page with the flow sealed into a per-flow cookie. + A per-server RFC 8707 ``resource`` naming a gateway-managed oauth2 server scopes the + flow to that one server: the scope is sealed into the flow, carried into the code, and + bound into the session token, while the connect page interlude runs exactly as before. + Validation failures respond directly with 400 and never redirect: per RFC 6749 section 4.1.2.1 an unvalidated redirect URI must not receive an error redirect, and once the client is at fault there is no trusted place to send the browser. @@ -358,6 +412,7 @@ def aggregate_authorize( login_url: Final = f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}" return RedirectResponse(login_url, status_code=303) now: Final = datetime.now(timezone.utc) + scoped_server: Final = resolve_scoped_resource_server(request, resource) handle: Final = secrets.token_urlsafe(24) flow: Final = _ConnectFlow( user_id=session_user_id, @@ -367,6 +422,7 @@ def aggregate_authorize( code_challenge=code_challenge, jti=secrets.token_urlsafe(24), exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS, + resource_server_id=scoped_server.server_id if scoped_server is not None else None, ) connect_url: Final = _append_query_params( f"{base_url}/ui/connect", @@ -455,6 +511,7 @@ async def complete_connect_flow( jti=secrets.token_urlsafe(24), iat=int(now.timestamp()), exp=int(now.timestamp()) + code_ttl, + resource_server_id=flow.resource_server_id, ), ) params: Final = {"code": code, **({"state": flow.state} if flow.state else {})} @@ -587,6 +644,20 @@ def _reload_failure_response(failure: ReloadUserFailure) -> Response: assert_never(failure) +def _resource_conflicts_with_scope( + request: Request, resource: str | None, sealed_resource_server_id: str | None +) -> bool: + """True when a scoped grant is being redeemed for a DIFFERENT resource than the one + sealed into it (RFC 8707 section 2.2: reject with ``invalid_target``). An absent + ``resource`` never conflicts (the sealed scope still binds the minted session), and an + unscoped grant ignores the parameter entirely, exactly as the endpoint always has, so + no pre-existing client breaks.""" + if sealed_resource_server_id is None or resource is None: + return False + resolved: Final = resolve_scoped_resource_server(request, resource) + return resolved is None or resolved.server_id != sealed_resource_server_id + + async def aggregate_token( request: Request, grant_type: str, @@ -598,6 +669,7 @@ async def aggregate_token( master_key: str | None, reload_user: ReloadUser, cache: DualCache, + resource: str | None = None, ) -> Response: """The aggregate token verb: authorization_code and refresh_token grants for the identity-only session pair. Every path re-validates the litellm user live before @@ -609,10 +681,12 @@ async def aggregate_token( now: Final = datetime.now(timezone.utc) if grant_type == "authorization_code": return await _authorization_code_grant( + request=request, code=code, redirect_uri=redirect_uri, client_id=client_id, code_verifier=code_verifier, + resource=resource, keys=keys, now=now, reload_user=reload_user, @@ -620,8 +694,10 @@ async def aggregate_token( ) if grant_type == "refresh_token": return await _refresh_token_grant( + request=request, refresh_token=refresh_token, client_id=client_id, + resource=resource, keys=keys, now=now, reload_user=reload_user, @@ -631,10 +707,12 @@ async def aggregate_token( async def _authorization_code_grant( + request: Request, code: str | None, redirect_uri: str | None, client_id: str, code_verifier: str | None, + resource: str | None, keys: SessionKeys, now: datetime, reload_user: ReloadUser, @@ -651,6 +729,8 @@ async def _authorization_code_grant( return _oauth_error(400, "invalid_grant", "the authorization code has expired") if client_id != parsed.client_id or redirect_uri != parsed.redirect_uri: return _oauth_error(400, "invalid_grant", "the authorization code was issued to a different client") + if _resource_conflicts_with_scope(request, resource, parsed.resource_server_id): + return _oauth_error(400, "invalid_target", "resource does not match the scope this code was issued for") if not _pkce_verifier_matches(code_verifier, parsed.code_challenge): return _oauth_error(400, "invalid_grant", "PKCE verification failed") # Revalidate the user BEFORE claiming the code, so a transient DB outage (a retryable @@ -666,12 +746,18 @@ async def _authorization_code_grant( parsed.exp - int(now.timestamp()) + _CLAIM_TTL_BUFFER_SECONDS, ): return _oauth_error(400, "invalid_grant", "the authorization code was already used") - return _session_token_pair(SessionPrincipal(user_id=parsed.user_id, client_id=client_id), keys, now) + return _session_token_pair( + SessionPrincipal(user_id=parsed.user_id, client_id=client_id, resource_server_id=parsed.resource_server_id), + keys, + now, + ) async def _refresh_token_grant( + request: Request, refresh_token: str | None, client_id: str, + resource: str | None, keys: SessionKeys, now: datetime, reload_user: ReloadUser, @@ -682,6 +768,8 @@ async def _refresh_token_grant( opened: Final = open_session_refresh_bearer(refresh_token, keys, now, expected_client_id=client_id) if not isinstance(opened, SessionRefreshOpened): return _oauth_error(400, "invalid_grant", "the refresh token is invalid for this client") + if _resource_conflicts_with_scope(request, resource, opened.principal.resource_server_id): + return _oauth_error(400, "invalid_target", "resource does not match the scope this token was issued for") failure: Final = await reload_user(opened.principal.user_id) if failure is not None: return _reload_failure_response(failure) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 4f94f94acb5..c782f0dfa09 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2491,6 +2491,18 @@ class MCPServerManager: open_ids.update(submitted_server_ids) return open_ids + @staticmethod + def _admitted_session_resource_scope(user_api_key_auth: UserAPIKeyAuth | None) -> str | None: + """The single server an admitted session subject's bearer was scoped to at authorize + time (RFC 8707 resource), or None for every other principal shape and for unscoped + sessions. Read at every return path of :meth:`get_allowed_mcp_servers`, including + the exception fallback, and applied AFTER every union (grants, operator-open, + submitted) because the scope is a ceiling over the whole reachable set; a resolver + fault therefore never widens a scoped bearer to the allow-all set.""" + if user_api_key_auth is None or not _is_mcp_admitted_user_subject(user_api_key_auth): + return None + return user_api_key_auth.mcp_session_resource_server_id + async def get_allowed_mcp_servers(self, user_api_key_auth: UserAPIKeyAuth | None = None) -> list[str]: """ Get the allowed MCP Servers for the user. @@ -2600,13 +2612,19 @@ class MCPServerManager: if len(combined_servers) == 0: verbose_logger.debug("No allowed MCP Servers found for user api key auth.") - return list(combined_servers) + scope = MCPServerManager._admitted_session_resource_scope(user_api_key_auth) + return [server_id for server_id in combined_servers if scope is None or server_id == scope] except Exception: # noqa: BLE001 verbose_logger.exception( "Failed to get allowed MCP servers; team-level object_permission " "grants may be dropped. Falling back to global and submitted servers." ) - return list(dict.fromkeys(allow_all_server_ids + submitted_server_ids)) + scope = MCPServerManager._admitted_session_resource_scope(user_api_key_auth) + return [ + server_id + for server_id in dict.fromkeys(allow_all_server_ids + submitted_server_ids) + if scope is None or server_id == scope + ] async def resolve_toolset_tool_permissions( self, diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 84b40b72258..a30b5ee9e49 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -633,7 +633,7 @@ def canonicalize_url_identity(url: str) -> str: return urlunparse((scheme, netloc, parsed.path.rstrip("/"), "", "", "")) -def _canonical_resource_uri(url: str) -> str | None: +def canonical_resource_uri(url: str) -> str | None: """Canonicalize an upstream MCP server URL into an RFC 8707 resource identifier. Keeps only the scheme, host, port and path, which is the shape the MCP authorization spec's @@ -693,7 +693,7 @@ def resolve_upstream_resource(mcp_server: "MCPServer") -> str | None: mcp_server.server_id, ) return None - canonical: Final = _canonical_resource_uri(mcp_server.url) + canonical: Final = canonical_resource_uri(mcp_server.url) if canonical is None: verbose_logger.warning( "MCP server %s sets upstream_resource=auto but its url is not an absolute URI, so no " diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py index 3ef7327cda3..15f5f82c4b6 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py @@ -85,11 +85,18 @@ class SessionPrincipal(BaseModel): enforced at use time rather than frozen at mint time. ``client_id`` is the (stateless, gateway-sealed) DCR client identifier the token was issued to; the token endpoint requires it to match on the refresh grant. + + ``resource_server_id`` is the single MCP server this session was authorized for when + the client requested a per-server RFC 8707 resource at authorize time, or ``None`` for + the aggregate scope. It is a RESTRICTION carried for admission to intersect against + the live grant resolution, never a grant by itself; the refresh grant re-mints from + this principal so the restriction survives rotation. """ model_config = ConfigDict(frozen=True) user_id: str = Field(min_length=1) client_id: str = Field(min_length=1) + resource_server_id: str | None = None class SessionKeys(BaseModel): @@ -186,6 +193,7 @@ class _SessionClaims(BaseModel): kind: SessionTokenKind user_id: str = Field(min_length=1) client_id: str = Field(min_length=1) + resource_server_id: str | None = None def is_session_token(candidate: str) -> bool: @@ -286,9 +294,10 @@ def _mint( kind=kind, user_id=principal.user_id, client_id=principal.client_id, + resource_server_id=principal.resource_server_id, ) token: Final = prefix + jwt.encode( - claims.model_dump(), keys.signing_key.get_secret_value(), algorithm=_SESSION_JWT_ALGORITHM + claims.model_dump(exclude_none=True), keys.signing_key.get_secret_value(), algorithm=_SESSION_JWT_ALGORITHM ) size_bytes: Final = len(token.encode("utf-8")) if size_bytes > MAX_SESSION_TOKEN_BYTES: @@ -323,7 +332,10 @@ def _open( if now.timestamp() >= claims.exp: return SessionExpired() return OpenedSessionToken( - principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id), jti=claims.jti + principal=SessionPrincipal( + user_id=claims.user_id, client_id=claims.client_id, resource_server_id=claims.resource_server_id + ), + jti=claims.jti, ) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 934ac9ac3d7..a566d491597 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2762,6 +2762,13 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # key off. Server-only and stripped from validated input for the same reason as the marker # above: a forged entry would let a caller pick which team's rpm bucket it is charged against. mcp_source_team_rpm_limits: dict[str, dict[str, int]] | None = Field(default=None, exclude=True) + # The single MCP server_id a gateway session bearer was scoped to at authorize time (RFC 8707 + # resource), or None for an aggregate-scope session. A RESTRICTION intersected against the live + # grant resolution, never a grant. Server-only, set exclusively by the MCP gateway admission + # path via post-construction assignment and stripped from validated input like the markers + # above; a forged value could at most narrow, but the stripping keeps the field's provenance + # single-owner so its meaning stays trustworthy. + mcp_session_resource_server_id: str | None = Field(default=None, exclude=True) via_virtual_key: bool = Field( default=False, exclude=True, @@ -2798,6 +2805,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # kwargs, model_validate, a JWT/key claim splat) so it can never be forged from caller data. values.pop("mcp_admitted_user_subject", None) values.pop("mcp_source_team_rpm_limits", None) + values.pop("mcp_session_resource_server_id", None) values.pop("via_virtual_key", None) if values.get("api_key") is not None: values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))}) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 52dc91ce24d..0209abee510 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -2646,7 +2646,7 @@ class TestMCPDelegateAuthToUpstream: def test_extract_target_server_names_matches_routing_parser(self): """ - Regression: _extract_target_server_names_from_path must match the + Regression: extract_target_server_names_from_path must match the downstream regex parser in server.py::_get_mcp_servers_in_path. Previously, a request to ``/mcp//garbage`` was parsed as @@ -2682,7 +2682,7 @@ class TestMCPDelegateAuthToUpstream: ("/", []), ] for path_input, expected in cases: - assert MCPRequestHandler._extract_target_server_names_from_path(path_input) == expected, ( + assert MCPRequestHandler.extract_target_server_names_from_path(path_input) == expected, ( f"path={path_input!r} → expected {expected!r}" ) assert (_get_mcp_servers_in_path(path_input) or []) == expected, ( @@ -8365,3 +8365,64 @@ class TestEntitlementFaultSemantics: ): allowed = await MCPRequestHandler.get_allowed_mcp_servers(auth) assert set(allowed) == {"srv1"} + + +@pytest.mark.asyncio +class TestScopedSessionAdmission: + """LIT-4917: a session bearer sealed to one server (RFC 8707 resource at authorize) + carries that scope onto the admitted auth object, where the grant resolution intersects + it fail closed; an unscoped bearer carries None and is byte-identical to before.""" + + _MASTER_KEY = "sk-scoped-session-admission-master-key" + + def _bearer(self, resource_server_id): + from datetime import datetime, timezone + + from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + session_keys_from_master_key, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import ( + SessionPrincipal, + mint_session_token, + ) + + keys = session_keys_from_master_key(self._MASTER_KEY) + principal = SessionPrincipal( + user_id="scoped-user", client_id="llm_dcrc_abc", resource_server_id=resource_server_id + ) + return mint_session_token(principal, keys, datetime(2030, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() + + @pytest.mark.parametrize("scope", ["github-server-id", None]) + async def test_admission_carries_sealed_resource_scope(self, scope): + token = self._bearer(scope) + scope_dict = { + "type": "http", + "method": "POST", + "path": "/mcp/github", + "headers": [(b"host", b"testserver"), (b"authorization", f"Bearer {token}".encode())], + } + get_user_object = AsyncMock( + return_value=MagicMock( + user_id="scoped-user", + organization_id=None, + metadata={"scim_active": True}, + user_role=None, + object_permission=None, + object_permission_id=None, + tpm_limit=None, + rpm_limit=None, + ) + ) + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + ): + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope_dict) + assert auth_result.mcp_admitted_user_subject is True + assert auth_result.mcp_session_resource_server_id == scope + + 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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py index d8eab3ecb3b..cc65970a180 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py @@ -795,3 +795,247 @@ async def test_manual_delivery_page_renders_the_url_as_data_never_as_a_shell_com assert 'curl "' not in body assert "curl '" not in body assert 'value="' in body + + +def _scoped_mcp_server(name="github", **kw): + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + return MCPServer( + server_id=f"{name}-id", + name=name, + server_name=name, + alias=name, + url="https://upstream.example/mcp", + transport="http", + auth_type=MCPAuth.oauth2, + **kw, + ) + + +SCOPED_RESOURCE = "https://llm.example.com/mcp/github" +_MANAGER_PATCH = "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + + +def _scoped_authorize(client_id, resource, session_user_id="u1"): + return aggregate_authorize( + request=_request(query=f"client_id={client_id}"), + client_id=client_id, + redirect_uri=REDIRECT_URI, + state="client-state-123", + code_challenge=CODE_CHALLENGE, + code_challenge_method="S256", + response_type="code", + session_user_id=session_user_id, + resource=resource, + ) + + +async def _redeem(code, client_id, cache=None, **overrides): + arguments = { + "request": _request("/token", method="POST"), + "grant_type": "authorization_code", + "code": code, + "redirect_uri": REDIRECT_URI, + "client_id": client_id, + "code_verifier": CODE_VERIFIER, + "refresh_token": None, + "master_key": MASTER_KEY, + "reload_user": _reload_user_active, + "cache": cache or DualCache(), + } + return await aggregate_token(**{**arguments, **overrides}) + + +def _opened_principal(payload): + keys = session_keys_from_master_key(MASTER_KEY) + admitted = resolve_session_bearer(f"Bearer {payload['access_token']}", keys, datetime.now(timezone.utc)) + assert isinstance(admitted, SessionBearerAdmitted) + return admitted.principal + + +async def _finish_connect_page(response): + handle, cookies = _flow_cookie_from(response) + completed = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="u1", + cache=DualCache(), + ) + return parse_qs(urlparse(completed.headers["location"]).query)["code"][0] + + +def _sealed_wire_json(sealed, prefix, debug_key): + from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper + + raw = decrypt_value_helper(sealed.removeprefix(prefix), debug_key, return_original_value=False) + assert isinstance(raw, str) + return json.loads(raw) + + +@pytest.mark.asyncio +async def test_scoped_authorize_runs_connect_page_with_sealed_scope(): + """LIT-4917: a per-server RFC 8707 resource naming a gateway-managed oauth2 server + seals that server into the flow. The connect page interlude runs exactly as before + (the scope restricts, it never skips consent), and the code minted at the finish step + and the session pair it redeems for are both scoped.""" + from unittest.mock import patch + + client_id = (await _register([REDIRECT_URI]))["client_id"] + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = _scoped_mcp_server() + response = _scoped_authorize(client_id, SCOPED_RESOURCE) + assert response.status_code == 303 + assert "/ui/connect" in response.headers["location"] + _, cookies = _flow_cookie_from(response) + assert _sealed_wire_json(next(iter(cookies.values())), "", "gateway_connect_flow")["resource_server_id"] == "github-id" + code = await _finish_connect_page(response) + assert _sealed_wire_json(code, GATEWAY_AUTH_CODE_PREFIX, "gateway_authorization_code")["resource_server_id"] == "github-id" + token_response = await _redeem(code, client_id) + assert token_response.status_code == 200 + principal = _opened_principal(json.loads(token_response.body)) + assert principal.resource_server_id == "github-id" + assert principal.user_id == "u1" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "resource, resolves", + [ + (None, False), + ("https://llm.example.com/mcp", False), + ("https://other.example.com/mcp/github", False), + ("https://llm.example.com/mcp/github,linear", False), + ("https://llm.example.com/mcp/unknown", None), + ("not a url", False), + ], +) +async def test_unscoped_resources_leave_flow_and_token_byte_identical(resource, resolves): + """Every resource shape outside 'exactly one gateway-managed server' keeps today's flow: + connect page interlude, and NONE of the minted artifacts carry the scope key on the + wire, not the flow cookie, not the code, not the session JWT, so an unscoped flow + started on a new pod completes on a pod whose strict models predate the claim.""" + import base64 + from unittest.mock import patch + + client_id = (await _register([REDIRECT_URI]))["client_id"] + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = None if resolves is None else _scoped_mcp_server() + response = _scoped_authorize(client_id, resource) + assert response.status_code == 303 + assert "/ui/connect" in response.headers["location"] + _, cookies = _flow_cookie_from(response) + assert "resource_server_id" not in _sealed_wire_json(next(iter(cookies.values())), "", "gateway_connect_flow") + code = await _finish_connect_page(response) + assert "resource_server_id" not in _sealed_wire_json(code, GATEWAY_AUTH_CODE_PREFIX, "gateway_authorization_code") + token_response = await _redeem(code, client_id) + payload = json.loads(token_response.body) + assert _opened_principal(payload).resource_server_id is None + jwt_payload_segment = payload["access_token"].removeprefix("llm_session_").split(".")[1] + claims = json.loads(base64.urlsafe_b64decode(jwt_payload_segment + "=" * (-len(jwt_payload_segment) % 4))) + assert "resource_server_id" not in claims + + +@pytest.mark.asyncio +async def test_scoped_authorize_delegate_server_stays_unscoped(): + """A delegate-auth oauth2 server is outside the gateway-managed set (its keyless flow is + upstream PKCE via the relay), so a resource naming it never scopes the gateway flow.""" + from unittest.mock import patch + + client_id = (await _register([REDIRECT_URI]))["client_id"] + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = _scoped_mcp_server(delegate_auth_to_upstream=True) + response = _scoped_authorize(client_id, SCOPED_RESOURCE) + assert "/ui/connect" in response.headers["location"] + code = await _finish_connect_page(response) + token_response = await _redeem(code, client_id) + assert _opened_principal(json.loads(token_response.body)).resource_server_id is None + + +@pytest.mark.asyncio +async def test_token_rejects_resource_conflicting_with_sealed_scope(): + """RFC 8707 section 2.2: redeeming a scoped code (or rotating a scoped refresh token) + for a DIFFERENT resource fails with invalid_target; an absent resource redeems fine and + the sealed scope still binds the minted pair, surviving refresh rotation.""" + from unittest.mock import patch + + client_id = (await _register([REDIRECT_URI]))["client_id"] + github = _scoped_mcp_server() + linear = _scoped_mcp_server(name="linear") + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = github + response = _scoped_authorize(client_id, SCOPED_RESOURCE) + code = await _finish_connect_page(response) + + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = linear + mismatched = await _redeem(code, client_id, resource="https://llm.example.com/mcp/linear") + assert json.loads(mismatched.body)["error"] == "invalid_target" + + cache = DualCache() + token_response = await _redeem(code, client_id, cache=cache) + payload = json.loads(token_response.body) + assert _opened_principal(payload).resource_server_id == "github-id" + + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = linear + refresh_mismatch = await _redeem( + None, + client_id, + cache=cache, + grant_type="refresh_token", + refresh_token=payload["refresh_token"], + resource="https://llm.example.com/mcp/linear", + ) + assert json.loads(refresh_mismatch.body)["error"] == "invalid_target" + + rotated = await _redeem( + None, client_id, cache=cache, grant_type="refresh_token", refresh_token=payload["refresh_token"] + ) + assert rotated.status_code == 200 + assert _opened_principal(json.loads(rotated.body)).resource_server_id == "github-id" + + +@pytest.mark.asyncio +async def test_resolve_scoped_resource_server_matrix(): + """Unit pin of the resource resolver: both per-server URL spellings resolve; the + aggregate resource, foreign hosts, CSV paths, unknown names, and non-gateway-managed + modes all return None so nothing outside the served set can enter the scoped flow.""" + from unittest.mock import patch + + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import resolve_scoped_resource_server + + request = _request() + github = _scoped_mcp_server() + for resource, resolved_server, expected in [ + ("https://llm.example.com/mcp/github", github, "github-id"), + ("https://llm.example.com/github/mcp", github, "github-id"), + ("https://LLM.example.com/mcp/github/", github, "github-id"), + ("https://llm.example.com/mcp", github, None), + ("https://other.example.com/mcp/github", github, None), + ("https://llm.example.com/mcp/a,b", github, None), + ("https://llm.example.com/mcp/github", None, None), + ("https://llm.example.com/mcp/github", _scoped_mcp_server(delegate_auth_to_upstream=True), None), + (None, github, None), + ]: + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = resolved_server + result = resolve_scoped_resource_server(request, resource) + assert (result.server_id if result is not None else None) == expected, resource + + +@pytest.mark.asyncio +async def test_resource_resolution_is_identity_not_ip_filtered_access(): + """The resolver decides which server a resource NAMES; per-IP visibility filtering + belongs to the MCP routes and grant intersection. Filtering here would mint an + entitlement-wide unscoped bearer exactly when the caller asked to narrow, and IP drift + between authorize and token would turn a matching redemption into invalid_target.""" + from unittest.mock import patch + + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import resolve_scoped_resource_server + + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = _scoped_mcp_server() + result = resolve_scoped_resource_server(_request(), SCOPED_RESOURCE) + assert result is not None + manager.get_mcp_server_by_name.assert_called_once_with("github") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py index 3e934577a66..f25d3baea0a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py @@ -137,7 +137,7 @@ def test_is_mcp_passthrough_cold_start_false_for_empty_servers(): [ ("/mcp/sample_docs", ["sample_docs"]), # Server names may contain at most one slash (mirrors - # ``_extract_target_server_names_from_path``), so when more than two + # ``extract_target_server_names_from_path``), so when more than two # segments follow ``/mcp/`` the first two are treated as the name. ("/mcp/sample_docs/tools/list", ["sample_docs/tools"]), ("/mcp/custom_solutions/user_123", ["custom_solutions/user_123"]), diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 55cc8565a71..99181f0f087 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -9859,3 +9859,66 @@ class TestToolAuthorizationIsNotConditionalOnLogging: ) upstream.assert_awaited_once() + + +class TestSessionResourceScopeIntersect: + """LIT-4917: the sealed session scope intersects the admitted subject's resolved server + set at the single convergence point every fan-out and tool call reads, covering the + exception fallback so a resolver fault never widens a scoped bearer.""" + + def _admitted_auth(self, scope): + from litellm.proxy._types import UserAPIKeyAuth + + auth = UserAPIKeyAuth(user_id="scoped-user") + auth.mcp_admitted_user_subject = True + auth.mcp_session_resource_server_id = scope + return auth + + def test_scope_reader_is_none_for_keys_and_unscoped_subjects(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy._types import UserAPIKeyAuth + + assert MCPServerManager._admitted_session_resource_scope(None) is None + assert MCPServerManager._admitted_session_resource_scope(UserAPIKeyAuth(user_id="u")) is None + assert MCPServerManager._admitted_session_resource_scope(self._admitted_auth(None)) is None + + def test_scope_reader_returns_sealed_scope_for_admitted_subjects(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + assert MCPServerManager._admitted_session_resource_scope(self._admitted_auth("b")) == "b" + + @pytest.mark.asyncio + async def test_get_allowed_mcp_servers_scopes_past_operator_open_union(self): + """The intersect applies AFTER the operator-open (allow_all_keys) union, so a scoped + bearer cannot reach an allow-all server outside its scope, and applies on the + exception fallback so a resolver fault yields the scoped subset of allow-all rather + than the whole set.""" + from unittest.mock import AsyncMock, patch + + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + manager = MCPServerManager() + auth = self._admitted_auth("granted-id") + with ( + patch.object(MCPServerManager, "get_allow_all_keys_server_ids", return_value=["open-id", "granted-id"]), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=["granted-id", "other-id"], + ), + patch.object(MCPServerManager, "_get_active_submitted_mcp_server_ids_for_user", new_callable=AsyncMock, return_value=[]), + ): + allowed = await manager.get_allowed_mcp_servers(auth) + assert allowed == ["granted-id"] + + with ( + patch.object(MCPServerManager, "get_allow_all_keys_server_ids", return_value=["open-id", "granted-id"]), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers", + new_callable=AsyncMock, + side_effect=RuntimeError("resolver down"), + ), + patch.object(MCPServerManager, "_get_active_submitted_mcp_server_ids_for_user", new_callable=AsyncMock, return_value=[]), + ): + fallback = await manager.get_allowed_mcp_servers(auth) + assert fallback == ["granted-id"]