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:
Tin Chi Lo 2026-07-14 09:52:35 -07:00
parent 20cbadead1
commit 77774c1e69
5 changed files with 557 additions and 97 deletions

View file

@ -40,6 +40,7 @@ from litellm.proxy._experimental.mcp_server.faults import (
render_token_fault,
)
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import (
ReloadUserFailure,
aggregate_authorize,
aggregate_token,
complete_connect_flow,
@ -412,6 +413,276 @@ def _validate_token_response(
)
def _litellm_key_from_request(request: Request) -> Optional[str]:
"""Return the LiteLLM API key presented on the request, or ``None``.
Accepts the key from ``x-litellm-api-key`` (what MCP clients such as Claude Desktop/Code
send) as well as ``Authorization``; either may carry a bare token or ``Bearer <token>``.
``x-litellm-api-key`` wins when both are present, since ``Authorization`` may instead carry
an OAuth/upstream bearer.
"""
for header_value in (
request.headers.get("x-litellm-api-key"),
request.headers.get("Authorization") or request.headers.get("authorization"),
):
if not header_value:
continue
value = header_value.strip()
if value.lower().startswith("bearer "):
value = value[7:].strip()
if value:
return value
return None
def _key_is_active(key_obj: "UserAPIKeyAuth") -> bool:
"""``True`` when the presented key is neither blocked nor past its expiry.
The OAuth token endpoint is unauthenticated, so the presented key is validated here before it is
trusted; a revoked or expired key must not mint a bridge envelope or write a stored credential.
``get_key_object`` resolves a row without these checks (the main ``user_api_key_auth`` pipeline
enforces them downstream, which this endpoint bypasses), so they are applied here. Deleted keys
are already rejected upstream, where ``get_key_object`` raises on a row that no longer exists.
This is an active-state gate only; it deliberately does not require a ``user_id``. A valid
team-scoped or service-account key has no ``user_id`` yet is a legitimate credential, so gating
on ``user_id`` presence would wrongly reject it. Callers that need the user (the per-user token
store) derive it separately via :func:`_active_key_user_id`.
Total by design: ``expires`` is typed ``str | datetime``, and an unparseable string would make
``datetime.fromisoformat`` raise. Since the callers run this outside their key-resolution
``try``, an uncaught parse error would surface as a 500 instead of the endpoint's fail-closed
behavior, so a malformed expiry is treated as inactive (return ``False``) rather than raising.
"""
if key_obj.blocked is True:
return False
expires = key_obj.expires
if expires is not None:
if isinstance(expires, datetime):
expiry = expires
else:
try:
expiry = datetime.fromisoformat(expires)
except (ValueError, TypeError):
return False
if expiry.tzinfo is None or expiry.tzinfo.utcoffset(expiry) is None:
expiry = expiry.replace(tzinfo=timezone.utc)
if expiry < datetime.now(timezone.utc):
return False
return True
def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> str | None:
"""The active key's ``user_id``, or ``None`` when the key is blocked/expired or simply has no
``user_id`` (a team-scoped or service-account key). Used only by the per-user token store, which
needs a user to key the stored credential; the bridge mint uses the key hash and does not."""
return key_obj.user_id if _key_is_active(key_obj) else None
@dataclass(frozen=True, slots=True)
class _ResolvedKey:
"""An active litellm key resolved from the token request: its hash (the value ``get_key_object``
and the cache/DB layer key the record by) and the live record."""
key_hash: str
key: "UserAPIKeyAuth"
# The token endpoint injects `_reload_active_user_by_id` as the flow's `ReloadUser`, so the
# two must share one failure type; alias the flow's canonical union rather than redeclare it.
_KeyResolutionFailure = ReloadUserFailure
"""Why a token request yielded no active litellm key, kept distinct so a caller statuses each truthfully
instead of blaming the client for a gateway problem:
- ``no_active_key``: none was presented, or the presented key is unknown / blocked / expired (the
caller's request is at fault)
- ``unavailable``: the auth database was transiently unreachable while resolving (retryable)
- ``unresolvable``: the gateway cannot resolve identity right now (no DB connection, or an unexpected
error) -- a gateway fault, not the caller's
The classification mirrors admission's ``_reload_admitted_key`` so the mint (ingress) and admission
(egress) never disagree on the status of the same outage."""
async def _resolve_active_litellm_key(request: Request) -> "_ResolvedKey | _KeyResolutionFailure":
"""Resolve the presented litellm key to an active key record, or say precisely why not.
Single resolution path the OAuth token endpoint reuses, resolving authoritatively via
``get_key_object`` (cache first, then DB). The failure is a value, not a bare ``None``, so a caller
can tell "the client sent no usable credential" (a request error) apart from "the gateway could not
check" (an infrastructure error) and status each truthfully; collapsing both to ``None`` is what let
a DB outage read as a 400. A resolved key is still gated by ``_key_is_active``, so a blocked or
expired key is ``no_active_key`` while a valid team-scoped or service-account key (no ``user_id``)
resolves. Classification mirrors admission's ``_reload_admitted_key``: no DB connection is a gateway
fault, a ``ProxyException`` / ``HTTPException`` from ``get_key_object`` is an unknown or invalid key,
a database-service-unavailable error is a retryable outage, and anything else is an unexpected
gateway fault."""
token = _litellm_key_from_request(request)
if not token:
return "no_active_key"
from litellm.proxy._types import hash_token # noqa: PLC0415 # inline import avoids a module-load circular import
return await _reload_active_key_by_hash(hash_token(token))
async def _reload_active_key_by_hash(key_hash: str) -> "_ResolvedKey | _KeyResolutionFailure":
"""Reload the live key record for ``key_hash`` (cache first, then DB) and gate it on active state,
returning the resolved key or a precise failure. Shared by the token request's presented-key
resolution (:func:`_resolve_active_litellm_key`, which hashes the presented key) and the refresh
path (which already holds the hash sealed in the refresh envelope), so both re-validate identity
through one active-key gate and one failure classification. Classification mirrors admission's
``_reload_admitted_key``: no DB connection is a gateway fault, a ``ProxyException`` / ``HTTPException``
from ``get_key_object`` is an unknown or invalid key, a database-service-unavailable error is a
retryable outage, and anything else is an unexpected gateway fault. A blocked or expired key is
``no_active_key``, so a revoked key can neither mint nor refresh a bridge envelope."""
from litellm.proxy._types import (
ProxyException, # noqa: PLC0415 # inline import avoids a module-load circular import
)
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
get_key_object,
)
from litellm.proxy.db.exception_handler import ( # noqa: PLC0415 # inline import avoids a module-load circular import
PrismaDBExceptionHandler,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
prisma_client,
user_api_key_cache,
)
if prisma_client is None:
return "unresolvable"
try:
key_obj = await get_key_object(
hashed_token=key_hash,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
except (ProxyException, HTTPException):
return "no_active_key"
except Exception as exc: # noqa: BLE001 # classify: a DB outage is retryable, anything else is an opaque gateway fault
if PrismaDBExceptionHandler.is_database_service_unavailable_error(exc):
return "unavailable"
verbose_logger.debug(
"_reload_active_key_by_hash: unexpected key-resolution error (%s)",
type(exc).__name__,
)
return "unresolvable"
if not _key_is_active(key_obj):
return "no_active_key"
return _ResolvedKey(key_hash=key_hash, key=key_obj)
async def _reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | None":
"""Re-validate a live litellm user by id, returning ``None`` when the user is active or a precise
failure otherwise. The interactive DCR client authenticates via SSO, so its refresh envelope seals a
user subject; renewing it must re-check the user is still live (present and not SCIM-deactivated) so a
deactivated user cannot keep refreshing, mirroring how admission re-validates the same user subject on
the egress side. No DB connection is a gateway fault (``unresolvable``) and a
database-service-unavailable error is a retryable outage (``unavailable``). Everything else fails
closed as ``no_active_key`` (the caller maps it to invalid_grant): a ``ProxyException`` /
``HTTPException``, a SCIM-deactivated user, and, unlike the key path, a missing user. ``get_user_object``
catches every DB failure and re-raises a bare ``ValueError`` (a deleted user and a real outage look
identical, the original error surviving only as ``__context__``), so the outage check walks the cause
chain, and a missing user falls through to ``no_active_key`` rather than an opaque gateway fault."""
from litellm.proxy._types import (
ProxyException, # noqa: PLC0415 # inline import avoids a module-load circular import
)
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
get_user_object,
)
from litellm.proxy.db.exception_handler import ( # noqa: PLC0415 # inline import avoids a module-load circular import
PrismaDBExceptionHandler,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
prisma_client,
user_api_key_cache,
)
if prisma_client is None:
return "unresolvable"
try:
user_object = await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
)
except (ProxyException, HTTPException):
return "no_active_key"
except Exception as exc: # noqa: BLE001 # a DB outage is retryable; a missing user (get_user_object's wrapped ValueError) or any other resolution failure fails closed as no_active_key, never a 500
if PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(exc):
return "unavailable"
verbose_logger.debug("_reload_active_user_by_id: user-resolution error (%s)", type(exc).__name__)
return "no_active_key"
if user_object is None:
return "no_active_key"
if isinstance(user_object.metadata, dict) and user_object.metadata.get("scim_active") is False:
return "no_active_key"
return None
async def _key_owner_scim_deactivated(key: "UserAPIKeyAuth") -> bool:
"""True only when the key's owning user was explicitly SCIM-deactivated, so a refresh revokes an
offboarded owner's key exactly as admission does via ``_reject_if_admitted_owner_scim_deactivated``.
A key with no owner, a missing owner record, or a failed lookup fails OPEN (returns ``False``),
matching admission and the standard builder: a key may outlive its owner record, and a transient DB
blip must not revoke a live key. Only an explicit ``scim_active`` of ``False`` gates renewal."""
if key.user_id is None:
return False
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
get_user_object,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
prisma_client,
user_api_key_cache,
)
if prisma_client is None:
return False
try:
owner = await get_user_object(
user_id=key.user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
)
except Exception as exc: # noqa: BLE001 # fail open: a missing owner (get_user_object's wrapped ValueError) or a DB blip must not revoke a live key
verbose_logger.debug("refresh: key-owner SCIM lookup failed, not revoking (%s)", type(exc).__name__)
return False
return owner is not None and isinstance(owner.metadata, dict) and owner.metadata.get("scim_active") is False
async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResolutionFailure | None":
"""Re-validate that the subject sealed in a refresh envelope is still live, dispatching on its type:
a key_hash reloads the virtual key, a user_id reloads the user. Returns ``None`` when the subject is
active or a precise failure otherwise, so revocation gates renewal for either identity source the same
way admission gates the egress: a blocked or expired key, a SCIM-deactivated key owner (mirroring
admission's owner check, so an offboarded user cannot keep renewing a still-active key), and a
deactivated or deleted user all fail closed to ``no_active_key``."""
match identity.subject_type:
case "key_hash":
reloaded = await _reload_active_key_by_hash(identity.subject)
if not isinstance(reloaded, _ResolvedKey):
return reloaded
if await _key_owner_scim_deactivated(reloaded.key):
return "no_active_key"
return None
case "user_id":
return await _reload_active_user_by_id(identity.subject)
case _:
assert_never(identity.subject_type)
async def _extract_user_id_from_request(request: Request) -> str | None:
"""The litellm ``user_id`` for the token request, so a per-user token is stored under the same
identity the egress later reads it by. Storage is best-effort, so every non-resolved outcome
(including a transient DB outage) collapses to ``None`` here and the caller simply skips the store;
the bridge mint, which must status those outcomes differently, consumes
:func:`_resolve_active_litellm_key` directly."""
resolved = await _resolve_active_litellm_key(request)
if not isinstance(resolved, _ResolvedKey):
return None
return _active_key_user_id(resolved.key)
async def _store_per_user_token_server_side(
server: MCPServer,
user_id: str,
@ -1340,7 +1611,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,
@ -1415,7 +1686,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,
@ -1457,16 +1728,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,
)
@ -2194,8 +2465,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(

View file

@ -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)

View file

@ -2384,12 +2384,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:

View file

@ -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()
@ -7630,6 +7637,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
@ -7682,8 +7690,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"},

View file

@ -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"