mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(mcp): scope gateway session bearers to the RFC 8707 resource (#35045)
This commit is contained in:
parent
8035fb3d27
commit
4eadf92ade
11 changed files with 515 additions and 16 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"))})
|
||||
|
|
|
|||
|
|
@ -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/<delegated>/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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"]),
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue