mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
refactor(mcp): harden gateway DCR flow per adversarial review
- atomic single-use guard (async_increment_cache) + reload-before-claim so a transient DB blip does not burn a valid code - PKCE verify over bytes so a non-ASCII code_challenge fails invalid_grant instead of raising a 500; validate code_verifier length (RFC 7636) - flag-off byte-identical for a server literally named mcp (AS well-known delegates to the named-server document) - connect flow is single-use (atomic jti claim) so a double-submit cannot mint two codes - extra=forbid on the sealed models; bound state length; drop unused request param and coarse dict on register - _reload_failure_response exhaustive match+assert_never; dedupe ReloadUserFailure with _KeyResolutionFailure - reject control/whitespace chars in the same-origin return_to
This commit is contained in:
parent
f546cfcb5f
commit
05c55d016b
5 changed files with 287 additions and 97 deletions
|
|
@ -29,6 +29,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
|
|||
_finish_bridge_mint,
|
||||
_prepare_bridge_mint,
|
||||
_prepare_bridge_refresh,
|
||||
_reload_active_user_by_id,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.faults import (
|
||||
CallerRejected,
|
||||
|
|
@ -1272,7 +1273,7 @@ async def authorize(
|
|||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
if mcp_server_name is None and client_id and is_gateway_dcr_client_id(client_id) and is_mcp_gateway_dcr_enabled():
|
||||
if mcp_server_name is None and client_id and is_gateway_dcr_client_id(client_id):
|
||||
return aggregate_authorize(
|
||||
request=request,
|
||||
client_id=client_id,
|
||||
|
|
@ -1347,7 +1348,7 @@ async def token_endpoint(
|
|||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
if mcp_server_name is None and is_gateway_dcr_client_id(client_id) and is_mcp_gateway_dcr_enabled():
|
||||
if mcp_server_name is None and is_gateway_dcr_client_id(client_id):
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # circular import at module load
|
||||
master_key,
|
||||
user_api_key_cache,
|
||||
|
|
@ -1389,16 +1390,16 @@ async def token_endpoint(
|
|||
|
||||
@router.post("/authorize/complete")
|
||||
async def authorize_complete(request: Request, flow: str = Form(...)):
|
||||
"""Finish an aggregate connect flow (``mcp_gateway_dcr``): mint the gateway
|
||||
authorization code for the signed-in user and redirect back to the DCR client. POST
|
||||
plus the per-flow HttpOnly cookie set at /authorize; 404 when the flag is off so the
|
||||
route is byte-invisible to existing deployments."""
|
||||
if not is_mcp_gateway_dcr_enabled():
|
||||
raise HTTPException(status_code=404, detail="Not Found")
|
||||
return complete_connect_flow(
|
||||
"""Finish an aggregate connect flow: mint the gateway authorization code for the
|
||||
signed-in user and redirect back to the DCR client. POST plus the per-flow HttpOnly
|
||||
cookie set at /authorize; an anonymous or bad-flow request just 400s."""
|
||||
from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 # circular import at module load
|
||||
|
||||
return await complete_connect_flow(
|
||||
request=request,
|
||||
flow_handle=flow,
|
||||
session_user_id=_session_cookie_user_id(request),
|
||||
cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -2126,8 +2127,13 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non
|
|||
}
|
||||
client_ip = IPAddressUtils.get_mcp_client_ip(request)
|
||||
if not mcp_server_name:
|
||||
if is_mcp_gateway_dcr_enabled():
|
||||
return await register_aggregate_client(request=request, request_body=data)
|
||||
# A real DCR request carries redirect_uris (RFC 7591): route it to the aggregate DCR
|
||||
# endpoint the aggregate authorization-server metadata advertises. A single-server
|
||||
# deployment registers at /{server}/register instead (its bare-origin discovery
|
||||
# advertises that), so this does not affect it. A request without redirect_uris is not
|
||||
# a DCR request, so the legacy single-server-or-dummy fallback is kept for it.
|
||||
if data.get("redirect_uris"):
|
||||
return await register_aggregate_client(request_body=data)
|
||||
resolved = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
|
||||
if resolved:
|
||||
return await register_client_with_server(
|
||||
|
|
|
|||
|
|
@ -41,6 +41,7 @@ import hashlib
|
|||
import hmac
|
||||
import secrets
|
||||
from base64 import urlsafe_b64encode
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime, timezone
|
||||
from typing import Awaitable, Callable, Literal, TypeVar
|
||||
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
|
||||
|
|
@ -48,6 +49,7 @@ from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
|
|||
from fastapi import Request
|
||||
from fastapi.responses import JSONResponse, RedirectResponse, Response
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
|
|
@ -89,7 +91,9 @@ server-side session store, and the sealed value never appears in a URL)."""
|
|||
|
||||
CONNECT_FLOW_TTL_SECONDS = 600
|
||||
GATEWAY_AUTH_CODE_TTL_SECONDS = 120
|
||||
_CLAIM_TTL_BUFFER_SECONDS = 60
|
||||
_USED_CODE_CACHE_PREFIX = "mcp_gateway_dcr_code_used:"
|
||||
_USED_FLOW_CACHE_PREFIX = "mcp_gateway_dcr_flow_used:"
|
||||
|
||||
MAX_REDIRECT_URIS = 3
|
||||
MAX_REDIRECT_URI_LENGTH = 256
|
||||
|
|
@ -99,6 +103,22 @@ every session-token claim set: 3 URIs of 256 bytes seal to roughly 1.2KB, comfor
|
|||
under this cap and under the session token's own 4KB ceiling. Claude Desktop and MCP
|
||||
Inspector register one or two redirect URIs."""
|
||||
|
||||
MAX_STATE_LENGTH = 1024
|
||||
"""Bound on the client ``state`` sealed into the flow cookie and echoed on the auth-code
|
||||
redirect. An unbounded ``state`` can push the sealed cookie past the browser's ~4KB cap
|
||||
(silently dropped, breaking the flow); spec clients send a short opaque value."""
|
||||
|
||||
MIN_CODE_VERIFIER_LENGTH = 43
|
||||
MAX_CODE_VERIFIER_LENGTH = 128
|
||||
"""RFC 7636 section 4.1 bounds for the PKCE ``code_verifier``. Enforced so an out-of-range
|
||||
verifier gets a clean ``invalid_request`` instead of an opaque PKCE-mismatch."""
|
||||
|
||||
_UNPREFIXED = ""
|
||||
"""Prefix for a sealed value that carries no wire marker because it is never routed by
|
||||
prefix (the connect flow lives only in its own per-handle cookie, opened by that one
|
||||
handle). Named so the empty-string argument to ``_seal`` / ``_open_sealed`` reads as
|
||||
deliberate rather than a typo."""
|
||||
|
||||
_CLIENT_RECORD_DEBUG_KEY = "gateway_dcr_client"
|
||||
_CONNECT_FLOW_DEBUG_KEY = "gateway_connect_flow"
|
||||
_AUTH_CODE_DEBUG_KEY = "gateway_authorization_code"
|
||||
|
|
@ -111,32 +131,41 @@ else fails the grant closed."""
|
|||
|
||||
|
||||
class GatewayDcrClient(BaseModel):
|
||||
"""The registration record sealed into a gateway DCR ``client_id``."""
|
||||
"""The registration record sealed into a gateway DCR ``client_id``.
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
``extra="forbid"`` so a sealed value of another type (an auth code, a connect flow)
|
||||
that happened to decrypt under the shared key can never validate as a client record:
|
||||
cross-type confusion is rejected at the model boundary, not left to differing required
|
||||
fields."""
|
||||
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
redirect_uris: tuple[str, ...] = Field(min_length=1, max_length=MAX_REDIRECT_URIS)
|
||||
iat: 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."""
|
||||
needed to mint the code at the finish step. Sealed into the per-flow cookie. ``jti``
|
||||
makes the flow single-use at complete; ``extra="forbid"`` rejects cross-type
|
||||
confusion."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
user_id: str = Field(min_length=1)
|
||||
client_id: str = Field(min_length=1)
|
||||
redirect_uri: str = Field(min_length=1)
|
||||
state: str
|
||||
code_challenge: str = Field(min_length=1)
|
||||
jti: str = Field(min_length=1)
|
||||
exp: int
|
||||
|
||||
|
||||
class _GatewayAuthCode(BaseModel):
|
||||
"""The gateway-sealed authorization code: the user consent it represents and the
|
||||
bindings the token endpoint must verify (client, redirect URI, PKCE challenge),
|
||||
plus a ``jti`` for the single-use guard."""
|
||||
plus a ``jti`` for the single-use guard. ``extra="forbid"`` rejects cross-type
|
||||
confusion."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
user_id: str = Field(min_length=1)
|
||||
client_id: str = Field(min_length=1)
|
||||
redirect_uri: str = Field(min_length=1)
|
||||
|
|
@ -149,7 +178,7 @@ class _GatewayAuthCode(BaseModel):
|
|||
def is_gateway_dcr_client_id(client_id: str | None) -> bool:
|
||||
"""Cheap prefix routing test so the root endpoints only enter the aggregate arm for
|
||||
clients this flow registered; every other client_id keeps today's behavior."""
|
||||
return bool(client_id) and str(client_id).startswith(GATEWAY_DCR_CLIENT_ID_PREFIX)
|
||||
return client_id is not None and client_id.startswith(GATEWAY_DCR_CLIENT_ID_PREFIX)
|
||||
|
||||
|
||||
def _oauth_error(status_code: int, error: str, description: str) -> JSONResponse:
|
||||
|
|
@ -200,7 +229,7 @@ def _redirect_uri_acceptable(uri: str) -> bool:
|
|||
return parsed.scheme == "http" and (parsed.hostname or "").lower() in ("localhost", "127.0.0.1", "::1")
|
||||
|
||||
|
||||
async def register_aggregate_client(request: Request, request_body: dict) -> Response:
|
||||
async def register_aggregate_client(request_body: Mapping[str, object]) -> Response:
|
||||
"""RFC 7591 dynamic registration against the gateway itself, statelessly.
|
||||
|
||||
Only ``redirect_uris`` is authoritative; every client is registered as a public
|
||||
|
|
@ -296,6 +325,8 @@ def aggregate_authorize(
|
|||
"invalid_request",
|
||||
"PKCE is required: send code_challenge with code_challenge_method=S256",
|
||||
)
|
||||
if len(state) > MAX_STATE_LENGTH:
|
||||
return _oauth_error(400, "invalid_request", f"state must be at most {MAX_STATE_LENGTH} characters")
|
||||
base_url = get_request_base_url(request)
|
||||
if session_user_id is None:
|
||||
login_url = f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}"
|
||||
|
|
@ -308,6 +339,7 @@ def aggregate_authorize(
|
|||
redirect_uri=redirect_uri,
|
||||
state=state,
|
||||
code_challenge=code_challenge,
|
||||
jti=secrets.token_urlsafe(24),
|
||||
exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS,
|
||||
)
|
||||
connect_url = _append_query_params(
|
||||
|
|
@ -318,7 +350,7 @@ def aggregate_authorize(
|
|||
path, secure = _cookie_path_and_secure(request)
|
||||
response.set_cookie(
|
||||
key=_flow_cookie_name(handle),
|
||||
value=_seal("", flow),
|
||||
value=_seal(_UNPREFIXED, flow),
|
||||
max_age=CONNECT_FLOW_TTL_SECONDS,
|
||||
path=path,
|
||||
secure=secure,
|
||||
|
|
@ -335,10 +367,11 @@ def _origin_only(url: str) -> str:
|
|||
return f"{parsed.scheme}://{parsed.netloc}" if parsed.netloc else ""
|
||||
|
||||
|
||||
def complete_connect_flow(
|
||||
async def complete_connect_flow(
|
||||
request: Request,
|
||||
flow_handle: str,
|
||||
session_user_id: str | None,
|
||||
cache: DualCache,
|
||||
) -> Response:
|
||||
"""The deliberate finish step of the connect flow: mint the gateway authorization
|
||||
code and send the browser back to the client.
|
||||
|
|
@ -346,12 +379,13 @@ def complete_connect_flow(
|
|||
Reached by POST so a cross-site GET cannot trigger it, and bound to the HttpOnly
|
||||
per-flow cookie plus an exact match between the signed-in user and the user sealed
|
||||
into the flow: a link crafted by another party dies here with ``access_denied``
|
||||
instead of minting a code for the victim's identity.
|
||||
instead of minting a code for the victim's identity. The flow is single-use (an atomic
|
||||
claim on its ``jti``), so a double-submit cannot mint two codes from one sign-in.
|
||||
"""
|
||||
sealed_flow = request.cookies.get(_flow_cookie_name(flow_handle))
|
||||
if sealed_flow is None:
|
||||
return _oauth_error(400, "invalid_request", "unknown or expired connect flow")
|
||||
flow = _open_sealed(sealed_flow, "", _ConnectFlow, _CONNECT_FLOW_DEBUG_KEY)
|
||||
flow = _open_sealed(sealed_flow, _UNPREFIXED, _ConnectFlow, _CONNECT_FLOW_DEBUG_KEY)
|
||||
if flow is None:
|
||||
return _oauth_error(400, "invalid_request", "unknown or expired connect flow")
|
||||
now = datetime.now(timezone.utc)
|
||||
|
|
@ -361,6 +395,10 @@ def complete_connect_flow(
|
|||
return _oauth_error(401, "login_required", "sign in to LiteLLM to finish connecting")
|
||||
if session_user_id != flow.user_id:
|
||||
return _oauth_error(403, "access_denied", "the signed-in user does not match this connect flow")
|
||||
if not await _SingleUseGuard(cache).claim(
|
||||
f"{_USED_FLOW_CACHE_PREFIX}{flow.jti}", CONNECT_FLOW_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS
|
||||
):
|
||||
return _oauth_error(400, "invalid_request", "this connect flow was already completed; restart the connection")
|
||||
code = _seal(
|
||||
GATEWAY_AUTH_CODE_PREFIX,
|
||||
_GatewayAuthCode(
|
||||
|
|
@ -381,27 +419,37 @@ def complete_connect_flow(
|
|||
|
||||
|
||||
def _pkce_verifier_matches(code_verifier: str, code_challenge: str) -> bool:
|
||||
"""RFC 7636 S256 verification, total over hostile input. The comparison is over bytes
|
||||
so a non-ASCII ``code_challenge`` (which reaches here unvalidated from the client's
|
||||
authorize request) simply fails to match instead of raising ``TypeError`` the way
|
||||
``hmac.compare_digest`` does on two ``str`` with non-ASCII content. The verifier is
|
||||
ASCII per spec; a compliant client's challenge is base64url and matches."""
|
||||
digest = hashlib.sha256(code_verifier.encode("ascii", "replace")).digest()
|
||||
computed = urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
|
||||
return hmac.compare_digest(computed, code_challenge)
|
||||
computed = urlsafe_b64encode(digest).rstrip(b"=")
|
||||
return hmac.compare_digest(computed, code_challenge.encode("utf-8"))
|
||||
|
||||
|
||||
class _SingleUseGuard:
|
||||
"""Best-effort single-use marking for gateway authorization codes over the injected
|
||||
proxy cache (in-memory always, Redis when the deployment wires it, in which case the
|
||||
guard holds across replicas). The code's 120s TTL is the hard bound either way; the
|
||||
guard exists so a same-process or shared-cache replay fails ``invalid_grant``."""
|
||||
"""Atomic single-use claim for a one-time id (an auth-code or connect-flow ``jti``) over
|
||||
the injected proxy cache.
|
||||
|
||||
Uses an atomic increment rather than a get-then-set: two concurrent redemptions of the
|
||||
same id cannot both observe "unused", because exactly one increment returns 1. With
|
||||
Redis wired this holds across replicas (``INCR`` is atomic); single-replica it holds in
|
||||
the in-memory cache. The id's own TTL is the outer bound. A claim is the gate, not a
|
||||
marker to check separately, so it fails closed: if the cache cannot record the claim
|
||||
(no backend at all) the id is refused rather than admitted. For the auth code, PKCE
|
||||
binding is the primary defense against interception; this makes the RFC 6749 4.1.2
|
||||
single-use property reliable on top of it."""
|
||||
|
||||
def __init__(self, cache: DualCache) -> None:
|
||||
self._cache = cache
|
||||
|
||||
async def already_used(self, jti: str) -> bool:
|
||||
return await self._cache.async_get_cache(f"{_USED_CODE_CACHE_PREFIX}{jti}") is not None
|
||||
|
||||
async def mark_used(self, jti: str) -> None:
|
||||
await self._cache.async_set_cache(
|
||||
f"{_USED_CODE_CACHE_PREFIX}{jti}", "1", ttl=GATEWAY_AUTH_CODE_TTL_SECONDS + 60
|
||||
)
|
||||
async def claim(self, key: str, ttl_seconds: int) -> bool:
|
||||
"""Atomically claim ``key``. ``True`` iff this caller is the first (increment to 1);
|
||||
``False`` on a replay (>1) or when the claim could not be recorded (fail closed)."""
|
||||
count = await self._cache.async_increment_cache(key, 1, ttl=ttl_seconds)
|
||||
return count == 1
|
||||
|
||||
|
||||
def _session_token_pair(principal: SessionPrincipal, keys: SessionKeys, now: datetime) -> Response:
|
||||
|
|
@ -422,11 +470,17 @@ def _session_token_pair(principal: SessionPrincipal, keys: SessionKeys, now: dat
|
|||
|
||||
|
||||
def _reload_failure_response(failure: ReloadUserFailure) -> Response:
|
||||
if failure == "unavailable":
|
||||
return _oauth_error(503, "temporarily_unavailable", "the gateway database is unavailable; retry")
|
||||
if failure == "unresolvable":
|
||||
return _oauth_error(500, "server_error", "the gateway is not configured to resolve users")
|
||||
return _oauth_error(400, "invalid_grant", "the user for this grant is no longer active")
|
||||
"""Map the live-user revalidation failure onto its OAuth error, exhaustively, so a new
|
||||
``ReloadUserFailure`` member is a type error here rather than silently 400ing."""
|
||||
match failure:
|
||||
case "unavailable":
|
||||
return _oauth_error(503, "temporarily_unavailable", "the gateway database is unavailable; retry")
|
||||
case "unresolvable":
|
||||
return _oauth_error(500, "server_error", "the gateway is not configured to resolve users")
|
||||
case "no_active_key":
|
||||
return _oauth_error(400, "invalid_grant", "the user for this grant is no longer active")
|
||||
case _:
|
||||
assert_never(failure)
|
||||
|
||||
|
||||
async def aggregate_token(
|
||||
|
|
@ -483,6 +537,8 @@ async def _authorization_code_grant(
|
|||
) -> Response:
|
||||
if not code or not redirect_uri or not code_verifier:
|
||||
return _oauth_error(400, "invalid_request", "code, redirect_uri, and code_verifier are required")
|
||||
if not MIN_CODE_VERIFIER_LENGTH <= len(code_verifier) <= MAX_CODE_VERIFIER_LENGTH:
|
||||
return _oauth_error(400, "invalid_request", "code_verifier must be 43 to 128 characters (RFC 7636)")
|
||||
parsed = _open_sealed(code, GATEWAY_AUTH_CODE_PREFIX, _GatewayAuthCode, _AUTH_CODE_DEBUG_KEY)
|
||||
if parsed is None:
|
||||
return _oauth_error(400, "invalid_grant", "the authorization code is invalid")
|
||||
|
|
@ -492,12 +548,17 @@ async def _authorization_code_grant(
|
|||
return _oauth_error(400, "invalid_grant", "the authorization code was issued to a different client")
|
||||
if not _pkce_verifier_matches(code_verifier, parsed.code_challenge):
|
||||
return _oauth_error(400, "invalid_grant", "PKCE verification failed")
|
||||
if await guard.already_used(parsed.jti):
|
||||
return _oauth_error(400, "invalid_grant", "the authorization code was already used")
|
||||
await guard.mark_used(parsed.jti)
|
||||
# Revalidate the user BEFORE claiming the code, so a transient DB outage (a retryable
|
||||
# 503) does not consume a still-valid code and force the client to restart sign-in.
|
||||
failure = await reload_user(parsed.user_id)
|
||||
if failure is not None:
|
||||
return _reload_failure_response(failure)
|
||||
# Atomic single-use claim is the gate: on a concurrent double-redeem exactly one caller
|
||||
# wins, and a claim that cannot be recorded fails closed.
|
||||
if not await guard.claim(
|
||||
f"{_USED_CODE_CACHE_PREFIX}{parsed.jti}", GATEWAY_AUTH_CODE_TTL_SECONDS + _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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -2420,12 +2420,19 @@ async def sso_readiness():
|
|||
|
||||
|
||||
def _is_same_origin_return_path(return_to: str) -> bool:
|
||||
"""True for a strictly relative return path (starts with ``/``, not
|
||||
protocol-relative ``//``, no backslash tricks browsers normalize to slashes), which
|
||||
stays on the gateway's own origin by construction and is therefore safe to honor
|
||||
without a configured ``control_plane_url``. Used by the MCP gateway DCR authorize
|
||||
round-trip so a browser sent through login lands back on the authorize request."""
|
||||
return return_to.startswith("/") and not return_to.startswith("//") and "\\" not in return_to
|
||||
"""True for a strictly relative return path that stays on the gateway's own origin by
|
||||
construction, and is therefore safe to honor without a configured ``control_plane_url``.
|
||||
Used by the MCP gateway DCR authorize round-trip so a browser sent through login lands
|
||||
back on the authorize request.
|
||||
|
||||
Requires a single leading ``/`` (not protocol-relative ``//``), no backslash (browsers
|
||||
fold ``\\`` to ``/``, so ``/\\evil.com`` would escape the origin), and no control or
|
||||
whitespace characters. Rejecting control chars keeps a ``\\r\\n``/tab-bearing value out
|
||||
of the redirect ``Location`` and the ``litellm_cp_return_to`` cookie entirely, rather
|
||||
than relying on downstream header encoding to neutralize it."""
|
||||
if not return_to.startswith("/") or return_to.startswith("//") or "\\" in return_to:
|
||||
return False
|
||||
return not any(ord(ch) < 0x20 or ch in (" ", "\x7f") for ch in return_to)
|
||||
|
||||
|
||||
class SSOAuthenticationHandler:
|
||||
|
|
|
|||
|
|
@ -2899,19 +2899,20 @@ async def test_token_root_does_not_resolve_private_server_for_external_client():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_root_resolves_single_oauth2_server():
|
||||
"""When /register is hit without server name and exactly 1 OAuth2 server exists, resolve it."""
|
||||
try:
|
||||
from fastapi import Request
|
||||
async def test_register_root_does_aggregate_dcr_not_single_server_resolution():
|
||||
"""Root /register is the aggregate DCR endpoint: it mints a stateless llm_dcrc_ client
|
||||
from the request's redirect_uris and does NOT resolve a single configured oauth2 server
|
||||
(a single-server deployment registers at /{server}/register instead)."""
|
||||
import json
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
register_client,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
register_client,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
oauth2_server = _create_oauth2_server()
|
||||
|
|
@ -2922,33 +2923,37 @@ async def test_register_root_resolves_single_oauth2_server():
|
|||
mock_request.headers = {}
|
||||
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
|
||||
new=AsyncMock(return_value={}),
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
|
||||
new=AsyncMock(return_value={"redirect_uris": ["https://claude.ai/cb"]}),
|
||||
),
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637"),
|
||||
):
|
||||
result = await register_client(request=mock_request, mcp_server_name=None)
|
||||
response = await register_client(request=mock_request, mcp_server_name=None)
|
||||
|
||||
# Should resolve to the single server and return its name as client_id
|
||||
assert result["client_id"] == "test_oauth"
|
||||
assert "redirect_uris" in result
|
||||
body = json.loads(response.body)
|
||||
assert body["client_id"].startswith("llm_dcrc_")
|
||||
assert body["client_id"] != "test_oauth"
|
||||
assert body["token_endpoint_auth_method"] == "none"
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_root_does_not_resolve_private_server_for_external_client():
|
||||
"""Root /register must not reveal or use a hidden MCP server."""
|
||||
try:
|
||||
from fastapi import Request
|
||||
async def test_register_root_does_not_leak_a_private_server():
|
||||
"""Root /register never resolves or reveals a configured server, so a private one cannot
|
||||
leak to an external caller: it always mints the aggregate DCR client instead."""
|
||||
import json
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
register_client,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
register_client,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
oauth2_server = _create_oauth2_server(available_on_public_internet=False)
|
||||
|
|
@ -2962,17 +2967,19 @@ async def test_register_root_does_not_resolve_private_server_for_external_client
|
|||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
|
||||
new=AsyncMock(return_value={}),
|
||||
new=AsyncMock(return_value={"redirect_uris": ["https://claude.ai/cb"]}),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip",
|
||||
return_value="198.51.100.10",
|
||||
),
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637"),
|
||||
):
|
||||
result = await register_client(request=mock_request, mcp_server_name=None)
|
||||
response = await register_client(request=mock_request, mcp_server_name=None)
|
||||
|
||||
assert result["client_id"] == "dummy_client"
|
||||
assert result["redirect_uris"] == ["https://llm.example.com/callback"]
|
||||
body = json.loads(response.body)
|
||||
assert body["client_id"].startswith("llm_dcrc_")
|
||||
assert "test_oauth" not in body["client_id"]
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
|
@ -7260,6 +7267,7 @@ async def test_bare_origin_discovery_resolves_single_server_not_aggregate():
|
|||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
||||
|
||||
def test_gateway_dcr_flow_routing_engages_only_for_llm_dcrc_clients(monkeypatch):
|
||||
"""The aggregate DCR arms engage for llm_dcrc_ client_ids (register always mints one,
|
||||
authorize/token route into the aggregate flow); a non-gateway client_id keeps the
|
||||
|
|
@ -7312,8 +7320,6 @@ def test_gateway_dcr_flow_routing_engages_only_for_llm_dcrc_clients(monkeypatch)
|
|||
assert token_response.status_code == 400
|
||||
assert token_response.json()["error"] == "invalid_grant"
|
||||
|
||||
# a non-gateway (upstream-issued) client_id is not routed into the aggregate arm; it
|
||||
# falls to the per-server exchange, which 404s for an unknown server
|
||||
upstream_shaped = client.post(
|
||||
"/token",
|
||||
data={"grant_type": "authorization_code", "client_id": "regular-upstream-client", "code": "x"},
|
||||
|
|
|
|||
|
|
@ -62,7 +62,7 @@ def _request(path="/authorize", query="", cookies=None, method="GET"):
|
|||
|
||||
|
||||
async def _register(redirect_uris) -> dict:
|
||||
response = await register_aggregate_client(request=_request("/register"), request_body={"redirect_uris": redirect_uris})
|
||||
response = await register_aggregate_client(request_body={"redirect_uris": redirect_uris})
|
||||
return json.loads(response.body)
|
||||
|
||||
|
||||
|
|
@ -103,7 +103,7 @@ async def test_register_allows_loopback_http_for_dev_clients():
|
|||
],
|
||||
)
|
||||
async def test_register_rejects_bad_redirect_uris(redirect_uris):
|
||||
response = await register_aggregate_client(request=_request("/register"), request_body={"redirect_uris": redirect_uris})
|
||||
response = await register_aggregate_client(request_body={"redirect_uris": redirect_uris})
|
||||
assert response.status_code == 400
|
||||
assert json.loads(response.body)["error"] in ("invalid_redirect_uri", "invalid_client_metadata")
|
||||
|
||||
|
|
@ -117,7 +117,9 @@ async def test_tampered_client_id_does_not_open():
|
|||
assert open_gateway_dcr_client("other_prefix") is None
|
||||
|
||||
|
||||
def _authorize(client_id, session_user_id, redirect_uri=REDIRECT_URI, challenge=CODE_CHALLENGE, method="S256", response_type="code"):
|
||||
def _authorize(
|
||||
client_id, session_user_id, redirect_uri=REDIRECT_URI, challenge=CODE_CHALLENGE, method="S256", response_type="code"
|
||||
):
|
||||
return aggregate_authorize(
|
||||
request=_request(query=f"client_id={client_id}"),
|
||||
client_id=client_id,
|
||||
|
|
@ -188,24 +190,27 @@ async def test_full_walk_register_authorize_complete_token_and_replay():
|
|||
authorize_response = _authorize(client_id, session_user_id="u1")
|
||||
handle, cookies = _flow_cookie_from(authorize_response)
|
||||
|
||||
denied = complete_connect_flow(
|
||||
denied = await complete_connect_flow(
|
||||
request=_request("/authorize/complete", cookies=cookies, method="POST"),
|
||||
flow_handle=handle,
|
||||
session_user_id="attacker",
|
||||
cache=DualCache(),
|
||||
)
|
||||
assert denied.status_code == 403
|
||||
|
||||
anonymous = complete_connect_flow(
|
||||
anonymous = await complete_connect_flow(
|
||||
request=_request("/authorize/complete", cookies=cookies, method="POST"),
|
||||
flow_handle=handle,
|
||||
session_user_id=None,
|
||||
cache=DualCache(),
|
||||
)
|
||||
assert anonymous.status_code == 401
|
||||
|
||||
completed = complete_connect_flow(
|
||||
completed = await complete_connect_flow(
|
||||
request=_request("/authorize/complete", cookies=cookies, method="POST"),
|
||||
flow_handle=handle,
|
||||
session_user_id="u1",
|
||||
cache=DualCache(),
|
||||
)
|
||||
assert completed.status_code == 303
|
||||
redirect = urlparse(completed.headers["location"])
|
||||
|
|
@ -269,15 +274,19 @@ async def test_full_walk_register_authorize_complete_token_and_replay():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_complete_rejects_missing_tampered_and_expired_flows():
|
||||
missing = complete_connect_flow(
|
||||
request=_request("/authorize/complete", method="POST"), flow_handle="nope", session_user_id="u1"
|
||||
missing = await complete_connect_flow(
|
||||
request=_request("/authorize/complete", method="POST"),
|
||||
flow_handle="nope",
|
||||
session_user_id="u1",
|
||||
cache=DualCache(),
|
||||
)
|
||||
assert missing.status_code == 400
|
||||
|
||||
tampered = complete_connect_flow(
|
||||
tampered = await complete_connect_flow(
|
||||
request=_request("/authorize/complete", cookies={f"{CONNECT_FLOW_COOKIE_PREFIX}h1": "garbage"}, method="POST"),
|
||||
flow_handle="h1",
|
||||
session_user_id="u1",
|
||||
cache=DualCache(),
|
||||
)
|
||||
assert tampered.status_code == 400
|
||||
|
||||
|
|
@ -353,10 +362,11 @@ async def test_token_gates_on_live_user_revalidation(failure, expected_status, e
|
|||
client_id = (await _register([REDIRECT_URI]))["client_id"]
|
||||
authorize_response = _authorize(client_id, session_user_id="deactivated-user")
|
||||
handle, cookies = _flow_cookie_from(authorize_response)
|
||||
completed = complete_connect_flow(
|
||||
completed = await complete_connect_flow(
|
||||
request=_request("/authorize/complete", cookies=cookies, method="POST"),
|
||||
flow_handle=handle,
|
||||
session_user_id="deactivated-user",
|
||||
cache=DualCache(),
|
||||
)
|
||||
code = parse_qs(urlparse(completed.headers["location"]).query)["code"][0]
|
||||
|
||||
|
|
@ -377,3 +387,103 @@ async def test_token_gates_on_live_user_revalidation(failure, expected_status, e
|
|||
)
|
||||
assert response.status_code == expected_status
|
||||
assert json.loads(response.body)["error"] == expected_error
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flow_is_single_use_shared_cache_rejects_second_complete():
|
||||
"""A double-submit of the finish step mints only ONE code: the second complete over the
|
||||
same cache fails invalid_request (atomic flow claim), so one sign-in cannot yield two codes."""
|
||||
cache = DualCache()
|
||||
client_id = (await _register([REDIRECT_URI]))["client_id"]
|
||||
handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1"))
|
||||
|
||||
first = await complete_connect_flow(
|
||||
request=_request("/authorize/complete", cookies=cookies, method="POST"),
|
||||
flow_handle=handle,
|
||||
session_user_id="u1",
|
||||
cache=cache,
|
||||
)
|
||||
assert first.status_code == 303
|
||||
second = await complete_connect_flow(
|
||||
request=_request("/authorize/complete", cookies=cookies, method="POST"),
|
||||
flow_handle=handle,
|
||||
session_user_id="u1",
|
||||
cache=cache,
|
||||
)
|
||||
assert second.status_code == 400
|
||||
assert json.loads(second.body)["error"] == "invalid_request"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_rejects_out_of_range_code_verifier():
|
||||
"""RFC 7636: a code_verifier outside 43-128 chars is invalid_request, not a confusing
|
||||
invalid_grant PKCE-mismatch."""
|
||||
for bad in ["short", "x" * 200]:
|
||||
response = await aggregate_token(
|
||||
request=_request("/token", method="POST"),
|
||||
grant_type="authorization_code",
|
||||
code="llm_gcode_whatever",
|
||||
redirect_uri=REDIRECT_URI,
|
||||
client_id="llm_dcrc_x",
|
||||
code_verifier=bad,
|
||||
refresh_token=None,
|
||||
master_key=MASTER_KEY,
|
||||
reload_user=_reload_user_active,
|
||||
cache=DualCache(),
|
||||
)
|
||||
assert response.status_code == 400
|
||||
assert json.loads(response.body)["error"] == "invalid_request"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authorize_rejects_over_long_state():
|
||||
client_id = (await _register([REDIRECT_URI]))["client_id"]
|
||||
response = aggregate_authorize(
|
||||
request=_request(query=f"client_id={client_id}"),
|
||||
client_id=client_id,
|
||||
redirect_uri=REDIRECT_URI,
|
||||
state="s" * 2000,
|
||||
code_challenge=CODE_CHALLENGE,
|
||||
code_challenge_method="S256",
|
||||
response_type="code",
|
||||
session_user_id="u1",
|
||||
)
|
||||
assert response.status_code == 400
|
||||
assert json.loads(response.body)["error"] == "invalid_request"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_ascii_code_challenge_fails_grant_not_500():
|
||||
"""A non-ASCII code_challenge (unvalidated from the client) must yield a clean
|
||||
invalid_grant, never a TypeError-driven 500 (bytes comparison, not str)."""
|
||||
client_id = (await _register([REDIRECT_URI]))["client_id"]
|
||||
# Seal a code carrying a non-ASCII challenge directly (authorize requires S256 shape,
|
||||
# but the challenge charset is not validated there, so this state is reachable).
|
||||
from datetime import datetime, timezone
|
||||
|
||||
code = _seal(
|
||||
GATEWAY_AUTH_CODE_PREFIX,
|
||||
_GatewayAuthCode(
|
||||
user_id="u1",
|
||||
client_id=client_id,
|
||||
redirect_uri=REDIRECT_URI,
|
||||
code_challenge="challenge-with-€-non-ascii",
|
||||
jti="jti-x",
|
||||
iat=int(datetime.now(timezone.utc).timestamp()),
|
||||
exp=int(datetime.now(timezone.utc).timestamp()) + 120,
|
||||
),
|
||||
)
|
||||
response = await aggregate_token(
|
||||
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=DualCache(),
|
||||
)
|
||||
assert response.status_code == 400
|
||||
assert json.loads(response.body)["error"] == "invalid_grant"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue