mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): require challenged upstream consent before completing OAuth
This commit is contained in:
parent
cb6050b521
commit
3a8ed9773b
6 changed files with 175 additions and 6 deletions
|
|
@ -280,6 +280,7 @@ def _gateway_dcr_challenge(
|
|||
route: str,
|
||||
mcp_servers: list[str] | None,
|
||||
invalid_token: bool,
|
||||
oauth_scope: str | None = None,
|
||||
) -> HTTPException:
|
||||
"""The RFC 9728 challenge pointing the client at the protected-resource metadata
|
||||
matching the scope it requested: the per-server document (same URL spelling the
|
||||
|
|
@ -298,13 +299,14 @@ def _gateway_dcr_challenge(
|
|||
else f"{get_request_base_url(request)}/.well-known/oauth-protected-resource{well_known_root_suffix()}/mcp"
|
||||
)
|
||||
error_attr: Final = 'error="invalid_token", ' if invalid_token else ""
|
||||
scope_attr: Final = f', scope="{oauth_scope}"' if oauth_scope else ""
|
||||
return HTTPException(
|
||||
status_code=401,
|
||||
detail={
|
||||
"error": "authentication_required",
|
||||
"message": "Authenticate with the gateway to use the MCP endpoint.",
|
||||
},
|
||||
headers={"WWW-Authenticate": f'Bearer {error_attr}resource_metadata="{resource_metadata_url}"'},
|
||||
headers={"WWW-Authenticate": f'Bearer {error_attr}resource_metadata="{resource_metadata_url}"{scope_attr}'},
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -108,6 +108,7 @@ handle carried in the connect-page URL (the same handle-plus-cookie pattern as t
|
|||
``mcp_oauth_state_`` upstream relay, for the same reasons: replica-safe with no
|
||||
server-side session store, and the sealed value never appears in a URL)."""
|
||||
|
||||
UPSTREAM_AUTHORIZATION_SCOPE_PREFIX: Final = "litellm:mcp:connect:"
|
||||
CONNECT_FLOW_TTL_SECONDS: Final = 600
|
||||
GATEWAY_AUTH_CODE_TTL_SECONDS: Final = 120
|
||||
MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS: Final = 300
|
||||
|
|
@ -286,6 +287,13 @@ class GatewayDcrClient(BaseModel):
|
|||
iat: int
|
||||
|
||||
|
||||
class _UpstreamAuthorizationRequirement(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
server_id: str = Field(min_length=1)
|
||||
user_id: str | None = None
|
||||
exp: int
|
||||
|
||||
|
||||
class _ConnectFlow(BaseModel):
|
||||
"""One in-flight authorize: the SSO user it belongs to and the client parameters
|
||||
needed to mint the code at the finish step. Sealed into the per-flow cookie. ``jti``
|
||||
|
|
@ -301,6 +309,7 @@ class _ConnectFlow(BaseModel):
|
|||
jti: str = Field(min_length=1)
|
||||
exp: int
|
||||
resource_server_id: str | None = None
|
||||
required_upstream_server_id: str | None = None
|
||||
audience: SessionAudience | None = None
|
||||
|
||||
|
||||
|
|
@ -502,6 +511,17 @@ def resolve_scoped_resource_server(request: Request, resource: str | None) -> MC
|
|||
return server
|
||||
|
||||
|
||||
def upstream_authorization_scope(server_id: str, user_id: str | None) -> str:
|
||||
return _seal(
|
||||
UPSTREAM_AUTHORIZATION_SCOPE_PREFIX,
|
||||
_UpstreamAuthorizationRequirement(
|
||||
server_id=server_id,
|
||||
user_id=user_id,
|
||||
exp=int(datetime.now(timezone.utc).timestamp()) + CONNECT_FLOW_TTL_SECONDS,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def aggregate_authorize(
|
||||
request: Request,
|
||||
client_id: str,
|
||||
|
|
@ -537,6 +557,30 @@ def aggregate_authorize(
|
|||
if session_user_id is None:
|
||||
return _login_redirect(base_url, request)
|
||||
scoped_server: Final = resolve_scoped_resource_server(request, resource)
|
||||
requested: Final = tuple(
|
||||
value
|
||||
for value in request.query_params.get("scope", "").split()
|
||||
if value.startswith(UPSTREAM_AUTHORIZATION_SCOPE_PREFIX)
|
||||
)
|
||||
if len(requested) > 1:
|
||||
return _oauth_error(400, "invalid_scope", "only one upstream authorization requirement is supported")
|
||||
requirement: Final = (
|
||||
_open_sealed(
|
||||
requested[0],
|
||||
UPSTREAM_AUTHORIZATION_SCOPE_PREFIX,
|
||||
_UpstreamAuthorizationRequirement,
|
||||
"mcp_upstream_authorization",
|
||||
)
|
||||
if requested
|
||||
else None
|
||||
)
|
||||
if requested and (requirement is None or datetime.now(timezone.utc).timestamp() >= requirement.exp):
|
||||
return _oauth_error(400, "invalid_scope", "invalid or expired upstream authorization requirement; reconnect")
|
||||
if requirement is not None:
|
||||
if requirement.user_id is not None and requirement.user_id != session_user_id:
|
||||
return _oauth_error(403, "access_denied", "sign in as the user that requested this MCP connection")
|
||||
if scoped_server is not None and scoped_server.server_id != requirement.server_id:
|
||||
return _oauth_error(400, "invalid_scope", "upstream authorization requirement does not match the resource")
|
||||
handle: Final = secrets.token_urlsafe(24)
|
||||
flow: Final = _new_connect_flow(
|
||||
session_user_id=session_user_id,
|
||||
|
|
@ -546,6 +590,7 @@ def aggregate_authorize(
|
|||
code_challenge=code_challenge or "",
|
||||
resource_server_id=scoped_server.server_id if scoped_server is not None else None,
|
||||
audience=None,
|
||||
required_upstream_server_id=requirement.server_id if requirement is not None else None,
|
||||
)
|
||||
connect_url: Final = _append_query_params(f"{base_url}/ui/connect", (("connect_flow", handle),))
|
||||
response: Final = RedirectResponse(connect_url, status_code=303)
|
||||
|
|
@ -705,6 +750,7 @@ def _new_connect_flow(
|
|||
code_challenge: str,
|
||||
resource_server_id: str | None,
|
||||
audience: SessionAudience | None,
|
||||
required_upstream_server_id: str | None = None,
|
||||
) -> _ConnectFlow:
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
return _ConnectFlow(
|
||||
|
|
@ -716,6 +762,7 @@ def _new_connect_flow(
|
|||
jti=secrets.token_urlsafe(24),
|
||||
exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS,
|
||||
resource_server_id=resource_server_id,
|
||||
required_upstream_server_id=required_upstream_server_id,
|
||||
audience=audience,
|
||||
)
|
||||
|
||||
|
|
@ -773,14 +820,15 @@ def _open_flow_for(
|
|||
async def _flow_target(
|
||||
flow: _ConnectFlow, lookup_server_reachability: LookupServerReachability
|
||||
) -> tuple[Literal["unscoped", "interactive", "m2m", "stale"], MCPServer | None]:
|
||||
if flow.resource_server_id is None:
|
||||
target_id: Final = flow.required_upstream_server_id or flow.resource_server_id
|
||||
if target_id is None:
|
||||
return "unscoped", None
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # import cycle
|
||||
MCPServerManager,
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
server: Final = global_mcp_server_manager.get_mcp_server_by_id(flow.resource_server_id)
|
||||
server: Final = global_mcp_server_manager.get_mcp_server_by_id(target_id)
|
||||
if (
|
||||
server is None
|
||||
or not (server.is_gateway_managed_oauth2 or server.advertises_gateway_authorization_server)
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
|||
from litellm.proxy._experimental.mcp_server.exceptions import (
|
||||
MCPUpstreamAuthError,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import upstream_authorization_scope
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import (
|
||||
_mcp_active_toolset_id,
|
||||
_mcp_gateway_initialize_instructions,
|
||||
|
|
@ -1728,7 +1729,13 @@ if MCP_AVAILABLE:
|
|||
if results and all(isinstance(result, HTTPException) and result.status_code == 401 for result in results):
|
||||
if all(server.is_gateway_managed_oauth2 for server in eligible):
|
||||
raise _gateway_dcr_challenge(
|
||||
StarletteRequest(scope), get_route_relative_request_path(scope), None, invalid_token=False
|
||||
StarletteRequest(scope),
|
||||
get_route_relative_request_path(scope),
|
||||
None,
|
||||
invalid_token=False,
|
||||
oauth_scope=upstream_authorization_scope(
|
||||
eligible[0].server_id, user_api_key_auth.user_id if user_api_key_auth is not None else None
|
||||
),
|
||||
)
|
||||
first: Final = results[0]
|
||||
if isinstance(first, HTTPException):
|
||||
|
|
|
|||
|
|
@ -2354,3 +2354,113 @@ async def test_token_exchange_relays_a_mint_refusal(failure, status, error):
|
|||
response = await _exchange_native(client_id, _Minter(failure), _Exchanger())
|
||||
assert response.status_code == status
|
||||
assert json.loads(response.body)["error"] == error
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unified_challenge_requires_upstream_consent_without_narrowing_gateway_token(monkeypatch) -> None:
|
||||
from unittest.mock import AsyncMock
|
||||
from urllib.parse import urlencode
|
||||
from fastapi import HTTPException
|
||||
from litellm.proxy._experimental.mcp_server import operations, server
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
github = _scoped_mcp_server(oauth2_flow="authorization_code")
|
||||
manager = operations.global_mcp_server_manager
|
||||
monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[github]))
|
||||
monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: github)
|
||||
monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", AsyncMock(return_value=github))
|
||||
monkeypatch.setattr(manager, "has_user_oauth_token", AsyncMock(return_value=False))
|
||||
with pytest.raises(HTTPException) as challenged:
|
||||
await server._raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope=_request("/mcp").scope, mcp_servers=None, oauth2_headers=None,
|
||||
mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(user_id="u1", api_key="test-key"),
|
||||
client_ip=None,
|
||||
)
|
||||
headers = {k.lower(): v for k, v in (challenged.value.headers or {}).items()}
|
||||
requested = re.search(r'scope="([^"]+)"', headers["www-authenticate"])
|
||||
assert requested is not None, "The challenge must carry the upstream authorization requirement"
|
||||
client_id = (await _register([REDIRECT_URI]))["client_id"]
|
||||
response = aggregate_authorize(
|
||||
request=_request("/authorize/mcp-session", query=urlencode({"scope": requested.group(1)})),
|
||||
client_id=client_id, redirect_uri=REDIRECT_URI, state="client-state", code_challenge=CODE_CHALLENGE,
|
||||
code_challenge_method="S256", response_type="code", session_user_id="u1",
|
||||
resource="https://llm.example.com/mcp",
|
||||
)
|
||||
assert response.status_code == 303
|
||||
described = await _describe_page(response, scoped_server=github, vendor=_VendorCredential("absent"))
|
||||
assert json.loads(described.body) == {
|
||||
"state": "interactive", "client_origin": "https://claude.ai",
|
||||
"server_id": "github-id", "server_name": "github", "connected": False,
|
||||
}
|
||||
cache = DualCache()
|
||||
premature = await _complete_page(response, scoped_server=github, vendor=_VendorCredential("absent"), cache=cache)
|
||||
assert premature.status_code == 400
|
||||
assert "location" not in premature.headers
|
||||
vendor = _VendorCredential("present")
|
||||
completed = await _complete_page(response, scoped_server=github, vendor=vendor, cache=cache)
|
||||
assert completed.status_code == 303
|
||||
assert vendor.calls == [("u1", "github-id")]
|
||||
code = parse_qs(urlparse(completed.headers["location"]).query)["code"][0]
|
||||
tokens = await _redeem(code, client_id, resource="https://llm.example.com/mcp")
|
||||
assert tokens.status_code == 200
|
||||
assert _opened_principal(json.loads(tokens.body)).resource_server_id is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("invalid", ("tampered", "expired", "duplicate", "other_user", "other_resource"))
|
||||
async def test_upstream_authorization_requirement_rejects_invalid_binding(invalid: str, monkeypatch) -> None:
|
||||
from urllib.parse import urlencode
|
||||
from litellm.proxy._experimental.mcp_server import gateway_dcr_flow as flow
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
hint = flow.upstream_authorization_scope("github-id", "u2" if invalid == "other_user" else "u1")
|
||||
if invalid == "tampered":
|
||||
hint = flow.UPSTREAM_AUTHORIZATION_SCOPE_PREFIX + "invalid-ciphertext"
|
||||
elif invalid == "expired":
|
||||
hint = flow._seal(flow.UPSTREAM_AUTHORIZATION_SCOPE_PREFIX, flow._UpstreamAuthorizationRequirement(
|
||||
server_id="github-id", user_id="u1", exp=int(datetime.now(timezone.utc).timestamp()) - 1,
|
||||
))
|
||||
elif invalid == "duplicate":
|
||||
hint = f"{hint} {hint}"
|
||||
monkeypatch.setattr(global_mcp_server_manager, "get_mcp_server_by_name", lambda *args, **kwargs: _scoped_mcp_server("other"))
|
||||
client_id = (await _register([REDIRECT_URI]))["client_id"]
|
||||
response = aggregate_authorize(
|
||||
request=_request(query=urlencode({"scope": hint})),
|
||||
client_id=client_id, redirect_uri=REDIRECT_URI, state="client-state", code_challenge=CODE_CHALLENGE,
|
||||
code_challenge_method="S256", response_type="code", session_user_id="u1",
|
||||
resource="https://llm.example.com/mcp/other" if invalid == "other_resource" else "https://llm.example.com/mcp",
|
||||
)
|
||||
assert response.status_code == (403 if invalid == "other_user" else 400)
|
||||
assert json.loads(response.body)["error"] == ("access_denied" if invalid == "other_user" else "invalid_scope")
|
||||
assert "location" not in response.headers
|
||||
assert "set-cookie" not in response.headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("condition", ("deleted", "revoked", "vault_unavailable", "cancelled"))
|
||||
async def test_required_upstream_completion_preserves_failure_and_cancellation_guards(condition: str) -> None:
|
||||
from urllib.parse import urlencode
|
||||
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import upstream_authorization_scope
|
||||
|
||||
client_id = (await _register([REDIRECT_URI]))["client_id"]
|
||||
hint = upstream_authorization_scope("github-id", "u1")
|
||||
response = aggregate_authorize(
|
||||
request=_request(query=urlencode({"scope": hint})), client_id=client_id, redirect_uri=REDIRECT_URI,
|
||||
state="client-state", code_challenge=CODE_CHALLENGE, code_challenge_method="S256",
|
||||
response_type="code", session_user_id="u1", resource="https://llm.example.com/mcp",
|
||||
)
|
||||
assert response.status_code == 303
|
||||
vendor = _VendorCredential("unavailable" if condition == "vault_unavailable" else "absent")
|
||||
completed = await _complete_page(
|
||||
response, scoped_server=None if condition == "deleted" else _scoped_mcp_server(),
|
||||
reachable=_ServerReachability(condition != "revoked"), vendor=vendor,
|
||||
decision="deny" if condition == "cancelled" else None,
|
||||
)
|
||||
if condition == "cancelled":
|
||||
assert completed.status_code == 303
|
||||
assert parse_qs(urlparse(completed.headers["location"]).query) == {"error": ["access_denied"], "state": ["client-state"]}
|
||||
assert vendor.calls == []
|
||||
else:
|
||||
assert completed.status_code == (503 if condition == "vault_unavailable" else 400)
|
||||
assert "location" not in completed.headers
|
||||
assert vendor.calls == ([("u1", "github-id")] if condition == "vault_unavailable" else [])
|
||||
|
|
|
|||
|
|
@ -867,6 +867,7 @@ async def test_initialize_challenges_missing_upstream_credentials_before_creatin
|
|||
from litellm.proxy._experimental.mcp_server import server
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "test-mcp-oauth-signing")
|
||||
github: Final = MCPServer(
|
||||
server_id="github-id", name="github", alias="github", server_name="github",
|
||||
url="https://github.example/mcp", transport=MCPTransport.http,
|
||||
|
|
@ -916,8 +917,8 @@ async def test_initialize_challenges_missing_upstream_credentials_before_creatin
|
|||
else:
|
||||
assert response.headers["www-authenticate"].startswith("Bearer ")
|
||||
if path == "/mcp":
|
||||
assert response.headers["www-authenticate"] == (
|
||||
'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp"'
|
||||
assert response.headers["www-authenticate"].startswith(
|
||||
'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp", scope="litellm:mcp:connect:'
|
||||
)
|
||||
assert "mcp-session-id" not in response.headers
|
||||
finally:
|
||||
|
|
|
|||
|
|
@ -10714,6 +10714,7 @@ async def test_unified_preflight_challenges_only_when_all_authorized_servers_nee
|
|||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import server as server_module
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "test-mcp-oauth-signing")
|
||||
servers: Final = tuple(
|
||||
_make_oauth2_server(f"server-{index}").model_copy(update={"server_id": f"server-{index}"})
|
||||
for index in range(len(token_states))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue