From 2c691d3820e25fa66224b89bccfcaf8476416ae5 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 20 Aug 2026 04:10:10 -0700 Subject: [PATCH] feat(proxy): native CLI login with OAuth authorization code + PKCE The proxy's OAuth authorization server (dynamic registration, PKCE S256, loopback redirects, single-use codes, refresh rotation) gains a proxy-API audience: /authorize?resource= renders a consent page with team selection and /token mints the same per-user credential lite login mints, so a native CLI can sign a user in through the system browser and call /v1/* with user and team attribution. Adds GET /.well-known/litellm-cli-auth as the versioned discovery contract for non-Python clients, POST /revoke (RFC 7009) for logout, and lite login --pkce, lite logout, and lite auth print-token on the CLI side. Proxy-API grants only ever redirect to a loopback address and the server never picks a team on the user's behalf. Fixes #37332 --- litellm/litellm_core_utils/cli_token_utils.py | 3 + .../mcp_server/bridge_token_flow.py | 11 +- .../mcp_server/discoverable_endpoints.py | 53 +- .../mcp_server/gateway_dcr_flow.py | 500 +++++++++++++--- .../outbound_credentials/session_token.py | 19 +- .../mcp_server/proxy_api_credentials.py | 84 +++ litellm/proxy/_lazy_features.py | 2 + litellm/proxy/client/cli/commands/auth.py | 91 ++- .../proxy/client/cli/commands/pkce_login.py | 471 +++++++++++++++ .../html_forms/native_client_consent.py | 91 +++ litellm/proxy/management_endpoints/ui_sso.py | 22 +- .../test_cli_token_utils.py | 27 + .../test_session_token.py | 52 ++ .../mcp_server/test_discoverable_endpoints.py | 229 +++++++ .../mcp_server/test_gateway_dcr_flow.py | 564 +++++++++++++++++- .../mcp_server/test_proxy_api_credentials.py | 167 ++++++ .../proxy/client/cli/test_auth_commands.py | 323 +++++++++- .../proxy/client/cli/test_pkce_login.py | 553 +++++++++++++++++ .../html_forms/test_native_client_consent.py | 70 +++ .../proxy/management_endpoints/test_ui_sso.py | 20 +- 20 files changed, 3209 insertions(+), 143 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/proxy_api_credentials.py create mode 100644 litellm/proxy/client/cli/commands/pkce_login.py create mode 100644 litellm/proxy/common_utils/html_forms/native_client_consent.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py create mode 100644 tests/test_litellm/proxy/client/cli/test_pkce_login.py create mode 100644 tests/test_litellm/proxy/common_utils/html_forms/test_native_client_consent.py diff --git a/litellm/litellm_core_utils/cli_token_utils.py b/litellm/litellm_core_utils/cli_token_utils.py index a44ce431f4e..0043bc31204 100644 --- a/litellm/litellm_core_utils/cli_token_utils.py +++ b/litellm/litellm_core_utils/cli_token_utils.py @@ -79,6 +79,9 @@ def is_cli_token_fresh(token_data: Mapping[str, object], buffer_hours: float = 0 `LITELLM_CLI_JWT_EXPIRATION_HOURS`.""" from litellm.constants import CLI_JWT_EXPIRATION_HOURS + expires_at: Final = token_data.get("expires_at") + if isinstance(expires_at, (int, float)): + return time.time() < expires_at - buffer_hours * 3600 timestamp: Final = token_data.get("timestamp") if not isinstance(timestamp, (int, float)): return False diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index 1d8d545023d..b8c25236b0d 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -15,6 +15,7 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HE from litellm.types.mcp_server.mcp_server_manager import MCPServer if TYPE_CHECKING: + from litellm.models.user import LiteLLM_UserTable from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _BridgeAuthorizationCode from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( EnvelopeIdentity, @@ -181,7 +182,13 @@ async def _reload_active_key_by_hash(key_hash: str) -> "_ResolvedKey | _KeyResol 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 + """``None`` when the user is live, else the precise failure ``load_active_user_by_id`` found.""" + loaded: Final = await load_active_user_by_id(user_id) + return loaded if isinstance(loaded, str) else None + + +async def load_active_user_by_id(user_id: str) -> "LiteLLM_UserTable | _KeyResolutionFailure": + """Load a live litellm user by id, returning the record 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 @@ -226,7 +233,7 @@ async def _reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | No 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 + return user_object async def _key_owner_scim_deactivated(key: "UserAPIKeyAuth") -> bool: diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 1ca4c657706..f9c5db76fa0 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -47,8 +47,12 @@ from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( aggregate_token, complete_connect_flow, is_gateway_dcr_client_id, + is_proxy_api_resource, + native_client_auth_contract, + native_client_authorize, register_aggregate_client, relative_request_url, + revoke_refresh_token, ) from litellm.proxy._experimental.mcp_server.oauth_utils import ( TOKEN_NO_CACHE_HEADERS, @@ -58,6 +62,10 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import ( validate_trusted_redirect_uri, well_known_root_suffix, ) +from litellm.proxy._experimental.mcp_server.proxy_api_credentials import ( + lookup_consent_teams, + mint_proxy_credential, +) from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, @@ -1663,6 +1671,18 @@ async def authorize( ) if mcp_server_name is None and client_id and is_gateway_dcr_client_id(client_id): + if is_proxy_api_resource(request, resource): + return await native_client_authorize( + request=request, + client_id=client_id, + redirect_uri=redirect_uri, + state=state, + code_challenge=code_challenge, + code_challenge_method=code_challenge_method, + response_type=response_type, + session_user_id=_session_cookie_user_id(request), + lookup_consent_teams=lookup_consent_teams, + ) return aggregate_authorize( request=request, client_id=client_id, @@ -1764,6 +1784,7 @@ async def token_endpoint( reload_user=_reload_active_user_by_id, cache=user_api_key_cache, resource=resource, + mint_proxy_credential=mint_proxy_credential, ) lookup_name: Final = mcp_server_name or client_id @@ -1793,12 +1814,19 @@ async def token_endpoint( @router.post("/authorize/complete") -async def authorize_complete(request: Request, flow: str = Form(...), delivery: str | None = Form(None)): +async def authorize_complete( + request: Request, + flow: str = Form(...), + delivery: str | None = Form(None), + team_id: str | None = Form(None), + decision: str | None = Form(None), +) -> Response: """Finish an aggregate connect flow: mint the gateway authorization code for the signed-in user and hand it back to the DCR client, by 303 redirect (default) or, for a loopback client on a different machine, as a copyable callback URL (``delivery=manual``). POST plus the per-flow HttpOnly cookie set at /authorize; an - anonymous or bad-flow request just 400s.""" + anonymous or bad-flow request just 400s. The native-client consent page adds + ``decision`` (approve or deny) and the ``team_id`` the credential is attributed to.""" from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 # circular import at module load return await complete_connect_flow( @@ -1807,9 +1835,30 @@ async def authorize_complete(request: Request, flow: str = Form(...), delivery: session_user_id=_session_cookie_user_id(request), cache=user_api_key_cache, delivery=delivery, + team_id=team_id, + decision=decision, ) +@router.post("/revoke") +async def revoke_endpoint(request: Request, token: str = Form(...), client_id: str = Form(...)) -> Response: + """RFC 7009 revocation for the gateway's refresh tokens (``lite logout``). Always 200 + for a known client, whatever the token's state; access tokens expire on their own.""" + from litellm.proxy.proxy_server import ( # noqa: PLC0415 # circular import at module load + master_key, + user_api_key_cache, + ) + + return await revoke_refresh_token(token=token, client_id=client_id, master_key=master_key, cache=user_api_key_cache) + + +@router.get("/.well-known/litellm-cli-auth") +async def native_client_auth_discovery(request: Request) -> JSONResponse: + """The versioned contract a native client (``lite login --pkce``, or a CLI in any other + language) reads to sign a user in through the browser and obtain a proxy credential.""" + return JSONResponse(native_client_auth_contract(request), headers=TOKEN_NO_CACHE_HEADERS) + + # Per RFC 6749 ยง4.1.2.1, an IdP that rejects an OAuth authorization request # redirects back to the configured redirect URI with ``error`` / # ``error_description`` / ``error_uri`` query params and no ``code``. The MCP diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index 85885fc75f5..53db6fbf9d1 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -42,15 +42,16 @@ import hmac import html import secrets from base64 import urlsafe_b64encode -from collections.abc import Awaitable, Callable, Mapping +from collections.abc import Awaitable, Callable, Iterable, Mapping from datetime import datetime, timezone -from typing import Final, Literal, TypeVar +from types import MappingProxyType +from typing import Final, Literal, Protocol, TypeVar from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse from fastapi import HTTPException, Request from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response from pydantic import BaseModel, ConfigDict, Field, ValidationError -from typing_extensions import assert_never +from typing_extensions import ReadOnly, TypedDict, assert_never from litellm._logging import verbose_logger from litellm.caching.caching import DualCache @@ -70,6 +71,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credent from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import ( SESSION_REFRESH_TTL_SECONDS, MintedSessionToken, + SessionAudience, SessionKeys, SessionPrincipal, mint_session_refresh_token, @@ -79,6 +81,9 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) +from litellm.proxy.common_utils.html_forms.native_client_consent import ( + render_native_client_consent_page, +) from litellm.types.mcp_server.mcp_server_manager import MCPServer GATEWAY_DCR_CLIENT_ID_PREFIX: Final = "llm_dcrc_" @@ -144,6 +149,46 @@ ReloadUser = Callable[[str], Awaitable[ReloadUserFailure | None]] ``None`` means the user is active; ``unavailable`` is a retryable DB outage; anything else fails the grant closed.""" +PROXY_API_AUDIENCE: Final[SessionAudience] = "proxy_api" +"""The audience a native client (``lite login --pkce``, a Go CLI) asks for by sending the +proxy base URL itself as the RFC 8707 ``resource``: the grant then mints the proxy-API CLI +credential that LLM routes accept, instead of the MCP-only session pair.""" + +ProxyCredentialMintFailure = Literal[ReloadUserFailure, "not_a_member"] + + +class MintedProxyCredential(BaseModel): + model_config = ConfigDict(frozen=True) + key: str = Field(min_length=1) + expires_in: int = Field(gt=0) + user_id: str = Field(min_length=1) + team_id: str | None = None + + +class MintProxyCredential(Protocol): + """Injected proxy-API credential minter ``(user_id, team_id)``: reloads the user live, + checks team membership, and mints the same credential ``lite login`` mints.""" + + def __call__( + self, user_id: str, team_id: str | None, / + ) -> Awaitable[MintedProxyCredential | ProxyCredentialMintFailure]: ... + + +class ConsentTeam(BaseModel): + model_config = ConfigDict(frozen=True) + team_id: str = Field(min_length=1) + team_alias: str | None = None + + +class LookupConsentTeams(Protocol): + """Injected lookup of the teams a signed-in user may bind a proxy-API credential to.""" + + def __call__(self, user_id: str, /) -> Awaitable[tuple[ConsentTeam, ...] | ReloadUserFailure]: ... + + +async def _refuse_proxy_credential(user_id: str, team_id: str | None) -> ProxyCredentialMintFailure: + return "unresolvable" + class GatewayDcrClient(BaseModel): """The registration record sealed into a gateway DCR ``client_id``. @@ -173,6 +218,7 @@ class _ConnectFlow(BaseModel): jti: str = Field(min_length=1) exp: int resource_server_id: str | None = None + audience: SessionAudience | None = None class _GatewayAuthCode(BaseModel): @@ -190,6 +236,8 @@ class _GatewayAuthCode(BaseModel): iat: int exp: int resource_server_id: str | None = None + audience: SessionAudience | None = None + team_id: str | None = None def is_gateway_dcr_client_id(client_id: str | None) -> bool: @@ -318,9 +366,9 @@ def _cookie_path_and_secure(request: Request) -> tuple[str, bool]: return parsed.path or "/", parsed.scheme == "https" -def _append_query_params(url: str, params: dict[str, str]) -> str: +def _append_query_params(url: str, params: Iterable[tuple[str, str]]) -> str: parsed: Final = urlparse(url) - query: Final = parse_qsl(parsed.query, keep_blank_values=True) + list(params.items()) + query: Final = (*parse_qsl(parsed.query, keep_blank_values=True), *params) return urlunparse(parsed._replace(query=urlencode(query))) @@ -392,6 +440,155 @@ def aggregate_authorize( section 4.1.2.1 an unvalidated redirect URI must not receive an error redirect, and once the client is at fault there is no trusted place to send the browser. """ + rejected: Final = _rejected_authorize_request( + client_id, redirect_uri, state, code_challenge, code_challenge_method, response_type + ) + if rejected is not None: + return rejected + base_url: Final = get_request_base_url(request) + if session_user_id is None: + return _login_redirect(base_url, request) + scoped_server: Final = resolve_scoped_resource_server(request, resource) + handle: Final = secrets.token_urlsafe(24) + flow: Final = _new_connect_flow( + session_user_id=session_user_id, + client_id=client_id, + redirect_uri=redirect_uri, + state=state, + code_challenge=code_challenge or "", + resource_server_id=scoped_server.server_id if scoped_server is not None else None, + audience=None, + ) + connect_url: Final = _append_query_params( + f"{base_url}/ui/connect", + (("connect_flow", handle), ("connect_client", _origin_only(redirect_uri))), + ) + response: Final = RedirectResponse(connect_url, status_code=303) + _set_flow_cookie(response, request, handle, flow) + return response + + +async def native_client_authorize( + request: Request, + client_id: str, + redirect_uri: str, + state: str, + code_challenge: str | None, + code_challenge_method: str | None, + response_type: str | None, + session_user_id: str | None, + lookup_consent_teams: LookupConsentTeams, +) -> Response: + """The authorize verb for a native client that named the proxy API itself as its + RFC 8707 ``resource``: the same client, redirect, PKCE, and sign-in checks as the + aggregate verb plus a loopback-only redirect (the credential this grant mints is the + user's personal proxy key, which belongs on their own machine and never behind a hosted + callback), then the consent page rendered right here (no connect-page interlude, since + there is no per-server vaulting to do) with the flow sealed into the per-flow cookie + and its handle carried only in the form, never in a URL.""" + rejected: Final = _rejected_authorize_request( + client_id, redirect_uri, state, code_challenge, code_challenge_method, response_type + ) + if rejected is not None: + return rejected + if not is_loopback_redirect_host(urlparse(redirect_uri)): + return _oauth_error(400, "invalid_request", "a proxy-API grant may only redirect to a loopback address") + base_url: Final = get_request_base_url(request) + if session_user_id is None: + return _login_redirect(base_url, request) + teams: Final = await lookup_consent_teams(session_user_id) + if not isinstance(teams, tuple): + return _consent_lookup_failure_response(teams) + handle: Final = secrets.token_urlsafe(24) + flow: Final = _new_connect_flow( + session_user_id=session_user_id, + client_id=client_id, + redirect_uri=redirect_uri, + state=state, + code_challenge=code_challenge or "", + resource_server_id=None, + audience=PROXY_API_AUDIENCE, + ) + page: Final = render_native_client_consent_page( + client_origin=_origin_only(redirect_uri), + user_id=session_user_id, + teams=tuple((team.team_id, team.team_alias or team.team_id) for team in teams), + flow_handle=handle, + complete_url=f"{base_url}/authorize/complete", + ) + response: Final = HTMLResponse(page, headers=_CONSENT_PAGE_HEADERS) + _set_flow_cookie(response, request, handle, flow) + return response + + +_CONSENT_PAGE_HEADERS: Final = MappingProxyType( + { + **TOKEN_NO_CACHE_HEADERS, + "X-Frame-Options": "DENY", + "Content-Security-Policy": "frame-ancestors 'none'", + } +) + +NATIVE_CLIENT_AUTH_CONTRACT_VERSION: Final = 1 +"""The version a native client checks before trusting the rest of the discovery document. +Bump it only when an existing field changes meaning or goes away; adding fields is free.""" + + +class NativeClientAuthContract(TypedDict): + contract_version: ReadOnly[int] + issuer: ReadOnly[str] + authorization_endpoint: ReadOnly[str] + token_endpoint: ReadOnly[str] + registration_endpoint: ReadOnly[str] + revocation_endpoint: ReadOnly[str] + resource: ReadOnly[str] + response_types_supported: ReadOnly[tuple[str, ...]] + grant_types_supported: ReadOnly[tuple[str, ...]] + code_challenge_methods_supported: ReadOnly[tuple[str, ...]] + token_endpoint_auth_methods_supported: ReadOnly[tuple[str, ...]] + revocation_endpoint_auth_methods_supported: ReadOnly[tuple[str, ...]] + + +def native_client_auth_contract(request: Request) -> NativeClientAuthContract: + """The versioned discovery document at ``/.well-known/litellm-cli-auth``: everything a + native client (in any language) needs to run the sign-in without reading LiteLLM + source. ``resource`` is the exact value to send as the RFC 8707 ``resource`` parameter + on authorize and token requests so the grant is issued for the proxy API.""" + base_url: Final = get_request_base_url(request) + contract: Final[NativeClientAuthContract] = { + "contract_version": NATIVE_CLIENT_AUTH_CONTRACT_VERSION, + "issuer": base_url, + "authorization_endpoint": f"{base_url}/authorize", + "token_endpoint": f"{base_url}/token", + "registration_endpoint": f"{base_url}/register", + "revocation_endpoint": f"{base_url}/revoke", + "resource": base_url, + "response_types_supported": ("code",), + "grant_types_supported": ("authorization_code", "refresh_token"), + "code_challenge_methods_supported": ("S256",), + "token_endpoint_auth_methods_supported": ("none",), + "revocation_endpoint_auth_methods_supported": ("none",), + } + return contract + + +def is_proxy_api_resource(request: Request, resource: str | None) -> bool: + """True when the RFC 8707 ``resource`` names the proxy itself (its base URL), which is + how a native client asks for the proxy-API audience rather than an MCP session.""" + if resource is None: + return False + canonical: Final = canonical_resource_uri(resource) + return canonical is not None and canonical == canonicalize_url_identity(get_request_base_url(request)) + + +def _rejected_authorize_request( + client_id: str, + redirect_uri: str, + state: str, + code_challenge: str | None, + code_challenge_method: str | None, + response_type: str | None, +) -> Response | None: client: Final = open_gateway_dcr_client(client_id) if client is None: return _oauth_error(400, "invalid_client", "unknown or malformed client_id") @@ -407,14 +604,25 @@ def aggregate_authorize( ) if len(state) > MAX_STATE_LENGTH: return _oauth_error(400, "invalid_request", f"state must be at most {MAX_STATE_LENGTH} characters") - base_url: Final = get_request_base_url(request) - if session_user_id is None: - login_url: Final = f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}" - return RedirectResponse(login_url, status_code=303) + return None + + +def _login_redirect(base_url: str, request: Request) -> Response: + return_to: Final = urlencode((("return_to", relative_request_url(request)),)) + return RedirectResponse(f"{base_url}/sso/key/generate?{return_to}", status_code=303) + + +def _new_connect_flow( + session_user_id: str, + client_id: str, + redirect_uri: str, + state: str, + code_challenge: str, + resource_server_id: str | None, + audience: SessionAudience | None, +) -> _ConnectFlow: now: Final = datetime.now(timezone.utc) - scoped_server: Final = resolve_scoped_resource_server(request, resource) - handle: Final = secrets.token_urlsafe(24) - flow: Final = _ConnectFlow( + return _ConnectFlow( user_id=session_user_id, client_id=client_id, redirect_uri=redirect_uri, @@ -422,13 +630,12 @@ def aggregate_authorize( code_challenge=code_challenge, jti=secrets.token_urlsafe(24), exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS, - resource_server_id=scoped_server.server_id if scoped_server is not None else None, + resource_server_id=resource_server_id, + audience=audience, ) - connect_url: Final = _append_query_params( - f"{base_url}/ui/connect", - {"connect_flow": handle, "connect_client": _origin_only(redirect_uri)}, - ) - response: Final = RedirectResponse(connect_url, status_code=303) + + +def _set_flow_cookie(response: Response, request: Request, handle: str, flow: _ConnectFlow) -> None: path, secure = _cookie_path_and_secure(request) response.set_cookie( key=_flow_cookie_name(handle), @@ -439,7 +646,18 @@ def aggregate_authorize( httponly=True, samesite="lax", ) - return response + + +def _consent_lookup_failure_response(failure: ReloadUserFailure) -> Response: + 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(403, "access_denied", "the signed-in user is not active") + case _: + assert_never(failure) def _origin_only(url: str) -> str: @@ -455,6 +673,8 @@ async def complete_connect_flow( session_user_id: str | None, cache: DualCache, delivery: str | None = None, + team_id: str | None = None, + decision: str | None = None, ) -> Response: """The deliberate finish step of the connect flow: mint the gateway authorization code and send the browser back to the client. @@ -479,9 +699,16 @@ async def complete_connect_flow( party. Unknown ``delivery`` values are rejected rather than defaulted: a client that asked for manual delivery and got a dead redirect instead would silently lose its code. + + ``decision`` and ``team_id`` come from the native-client consent page. ``"deny"`` + burns the flow and sends the client ``error=access_denied`` so it stops waiting; + ``team_id`` is sealed into the code only for proxy-API flows, where it picks which of + the user's teams the minted credential is attributed to. """ if delivery not in (None, "redirect", "manual"): return _oauth_error(400, "invalid_request", "delivery must be 'redirect' or 'manual'") + if decision not in (None, "approve", "deny"): + return _oauth_error(400, "invalid_request", "decision must be 'approve' or 'deny'") sealed_flow: Final = request.cookies.get(_flow_cookie_name(flow_handle)) if sealed_flow is None: return _oauth_error(400, "invalid_request", "unknown or expired connect flow") @@ -499,6 +726,24 @@ async def complete_connect_flow( 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") + response: Final = ( + _denied_flow_response(flow) if decision == "deny" else _approved_flow_response(flow, delivery, team_id, now) + ) + path, secure = _cookie_path_and_secure(request) + response.delete_cookie(key=_flow_cookie_name(flow_handle), path=path, secure=secure, httponly=True, samesite="lax") + return response + + +def _state_param(flow: _ConnectFlow) -> tuple[tuple[str, str], ...]: + return (("state", flow.state),) if flow.state else () + + +def _denied_flow_response(flow: _ConnectFlow) -> Response: + params: Final = (("error", "access_denied"), *_state_param(flow)) + return RedirectResponse(_append_query_params(flow.redirect_uri, params), status_code=303) + + +def _approved_flow_response(flow: _ConnectFlow, delivery: str | None, team_id: str | None, now: datetime) -> Response: manual_delivery: Final = delivery == "manual" and is_loopback_redirect_host(urlparse(flow.redirect_uri)) code_ttl: Final = MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS if manual_delivery else GATEWAY_AUTH_CODE_TTL_SECONDS code: Final = _seal( @@ -512,16 +757,14 @@ async def complete_connect_flow( iat=int(now.timestamp()), exp=int(now.timestamp()) + code_ttl, resource_server_id=flow.resource_server_id, + audience=flow.audience, + team_id=(team_id or None) if flow.audience == PROXY_API_AUDIENCE else None, ), ) - params: Final = {"code": code, **({"state": flow.state} if flow.state else {})} - callback_url: Final = _append_query_params(flow.redirect_uri, params) - response: Final[Response] = ( - _manual_delivery_response(callback_url) if manual_delivery else RedirectResponse(callback_url, status_code=303) - ) - path, secure = _cookie_path_and_secure(request) - response.delete_cookie(key=_flow_cookie_name(flow_handle), path=path, secure=secure, httponly=True, samesite="lax") - return response + callback_url: Final = _append_query_params(flow.redirect_uri, (("code", code), *_state_param(flow))) + if manual_delivery: + return _manual_delivery_response(callback_url) + return RedirectResponse(callback_url, status_code=303) def _manual_delivery_response(callback_url: str) -> Response: @@ -630,6 +873,37 @@ def _session_token_pair(principal: SessionPrincipal, keys: SessionKeys, now: dat ) +class _ProxyCredentialTokenResponse(TypedDict): + access_token: ReadOnly[str] + token_type: ReadOnly[Literal["Bearer"]] + expires_in: ReadOnly[int] + refresh_token: ReadOnly[str] + user_id: ReadOnly[str] + team_id: ReadOnly[str | None] + + +def _proxy_credential_response( + minted: MintedProxyCredential, principal: SessionPrincipal, keys: SessionKeys, now: datetime +) -> Response: + """The proxy-API token response: the access token is the very credential ``lite + login`` stores (accepted on every proxy route with user and team attribution), and + the refresh token is a gateway-sealed rotating token bound to the team the credential + was minted for, so a renewal keeps the team the user consented to.""" + bound_principal: Final = principal.model_copy(update=MappingProxyType({"team_id": minted.team_id})) + refresh: Final = mint_session_refresh_token(bound_principal, keys, now) + if not isinstance(refresh, MintedSessionToken): + return _oauth_error(500, "server_error", "failed to mint the session credential") + body: Final[_ProxyCredentialTokenResponse] = { + "access_token": minted.key, + "token_type": "Bearer", + "expires_in": minted.expires_in, + "refresh_token": refresh.token.get_secret_value(), + "user_id": minted.user_id, + "team_id": minted.team_id, + } + return JSONResponse(status_code=200, content=body, headers=TOKEN_NO_CACHE_HEADERS) + + def _reload_failure_response(failure: ReloadUserFailure) -> Response: """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.""" @@ -644,6 +918,18 @@ def _reload_failure_response(failure: ReloadUserFailure) -> Response: assert_never(failure) +def _mint_failure_response(failure: ProxyCredentialMintFailure) -> Response: + match failure: + case "not_a_member": + return _oauth_error( + 400, "invalid_grant", "the user is no longer a member of the team this grant was issued for" + ) + case "unavailable" | "unresolvable" | "no_active_key": + return _reload_failure_response(failure) + case _: + assert_never(failure) + + def _resource_conflicts_with_scope( request: Request, resource: str | None, sealed_resource_server_id: str | None ) -> bool: @@ -670,15 +956,26 @@ async def aggregate_token( reload_user: ReloadUser, cache: DualCache, resource: str | None = None, + mint_proxy_credential: MintProxyCredential = _refuse_proxy_credential, ) -> Response: """The aggregate token verb: authorization_code and refresh_token grants for the - identity-only session pair. Every path re-validates the litellm user live before - minting, so a deactivated user cannot obtain or renew a session.""" + identity-only session pair, or for the proxy-API credential when the grant was issued + with that audience. Every path re-validates the litellm user live before minting, so a + deactivated user cannot obtain or renew a session.""" if master_key is None: verbose_logger.error("mcp_gateway_dcr token grant rejected: no master_key configured") return _oauth_error(500, "server_error", "the gateway has no master key configured") keys: Final = session_keys_from_master_key(master_key) now: Final = datetime.now(timezone.utc) + issue: Final = _GrantIssuer( + request=request, + resource=resource, + keys=keys, + now=now, + reload_user=reload_user, + mint_proxy_credential=mint_proxy_credential, + guard=_SingleUseGuard(cache), + ) if grant_type == "authorization_code": return await _authorization_code_grant( request=request, @@ -687,10 +984,8 @@ async def aggregate_token( client_id=client_id, code_verifier=code_verifier, resource=resource, - keys=keys, now=now, - reload_user=reload_user, - guard=_SingleUseGuard(cache), + issue=issue, ) if grant_type == "refresh_token": return await _refresh_token_grant( @@ -700,12 +995,71 @@ async def aggregate_token( resource=resource, keys=keys, now=now, - reload_user=reload_user, - guard=_SingleUseGuard(cache), + issue=issue, ) return _oauth_error(400, "unsupported_grant_type", "grant_type must be authorization_code or refresh_token") +class _GrantIssuer: + """The tail every grant shares once its own proof (code + PKCE, or a refresh token) + has checked out: revalidate the user live, claim the single-use marker, mint. The + claim comes AFTER revalidation and minting so a transient DB 503 never burns a + still-valid code or refresh token, and fails closed when it cannot be recorded.""" + + def __init__( + self, + request: Request, + resource: str | None, + keys: SessionKeys, + now: datetime, + reload_user: ReloadUser, + mint_proxy_credential: MintProxyCredential, + guard: _SingleUseGuard, + ) -> None: + self._request: Final = request + self._resource: Final = resource + self._keys: Final = keys + self._now: Final = now + self._reload_user: Final = reload_user + self._mint_proxy_credential: Final = mint_proxy_credential + self._guard: Final = guard + + async def __call__( + self, principal: SessionPrincipal, claim_key: str, claim_ttl_seconds: int, replayed: str + ) -> Response: + match principal.audience: + case None: + return await self._issue_session_pair(principal, claim_key, claim_ttl_seconds, replayed) + case "proxy_api": + return await self._issue_proxy_credential(principal, claim_key, claim_ttl_seconds, replayed) + case _: + assert_never(principal.audience) + + async def _issue_session_pair( + self, principal: SessionPrincipal, claim_key: str, claim_ttl_seconds: int, replayed: str + ) -> Response: + failure: Final = await self._reload_user(principal.user_id) + if failure is not None: + return _reload_failure_response(failure) + if not await self._guard.claim(claim_key, claim_ttl_seconds): + return _oauth_error(400, "invalid_grant", replayed) + return _session_token_pair(principal, self._keys, self._now) + + async def _issue_proxy_credential( + self, principal: SessionPrincipal, claim_key: str, claim_ttl_seconds: int, replayed: str + ) -> Response: + if self._resource is not None and not is_proxy_api_resource(self._request, self._resource): + return _oauth_error( + 400, "invalid_target", "resource does not match the proxy API this grant was issued for" + ) + minted: Final = await self._mint_proxy_credential(principal.user_id, principal.team_id) + if not isinstance(minted, MintedProxyCredential): + return _mint_failure_response(minted) + if not await self._guard.claim(claim_key, claim_ttl_seconds): + return _oauth_error(400, "invalid_grant", replayed) + return _proxy_credential_response(minted, principal, self._keys, self._now) + + async def _authorization_code_grant( request: Request, code: str | None, @@ -713,10 +1067,8 @@ async def _authorization_code_grant( client_id: str, code_verifier: str | None, resource: str | None, - keys: SessionKeys, now: datetime, - reload_user: ReloadUser, - guard: _SingleUseGuard, + issue: _GrantIssuer, ) -> 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") @@ -733,23 +1085,19 @@ async def _authorization_code_grant( return _oauth_error(400, "invalid_target", "resource does not match the scope this code was issued for") if not _pkce_verifier_matches(code_verifier, parsed.code_challenge): return _oauth_error(400, "invalid_grant", "PKCE verification failed") - # Revalidate the user BEFORE claiming the code, so a transient DB outage (a retryable - # 503) does not consume a still-valid code and force the client to restart sign-in. - failure: Final = 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. The marker's TTL derives from - # the code's own remaining lifetime so it outlives whichever lifetime the code was minted with. - if not await guard.claim( - f"{_USED_CODE_CACHE_PREFIX}{parsed.jti}", - parsed.exp - int(now.timestamp()) + _CLAIM_TTL_BUFFER_SECONDS, - ): - return _oauth_error(400, "invalid_grant", "the authorization code was already used") - return _session_token_pair( - SessionPrincipal(user_id=parsed.user_id, client_id=client_id, resource_server_id=parsed.resource_server_id), - keys, - now, + # The marker's TTL derives from the code's own remaining lifetime so it outlives + # whichever lifetime the code was minted with. + return await issue( + SessionPrincipal( + user_id=parsed.user_id, + client_id=client_id, + resource_server_id=parsed.resource_server_id, + audience=parsed.audience, + team_id=parsed.team_id, + ), + claim_key=f"{_USED_CODE_CACHE_PREFIX}{parsed.jti}", + claim_ttl_seconds=parsed.exp - int(now.timestamp()) + _CLAIM_TTL_BUFFER_SECONDS, + replayed="the authorization code was already used", ) @@ -760,8 +1108,7 @@ async def _refresh_token_grant( resource: str | None, keys: SessionKeys, now: datetime, - reload_user: ReloadUser, - guard: _SingleUseGuard, + issue: _GrantIssuer, ) -> Response: if not refresh_token: return _oauth_error(400, "invalid_request", "refresh_token is required") @@ -770,16 +1117,33 @@ async def _refresh_token_grant( return _oauth_error(400, "invalid_grant", "the refresh token is invalid for this client") if _resource_conflicts_with_scope(request, resource, opened.principal.resource_server_id): return _oauth_error(400, "invalid_target", "resource does not match the scope this token was issued for") - failure: Final = await reload_user(opened.principal.user_id) - if failure is not None: - return _reload_failure_response(failure) # Refresh-token rotation (OAuth 2.0 Security BCP section 4.13): the presented refresh token is - # single-use. Claim its jti before issuing the replacement pair, so a captured or replayed - # refresh token cannot mint a second pair after the legitimate holder rotated. Claimed AFTER - # user revalidation so a transient DB 503 does not burn a still-valid token; a claim that - # cannot be recorded fails closed, exactly like the authorization-code path. - if not await guard.claim( - f"{_USED_REFRESH_CACHE_PREFIX}{opened.jti}", SESSION_REFRESH_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS - ): - return _oauth_error(400, "invalid_grant", "the refresh token was already used") - return _session_token_pair(opened.principal, keys, now) + # single-use, so a captured or replayed refresh token cannot mint a second pair after the + # legitimate holder rotated. + return await issue( + opened.principal, + claim_key=f"{_USED_REFRESH_CACHE_PREFIX}{opened.jti}", + claim_ttl_seconds=SESSION_REFRESH_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS, + replayed="the refresh token was already used", + ) + + +async def revoke_refresh_token(token: str, client_id: str, master_key: str | None, cache: DualCache) -> Response: + """RFC 7009 revocation for the gateway's refresh tokens: burn the presented token's + ``jti`` so neither the holder nor a thief can rotate it again. Access tokens are + stateless and expire on their own (the proxy-API credential within + ``CLI_JWT_EXPIRATION_HOURS``), so per RFC 7009 section 2.2 an unrecognized or already + dead token still answers 200; only an unknown client is refused.""" + if not is_gateway_dcr_client_id(client_id) or open_gateway_dcr_client(client_id) is None: + return _oauth_error(401, "invalid_client", "unknown or malformed client_id") + if master_key is None: + verbose_logger.error("mcp_gateway_dcr revoke rejected: no master_key configured") + return _oauth_error(500, "server_error", "the gateway has no master key configured") + keys: Final = session_keys_from_master_key(master_key) + now: Final = datetime.now(timezone.utc) + opened: Final = open_session_refresh_bearer(token, keys, now, expected_client_id=client_id) + if isinstance(opened, SessionRefreshOpened): + _ = await _SingleUseGuard(cache).claim( + f"{_USED_REFRESH_CACHE_PREFIX}{opened.jti}", SESSION_REFRESH_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS + ) + return Response(content="{}", media_type="application/json", headers=TOKEN_NO_CACHE_HEADERS) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py index 15f5f82c4b6..d6b0a462062 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py @@ -76,6 +76,13 @@ SessionTokenKind = Literal["session", "session_refresh"] on open, so a signature-valid token of one kind cannot be replayed as the other even if its wire prefix is swapped (the prefix is not part of the signed payload; this claim is).""" +SessionAudience = Literal["proxy_api"] +"""The non-MCP audience a session REFRESH token can be minted for. ``None`` (the default and +the only value ever on an MCP wire) means the aggregate MCP gateway; ``"proxy_api"`` means the +refresh grant re-mints the proxy-API CLI credential instead of an MCP session pair. The audience +is read only from the signed claims, never from the request, so a token of one audience can +never be redeemed as the other.""" + class SessionPrincipal(BaseModel): """The litellm user a session token identifies and the DCR client it was issued to. @@ -97,6 +104,8 @@ class SessionPrincipal(BaseModel): user_id: str = Field(min_length=1) client_id: str = Field(min_length=1) resource_server_id: str | None = None + audience: SessionAudience | None = None + team_id: str | None = None class SessionKeys(BaseModel): @@ -194,6 +203,8 @@ class _SessionClaims(BaseModel): user_id: str = Field(min_length=1) client_id: str = Field(min_length=1) resource_server_id: str | None = None + audience: SessionAudience | None = None + team_id: str | None = None def is_session_token(candidate: str) -> bool: @@ -295,6 +306,8 @@ def _mint( user_id=principal.user_id, client_id=principal.client_id, resource_server_id=principal.resource_server_id, + audience=principal.audience, + team_id=principal.team_id, ) token: Final = prefix + jwt.encode( claims.model_dump(exclude_none=True), keys.signing_key.get_secret_value(), algorithm=_SESSION_JWT_ALGORITHM @@ -333,7 +346,11 @@ def _open( return SessionExpired() return OpenedSessionToken( principal=SessionPrincipal( - user_id=claims.user_id, client_id=claims.client_id, resource_server_id=claims.resource_server_id + user_id=claims.user_id, + client_id=claims.client_id, + resource_server_id=claims.resource_server_id, + audience=claims.audience, + team_id=claims.team_id, ), jti=claims.jti, ) diff --git a/litellm/proxy/_experimental/mcp_server/proxy_api_credentials.py b/litellm/proxy/_experimental/mcp_server/proxy_api_credentials.py new file mode 100644 index 00000000000..a24afc529bd --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/proxy_api_credentials.py @@ -0,0 +1,84 @@ +"""The proxy-API side of the native-client sign-in: turning a consented OAuth grant into +the same per-user credential ``lite login`` stores, so the bearer a CLI obtains through +the browser flow is accepted on every proxy route with user and team attribution.""" + +from __future__ import annotations + +from collections.abc import Sequence +from typing import Final + +from litellm.constants import CLI_JWT_EXPIRATION_HOURS +from litellm.proxy._experimental.mcp_server.bridge_token_flow import load_active_user_by_id +from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( + ConsentTeam, + MintedProxyCredential, + ProxyCredentialMintFailure, + ReloadUserFailure, +) +from litellm.proxy._types import LiteLLM_UserTable +from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken +from litellm.proxy.management_endpoints.ui_sso import ( + CliSsoTeamDetail, + fetch_cli_sso_team_details, + selected_cli_sso_team_detail, +) + + +async def lookup_consent_teams(user_id: str) -> tuple[ConsentTeam, ...] | ReloadUserFailure: + user: Final = await load_active_user_by_id(user_id) + if isinstance(user, str): + return user + details: Final = await _team_details(user.teams) + if details is None: + return "unavailable" + return tuple( + ConsentTeam(team_id=detail.team_id, team_alias=detail.team_alias) + for detail in details + if detail.team_id is not None + ) + + +async def mint_proxy_credential( + user_id: str, team_id: str | None +) -> MintedProxyCredential | ProxyCredentialMintFailure: + """Mint the ``lite login`` credential for a consented grant. Membership is checked + live, so a team the user left between consent and redemption (or between refreshes) + refuses the grant instead of minting a credential attributed to a team they are no + longer on. The team is exactly the one the consent page sealed into the grant; nothing + is picked on the user's behalf here, so a refresh can never move the credential. The + user row handed to the minter carries no team list, exactly like ``lite login``'s, so + the minter's own first-team fallback stays inert.""" + user: Final = await load_active_user_by_id(user_id) + if isinstance(user, str): + return user + if user.user_role is None: + return "no_active_key" + if team_id is not None and team_id not in user.teams: + return "not_a_member" + details: Final = await _team_details(user.teams) if team_id is not None else () + if details is None: + return "unavailable" + selected: Final = selected_cli_sso_team_detail(details, team_id) + if selected is None: + return "not_a_member" + key: Final = ExperimentalUIJWTToken.get_cli_jwt_auth_token( + user_info=LiteLLM_UserTable(user_id=user.user_id, user_role=user.user_role, models=user.models), + team_id=team_id, + team_alias=selected.team_alias, + team_models=selected.team_models, + team_model_aliases=selected.team_model_aliases, + ) + return MintedProxyCredential( + key=key, + expires_in=CLI_JWT_EXPIRATION_HOURS * 3600, + user_id=user.user_id, + team_id=team_id, + ) + + +async def _team_details(teams: Sequence[str]) -> tuple[CliSsoTeamDetail, ...] | None: + from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # rebound after startup, so read it per call + + if prisma_client is None: + return None + return await fetch_cli_sso_team_details(prisma_client, teams) diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 262d97e4579..fdd15a89aa5 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -155,10 +155,12 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( "/.well-known/oauth-", "/.well-known/openid-configuration", "/.well-known/jwks.json", + "/.well-known/litellm-cli-auth", "/authorize", "/token", "/callback", "/register", + "/revoke", ), # Catches the /{mcp_server_name}/authorize|token|register variants. path_suffixes=("/authorize", "/token", "/register"), diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 0a0bcf80ee5..aee715686b7 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -11,7 +11,7 @@ import click import requests from rich.console import Console from rich.table import Table -from typing_extensions import NotRequired, TypedDict +from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm.constants import CLI_JWT_EXPIRATION_HOURS from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh @@ -22,6 +22,13 @@ from .claude_settings import ( ClaudeSettingsError, write_claude_settings, ) +from .pkce_login import ( + PkceFailure, + fresh_api_key, + pkce_token_record, + revoke_stored_credential, + run_pkce_login, +) from .private_json import write_private_json @@ -34,6 +41,13 @@ class CliTokenData(TypedDict): auth_header_name: str jwt_token: str timestamp: float + expires_at: ReadOnly[NotRequired[float]] + refresh_token: ReadOnly[NotRequired[str]] + client_id: ReadOnly[NotRequired[str]] + token_endpoint: ReadOnly[NotRequired[str]] + revocation_endpoint: ReadOnly[NotRequired[str]] + resource: ReadOnly[NotRequired[str]] + team_id: ReadOnly[NotRequired[str | None]] class CliTeam(TypedDict, total=False): @@ -79,10 +93,7 @@ class CliAuthResult(TypedDict): # Token storage utilities def get_token_file_path() -> str: """Get the path to store the authentication token""" - home_dir: Final = Path.home() - config_dir: Final = home_dir / ".litellm" - config_dir.mkdir(exist_ok=True) - return str(config_dir / "token.json") + return str(Path.home() / ".litellm" / "token.json") def save_token(token_data: CliTokenData) -> None: @@ -115,11 +126,15 @@ def get_stored_api_key(expected_base_url: str | None = None) -> str | None: If expected_base_url is provided, the key is only returned when it was originally issued for that URL. This prevents credential leakage when the - CLI is pointed at a different (possibly malicious) server. + CLI is pointed at a different (possibly malicious) server. A key obtained by + ``lite login --pkce`` is refreshed here once it nears expiry. """ - from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_key - - return get_litellm_gateway_api_key(expected_base_url=expected_base_url) + token_data: Final = load_token() + if token_data is None: + return None + if expected_base_url is not None and token_data.get("base_url") != expected_base_url.rstrip("/"): + return None + return fresh_api_key(token_data, save_token, requests.Session(), reload=load_token) # Team selection utilities @@ -645,6 +660,27 @@ def _configure_claude_code(base_url: str) -> None: click.echo("Your other Claude Code settings were left untouched. Restart Claude Code to pick this up.") +def _finish_login(base_url: str, api_key: str, config_claude: bool) -> None: + from litellm.proxy.client.cli.interface import show_commands + + click.echo("\nLogin successful!") + click.echo(f"JWT Token: {api_key[:20]}...") + click.echo("You can now use the CLI without specifying --api-key") + if config_claude: + _configure_claude_code(base_url) + click.echo("\n" + "=" * 60) + show_commands() + + +def _pkce_login(base_url: str, config_claude: bool) -> None: + credential: Final = run_pkce_login(base_url, requests.Session(), echo=click.echo) + if isinstance(credential, PkceFailure): + click.echo(f"Authentication failed: {credential.reason}") + return + save_token(pkce_token_record(base_url, credential)) + _finish_login(base_url, credential.access_token, config_claude) + + @click.command(name="login") @click.option( "--config-claude", @@ -655,16 +691,28 @@ def _configure_claude_code(base_url: str) -> None: "Unrelated settings are preserved." ), ) +@click.option( + "--pkce", + is_flag=True, + default=False, + help=( + "Sign in with OAuth authorization code + PKCE through your system browser (loopback redirect), " + "with a refresh token that renews the key automatically. Requires a proxy that serves " + "/.well-known/litellm-cli-auth." + ), +) @click.pass_context -def login(ctx: click.Context, config_claude: bool): +def login(ctx: click.Context, config_claude: bool, pkce: bool) -> None: """Login to LiteLLM proxy using SSO authentication""" from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER - from litellm.proxy.client.cli.interface import show_commands ctx_obj: Final[CliContextObj] = ctx.obj base_url: Final = ctx_obj["base_url"] try: + if pkce: + _pkce_login(base_url, config_claude) + return cli_sso_flow: Final = _start_cli_sso_flow(base_url=base_url) key_id: Final = cli_sso_flow["login_id"] poll_secret: Final = cli_sso_flow["poll_secret"] @@ -704,16 +752,7 @@ def login(ctx: click.Context, config_claude: bool): } ) - click.echo("\nLogin successful!") - click.echo(f"JWT Token: {api_key[:20]}...") - click.echo("You can now use the CLI without specifying --api-key") - - if config_claude: - _configure_claude_code(base_url) - - # Show available commands after successful login - click.echo("\n" + "=" * 60) - show_commands() + _finish_login(base_url, api_key, config_claude) return else: click.echo("Authentication timed out. Please try again.") @@ -738,7 +777,11 @@ def login(ctx: click.Context, config_claude: bool): @click.command(name="logout") def logout(): """Logout and clear stored authentication""" + token_data: Final = load_token() + revocation: Final = revoke_stored_credential(token_data, requests.Session()) if token_data is not None else None clear_token() + if revocation is not None: + click.echo(f"Could not revoke the refresh token on the proxy ({revocation.reason}); it expires on its own.") click.echo("Logged out successfully. Authentication token cleared.") @@ -769,13 +812,13 @@ def print_token(ctx: click.Context): click.echo("Not authenticated for this server. Run 'lite login'.", err=True) sys.exit(1) - if not is_cli_token_fresh(token_data): + if not is_cli_token_fresh(token_data) and "refresh_token" not in token_data: click.echo("Token expired. Run 'lite login' again.", err=True) sys.exit(1) - api_key: Final = token_data.get("key") + api_key: Final = fresh_api_key(token_data, save_token, requests.Session(), reload=load_token) if not api_key: - click.echo("No token available. Run 'lite login'.", err=True) + click.echo("Token expired. Run 'lite login' again.", err=True) sys.exit(1) click.echo(api_key) diff --git a/litellm/proxy/client/cli/commands/pkce_login.py b/litellm/proxy/client/cli/commands/pkce_login.py new file mode 100644 index 00000000000..6d6646a886c --- /dev/null +++ b/litellm/proxy/client/cli/commands/pkce_login.py @@ -0,0 +1,471 @@ +"""Browser sign-in for ``lite login --pkce``: OAuth 2.1 authorization code + PKCE S256 +against the proxy's own authorization server, as a public client on a loopback redirect. +The proxy publishes everything this needs at ``/.well-known/litellm-cli-auth``, so a CLI +in any other language can run the same steps from that document alone.""" + +from __future__ import annotations + +import hashlib +import secrets +import socket +import threading +import time +import webbrowser +from base64 import urlsafe_b64encode +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from http.server import BaseHTTPRequestHandler, HTTPServer +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal, Protocol +from urllib.parse import parse_qs, urlencode, urlparse + +import requests +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError +from typing_extensions import ReadOnly, TypedDict + +if TYPE_CHECKING: + from .auth import CliTokenData + +CLI_AUTH_DISCOVERY_PATH: Final = "/.well-known/litellm-cli-auth" +CALLBACK_PATH: Final = "/callback" +LOGIN_TIMEOUT_SECONDS: Final = 300 +REFRESH_LEEWAY_SECONDS: Final = 60 +_HTTP_TIMEOUT_SECONDS: Final = 15 +_CLIENT_NAME: Final = "litellm-cli" + + +class CliAuthContract(BaseModel): + model_config = ConfigDict(frozen=True) + + contract_version: Literal[1] + issuer: str + authorization_endpoint: str + token_endpoint: str + registration_endpoint: str + revocation_endpoint: str + resource: str + code_challenge_methods_supported: tuple[str, ...] + + +class _RegisteredClient(BaseModel): + model_config = ConfigDict(frozen=True) + + client_id: str = Field(min_length=1) + + +class _TokenResponse(BaseModel): + model_config = ConfigDict(frozen=True) + + access_token: str = Field(min_length=1) + expires_in: int = Field(gt=0) + refresh_token: str = Field(min_length=1) + user_id: str | None = None + team_id: str | None = None + + +@dataclass(frozen=True, slots=True) +class PkceFailure: + reason: str + + +@dataclass(frozen=True, slots=True) +class PkceCredential: + access_token: str + refresh_token: str + expires_at: float + client_id: str + token_endpoint: str + revocation_endpoint: str + resource: str + user_id: str | None + team_id: str | None + + +@dataclass(frozen=True, slots=True) +class CallbackCode: + code: str + + +@dataclass(frozen=True, slots=True) +class CallbackDenied: + error: str + description: str | None + + +CallbackOutcome = CallbackCode | CallbackDenied + + +class Http(Protocol): + def get(self, url: str, *, timeout: float) -> requests.Response: ... + + def post( + self, + url: str, + *, + data: Mapping[str, str] | None = None, + json: Mapping[str, object] | None = None, + timeout: float, + ) -> requests.Response: ... + + +class LoopbackServer(HTTPServer): + """The OS-assigned loopback listener the browser is sent back to. Only the response + carrying the pending sign-in's ``state`` settles it; anything else (a stray request, a + stale tab, an attacker poking the port) gets a 400 and the wait continues. A connection + that opens and then sends nothing is dropped after ``connection_timeout_seconds`` so it + cannot hold the single-threaded wait past its deadline.""" + + def __init__(self, expected_state: str, connection_timeout_seconds: float = 5) -> None: + super().__init__(("127.0.0.1", 0), _CallbackHandler) + self.expected_state: Final = expected_state + self.connection_timeout_seconds: Final = connection_timeout_seconds + self.outcome: CallbackOutcome | None = None + self.timeout = 1 + + @property + def redirect_uri(self) -> str: + return f"http://127.0.0.1:{self.server_address[1]}{CALLBACK_PATH}" + + def get_request(self) -> tuple[socket.socket, object]: + accepted: Final[tuple[socket.socket, object]] = super().get_request() + accepted[0].settimeout(self.connection_timeout_seconds) + return accepted + + def wait( + self, timeout_seconds: float, clock: Callable[[], float] = time.monotonic + ) -> CallbackOutcome | PkceFailure: + deadline: Final = clock() + timeout_seconds + while self.outcome is None: + if clock() >= deadline: + return PkceFailure("timed out waiting for the browser sign-in to finish") + self.handle_request() + return self.outcome + + +class _CallbackHandler(BaseHTTPRequestHandler): + server: LoopbackServer # pyright: ignore[reportIncompatibleVariableOverride] # only ever constructed by LoopbackServer + + def do_GET(self) -> None: + parsed: Final = urlparse(self.path) + if parsed.path != CALLBACK_PATH: + self._respond(404, "Not found.") + return + params: Final = parse_qs(parsed.query) + if _first(params, "state") != self.server.expected_state: + self._respond(400, "This response does not belong to the pending sign-in; still waiting.") + return + error: Final = _first(params, "error") + if error is not None: + self.server.outcome = CallbackDenied(error=error, description=_first(params, "error_description")) + self._respond(200, "Sign-in was not approved. You can close this window.") + return + code: Final = _first(params, "code") + if code is None: + self._respond(400, "The sign-in response carried no authorization code; still waiting.") + return + self.server.outcome = CallbackCode(code=code) + self._respond(200, "Signed in to LiteLLM. You can close this window and return to the terminal.") + + def log_message(self, format: str, *args: object) -> None: + return + + def _respond(self, status: int, text: str) -> None: + body: Final = text.encode("utf-8") + self.send_response(status) + self.send_header("Content-Type", "text/plain; charset=utf-8") + self.send_header("Content-Length", str(len(body))) + self.send_header("Cache-Control", "no-store") + self.end_headers() + self.wfile.write(body) + + +def _first(params: Mapping[str, Sequence[str]], key: str) -> str | None: + values: Final = params.get(key) + return values[0] if values else None + + +def discover_cli_auth(base_url: str, http: Http) -> CliAuthContract | PkceFailure: + url: Final = f"{base_url.rstrip('/')}{CLI_AUTH_DISCOVERY_PATH}" + try: + response: Final = http.get(url, timeout=_HTTP_TIMEOUT_SECONDS) + except requests.RequestException as exc: + return PkceFailure(f"could not reach {url}: {exc}") + if response.status_code != 200: + return PkceFailure( + f"{url} answered {response.status_code}; this proxy version does not support `lite login --pkce`" + ) + try: + contract: Final = CliAuthContract.model_validate(response.json()) + except (ValueError, ValidationError) as exc: + return PkceFailure(f"{url} returned an unsupported discovery document: {exc}") + if "S256" not in contract.code_challenge_methods_supported: + return PkceFailure("the proxy does not support PKCE S256") + return contract + + +class _ClientRegistration(TypedDict): + client_name: ReadOnly[str] + redirect_uris: ReadOnly[tuple[str, ...]] + grant_types: ReadOnly[tuple[str, ...]] + response_types: ReadOnly[tuple[str, ...]] + token_endpoint_auth_method: ReadOnly[Literal["none"]] + + +def _form(**fields: str) -> Mapping[str, str]: + return MappingProxyType(fields) + + +def register_client(contract: CliAuthContract, redirect_uri: str, http: Http) -> str | PkceFailure: + registration: Final[_ClientRegistration] = { + "client_name": _CLIENT_NAME, + "redirect_uris": (redirect_uri,), + "grant_types": ("authorization_code", "refresh_token"), + "response_types": ("code",), + "token_endpoint_auth_method": "none", + } + try: + response: Final = http.post(contract.registration_endpoint, json=registration, timeout=_HTTP_TIMEOUT_SECONDS) + except requests.RequestException as exc: + return PkceFailure(f"client registration failed: {exc}") + if response.status_code not in (200, 201): + return PkceFailure(f"client registration failed with {response.status_code}: {_error_detail(response)}") + try: + return _RegisteredClient.model_validate(response.json()).client_id + except (ValueError, ValidationError) as exc: + return PkceFailure(f"client registration returned an unexpected body: {exc}") + + +def pkce_pair() -> tuple[str, str]: + verifier: Final = secrets.token_urlsafe(64) + digest: Final = hashlib.sha256(verifier.encode("ascii")).digest() + return verifier, urlsafe_b64encode(digest).rstrip(b"=").decode("ascii") + + +def authorize_url(contract: CliAuthContract, client_id: str, redirect_uri: str, state: str, code_challenge: str) -> str: + query: Final = urlencode( + _form( + response_type="code", + client_id=client_id, + redirect_uri=redirect_uri, + state=state, + code_challenge=code_challenge, + code_challenge_method="S256", + resource=contract.resource, + ) + ) + return f"{contract.authorization_endpoint}?{query}" + + +def redeem_code( + contract: CliAuthContract, + client_id: str, + redirect_uri: str, + code: str, + code_verifier: str, + http: Http, + now: Callable[[], float] = time.time, +) -> PkceCredential | PkceFailure: + return _token_request( + token_endpoint=contract.token_endpoint, + revocation_endpoint=contract.revocation_endpoint, + resource=contract.resource, + client_id=client_id, + form=_form( + grant_type="authorization_code", + code=code, + redirect_uri=redirect_uri, + client_id=client_id, + code_verifier=code_verifier, + resource=contract.resource, + ), + http=http, + now=now, + ) + + +def refresh_credential( + token_endpoint: str, + revocation_endpoint: str, + resource: str, + client_id: str, + refresh_token: str, + http: Http, + now: Callable[[], float] = time.time, +) -> PkceCredential | PkceFailure: + return _token_request( + token_endpoint=token_endpoint, + revocation_endpoint=revocation_endpoint, + resource=resource, + client_id=client_id, + form=_form(grant_type="refresh_token", refresh_token=refresh_token, client_id=client_id, resource=resource), + http=http, + now=now, + ) + + +def _token_request( + token_endpoint: str, + revocation_endpoint: str, + resource: str, + client_id: str, + form: Mapping[str, str], + http: Http, + now: Callable[[], float], +) -> PkceCredential | PkceFailure: + try: + response: Final = http.post(token_endpoint, data=form, timeout=_HTTP_TIMEOUT_SECONDS) + except requests.RequestException as exc: + return PkceFailure(f"token request failed: {exc}") + if response.status_code != 200: + return PkceFailure(f"token request failed with {response.status_code}: {_error_detail(response)}") + try: + token: Final = _TokenResponse.model_validate(response.json()) + except (ValueError, ValidationError) as exc: + return PkceFailure(f"token endpoint returned an unexpected body: {exc}") + return PkceCredential( + access_token=token.access_token, + refresh_token=token.refresh_token, + expires_at=now() + token.expires_in, + client_id=client_id, + token_endpoint=token_endpoint, + revocation_endpoint=revocation_endpoint, + resource=resource, + user_id=token.user_id, + team_id=token.team_id, + ) + + +def revoke_credential(revocation_endpoint: str, client_id: str, refresh_token: str, http: Http) -> PkceFailure | None: + try: + response: Final = http.post( + revocation_endpoint, + data=_form(token=refresh_token, token_type_hint="refresh_token", client_id=client_id), + timeout=_HTTP_TIMEOUT_SECONDS, + ) + except requests.RequestException as exc: + return PkceFailure(f"revocation request failed: {exc}") + if response.status_code != 200: + return PkceFailure(f"revocation failed with {response.status_code}: {_error_detail(response)}") + return None + + +_ERROR_BODY: Final = TypeAdapter(Mapping[str, object]) + + +def _error_detail(response: requests.Response) -> str: + try: + body: Final = _ERROR_BODY.validate_json(response.content) + except ValidationError: + return response.text[:200] + return str(body.get("error_description") or body.get("error") or body.get("detail") or body)[:200] + + +def run_pkce_login( + base_url: str, + http: Http, + open_browser: Callable[[str], object] = webbrowser.open, + echo: Callable[[str], None] = print, + timeout_seconds: float = LOGIN_TIMEOUT_SECONDS, +) -> PkceCredential | PkceFailure: + contract: Final = discover_cli_auth(base_url, http) + if isinstance(contract, PkceFailure): + return contract + state: Final = secrets.token_urlsafe(32) + verifier, challenge = pkce_pair() + with LoopbackServer(state) as server: + client_id: Final = register_client(contract, server.redirect_uri, http) + if isinstance(client_id, PkceFailure): + return client_id + url: Final = authorize_url(contract, client_id, server.redirect_uri, state, challenge) + echo(f"Opening browser to: {url}") + echo("Approve the sign-in in your browser. Waiting...") + threading.Thread(target=open_browser, args=(url,), name="lite-login-browser", daemon=True).start() + outcome: Final = server.wait(timeout_seconds) + match outcome: + case PkceFailure(): + return outcome + case CallbackDenied(): + return PkceFailure(f"sign-in was not approved ({outcome.error}): {outcome.description or 'no details'}") + case CallbackCode(): + return redeem_code(contract, client_id, server.redirect_uri, outcome.code, verifier, http) + + +def pkce_token_record(base_url: str, credential: PkceCredential) -> CliTokenData: + record: Final[CliTokenData] = { + "base_url": base_url.rstrip("/"), + "key": credential.access_token, + "user_id": credential.user_id or "cli-user", + "user_email": "unknown", + "user_role": "cli", + "auth_header_name": "Authorization", + "jwt_token": "", + "timestamp": time.time(), + "expires_at": credential.expires_at, + "refresh_token": credential.refresh_token, + "client_id": credential.client_id, + "token_endpoint": credential.token_endpoint, + "revocation_endpoint": credential.revocation_endpoint, + "resource": credential.resource, + "team_id": credential.team_id, + } + return record + + +def fresh_api_key( + token_data: Mapping[str, object], + save: Callable[[CliTokenData], None], + http: Http, + *, + reload: Callable[[], Mapping[str, object] | None], + now: Callable[[], float] = time.time, +) -> str | None: + """The stored key, refreshed first when it is about to expire and a refresh token is + on file. The rotated pair is saved before the new key is returned, so a crash after + this point never strands the CLI with a burned refresh token. A refresh that fails + reads the record again, because a sibling ``lite`` process may have rotated the pair + first, in which case the key it saved is the live one. A record without + ``expires_at`` (the classic ``lite login`` credential) is returned as stored.""" + key: Final = token_data.get("key") + if not isinstance(key, str) or not key: + return None + expires_at: Final = token_data.get("expires_at") + if not isinstance(expires_at, (int, float)): + return key + if now() < expires_at - REFRESH_LEEWAY_SECONDS: + return key + still_valid: Final = key if now() < expires_at else None + refresh_inputs: Final = _refresh_inputs(token_data) + if refresh_inputs is None: + return still_valid + refreshed: Final = refresh_credential(*refresh_inputs, http=http, now=now) + if isinstance(refreshed, PkceFailure): + return _key_rotated_by_a_sibling(reload(), token_data.get("refresh_token")) or still_valid + base_url: Final = token_data.get("base_url") + save(pkce_token_record(base_url if isinstance(base_url, str) else "", refreshed)) + return refreshed.access_token + + +def _key_rotated_by_a_sibling(record: Mapping[str, object] | None, sent_refresh_token: object) -> str | None: + if record is None or record.get("refresh_token") == sent_refresh_token: + return None + key: Final = record.get("key") + return key if isinstance(key, str) and key else None + + +def _refresh_inputs(token_data: Mapping[str, object]) -> tuple[str, str, str, str, str] | None: + values: Final = tuple( + token_data.get(field) + for field in ("token_endpoint", "revocation_endpoint", "resource", "client_id", "refresh_token") + ) + if not all(isinstance(value, str) and value for value in values): + return None + token_endpoint, revocation_endpoint, resource, client_id, refresh_token = values + return str(token_endpoint), str(revocation_endpoint), str(resource), str(client_id), str(refresh_token) + + +def revoke_stored_credential(token_data: Mapping[str, object], http: Http) -> PkceFailure | None: + refresh_inputs: Final = _refresh_inputs(token_data) + if refresh_inputs is None: + return None + _, revocation_endpoint, _, client_id, refresh_token = refresh_inputs + return revoke_credential(revocation_endpoint, client_id, refresh_token, http) diff --git a/litellm/proxy/common_utils/html_forms/native_client_consent.py b/litellm/proxy/common_utils/html_forms/native_client_consent.py new file mode 100644 index 00000000000..dac92c4e787 --- /dev/null +++ b/litellm/proxy/common_utils/html_forms/native_client_consent.py @@ -0,0 +1,91 @@ +from collections.abc import Sequence +from html import escape +from typing import Final + +from litellm.constants import CLI_JWT_EXPIRATION_HOURS + + +def render_native_client_consent_page( + *, + client_origin: str, + user_id: str, + teams: Sequence[tuple[str, str]], + flow_handle: str, + complete_url: str, +) -> str: + """The consent page a native client's sign-in lands on: who is signed in, which + loopback client asked, which team the credential is attributed to, and an explicit + Approve or Deny that POSTs back to ``complete_url``. Every value is client- or + user-influenced and HTML-escaped; the flow handle travels only in the form body.""" + return f""" + + + + + +Authorize CLI access - LiteLLM + + + +
+

Authorize CLI access

+

A command-line client at {escape(client_origin)} wants to call LiteLLM as {escape(user_id)}.

+

Approving issues it a personal credential that expires within {CLI_JWT_EXPIRATION_HOURS} hours. lite logout stops it from being renewed. Only approve if you started this sign-in yourself.

+
+ +{_team_field(teams)} +
+ + +
+
+
+ + +""" + + +def _team_field(teams: Sequence[tuple[str, str]]) -> str: + if not teams: + return "" + if len(teams) == 1: + team_id, team_label = teams[0] + return ( + f'' + f"

Requests are attributed to team {escape(team_label)}.

" + ) + options: Final = "".join( + f'' for team_id, team_label in teams + ) + return ( + f'' + ) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 46af5dd80e1..1ebcb53fd6b 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -270,7 +270,7 @@ class _TeamRowGrants(BaseModel): litellm_model_table: _TeamModelAliasTable | None = None -class _CliSsoTeamDetail(BaseModel): +class CliSsoTeamDetail(BaseModel): """The per-team snapshot cached in the CLI SSO flow and echoed to the CLI on poll.""" team_id: str | None = None @@ -279,8 +279,8 @@ class _CliSsoTeamDetail(BaseModel): team_model_aliases: Mapping[str, str] | None = None -_CLI_SSO_TEAM_DETAILS_ADAPTER: Final = TypeAdapter(tuple[_CliSsoTeamDetail, ...]) -_TEAMLESS_CLI_SSO_TEAM_DETAIL: Final = _CliSsoTeamDetail(team_models=()) +_CLI_SSO_TEAM_DETAILS_ADAPTER: Final = TypeAdapter(tuple[CliSsoTeamDetail, ...]) +_TEAMLESS_CLI_SSO_TEAM_DETAIL: Final = CliSsoTeamDetail(team_models=()) class _CustomSsoCall(Protocol): @@ -2192,10 +2192,10 @@ async def _build_cli_sso_user_defined_values( ) -def _cli_sso_team_detail(team_row: Mapping[str, object]) -> _CliSsoTeamDetail: +def _cli_sso_team_detail(team_row: Mapping[str, object]) -> CliSsoTeamDetail: team: Final = _TeamRowGrants.model_validate(team_row) alias_table: Final = team.litellm_model_table - return _CliSsoTeamDetail( + return CliSsoTeamDetail( team_id=team.team_id, team_alias=team.team_alias, team_models=team.models, @@ -2203,10 +2203,10 @@ def _cli_sso_team_detail(team_row: Mapping[str, object]) -> _CliSsoTeamDetail: ) -async def _fetch_cli_sso_team_details( +async def fetch_cli_sso_team_details( prisma_client: PrismaClient, teams: Sequence[str], -) -> tuple[_CliSsoTeamDetail, ...] | None: +) -> tuple[CliSsoTeamDetail, ...] | None: """``None`` means the lookup itself failed, which is not the same as the user having no teams.""" if not teams: return () @@ -2221,7 +2221,7 @@ async def _fetch_cli_sso_team_details( return tuple(_cli_sso_team_detail(team_row.model_dump()) for team_row in prisma_teams) -def _cli_sso_session_teams(team_details: Sequence[_CliSsoTeamDetail]) -> list[str]: +def _cli_sso_session_teams(team_details: Sequence[CliSsoTeamDetail]) -> list[str]: """The teams a login may bind to: only those whose row still exists. A team deleted out from under a membership, which is what deleting an organization @@ -2231,7 +2231,7 @@ def _cli_sso_session_teams(team_details: Sequence[_CliSsoTeamDetail]) -> list[st return [detail.team_id for detail in team_details if detail.team_id is not None] -def _selected_cli_sso_team_detail(team_details: object, team_id: str | None) -> _CliSsoTeamDetail | None: +def selected_cli_sso_team_detail(team_details: object, team_id: str | None) -> CliSsoTeamDetail | None: """``None`` means the team's grants are unknown. An empty grant is a real value meaning unrestricted, so an unknown one must not be minted as empty.""" if team_id is None: @@ -2282,7 +2282,7 @@ async def _complete_cli_sso_callback_session( if hasattr(user_info, "teams") and user_info.teams: teams = user_info.teams if isinstance(user_info.teams, list) else [] - team_details: Final = await _fetch_cli_sso_team_details(prisma_client=prisma_client, teams=teams) + team_details: Final = await fetch_cli_sso_team_details(prisma_client=prisma_client, teams=teams) if team_details is None: raise HTTPException( status_code=500, @@ -2483,7 +2483,7 @@ async def cli_poll_key( # If no team_id provided and user has 0 or 1 team, use first team (or None) team_id = user_teams[0] if len(user_teams) > 0 else None - selected_team: Final = _selected_cli_sso_team_detail( + selected_team: Final = selected_cli_sso_team_detail( team_details=user_team_details, team_id=team_id, ) diff --git a/tests/test_litellm/litellm_core_utils/test_cli_token_utils.py b/tests/test_litellm/litellm_core_utils/test_cli_token_utils.py index 27fc5eb4bd0..5593211ba6f 100644 --- a/tests/test_litellm/litellm_core_utils/test_cli_token_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_cli_token_utils.py @@ -5,6 +5,7 @@ Unit tests for CLI token utilities import json import os import tempfile +import time from pathlib import Path from unittest.mock import mock_open, patch @@ -87,3 +88,29 @@ class TestCLITokenUtils: result = get_litellm_gateway_api_key() assert result is None + + +class TestIsCliTokenFreshWithExpiresAt: + """A ``lite login --pkce`` record carries the proxy's own ``expires_at``, which wins + over the age-based guess made from ``timestamp``.""" + + def test_future_expiry_is_fresh(self): + from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh + + assert is_cli_token_fresh({"expires_at": time.time() + 3600, "timestamp": 0}) is True + + def test_expiry_inside_the_buffer_is_stale(self): + from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh + + assert is_cli_token_fresh({"expires_at": time.time() + 100}) is False + assert is_cli_token_fresh({"expires_at": time.time() + 100}, buffer_hours=0) is True + + def test_past_expiry_is_stale_even_with_a_fresh_timestamp(self): + from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh + + assert is_cli_token_fresh({"expires_at": time.time() - 1, "timestamp": time.time()}) is False + + def test_non_numeric_expiry_falls_back_to_the_timestamp(self): + from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh + + assert is_cli_token_fresh({"expires_at": "soon", "timestamp": time.time()}) is True diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_session_token.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_session_token.py index a43592ebe18..36280530eac 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_session_token.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_session_token.py @@ -212,3 +212,55 @@ def test_minted_token_repr_never_leaks_value(): minted = mint_session_token(PRINCIPAL, KEYS, NOW) assert isinstance(minted, MintedSessionToken) assert minted.token.get_secret_value() not in repr(minted) + + +def _decoded_claims(token: str, prefix: str) -> dict: + return jwt.decode( + token.removeprefix(prefix), + KEYS.signing_key.get_secret_value(), + algorithms=["HS256"], + options={"verify_exp": False}, + ) + + +def test_mcp_principal_wire_claims_carry_no_audience_or_team_keys(): + access_claims = _decoded_claims(_mint_access(), SESSION_TOKEN_PREFIX) + refresh_claims = _decoded_claims(_mint_refresh(), SESSION_REFRESH_PREFIX) + for claims in (access_claims, refresh_claims): + assert "audience" not in claims + assert "team_id" not in claims + + +def test_legacy_signed_claims_open_with_no_audience_and_no_team(): + opened = open_session_token(_sign_claims(_valid_claims()), KEYS, NOW) + assert isinstance(opened, OpenedSessionToken) + assert opened.principal.audience is None + assert opened.principal.team_id is None + + +def test_proxy_api_audience_and_team_round_trip_through_the_refresh_token(): + principal = SessionPrincipal(user_id="user-123", client_id="llm_client_abc", audience="proxy_api", team_id="team-b") + minted = mint_session_refresh_token(principal, KEYS, NOW) + assert isinstance(minted, MintedSessionToken) + token = minted.token.get_secret_value() + claims = _decoded_claims(token, SESSION_REFRESH_PREFIX) + assert claims["audience"] == "proxy_api" + assert claims["team_id"] == "team-b" + opened = open_session_refresh_token(token, KEYS, NOW) + assert isinstance(opened, OpenedSessionToken) + assert opened.principal == principal + + +def test_signed_claims_with_an_unknown_audience_are_rejected(): + token = _sign_claims(_valid_claims(audience="bogus")) + assert isinstance(open_session_token(token, KEYS, NOW), SessionMalformed) + + +def test_signed_claims_with_a_non_string_team_are_rejected(): + token = _sign_claims(_valid_claims(team_id=42)) + assert isinstance(open_session_token(token, KEYS, NOW), SessionMalformed) + + +def test_principal_rejects_an_unknown_audience_at_construction(): + with pytest.raises(ValidationError): + SessionPrincipal(user_id="user-123", client_id="llm_client_abc", audience="mcp") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 424b993de85..bdaf1458fb0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -1,6 +1,9 @@ """Tests for MCP OAuth discoverable endpoints""" +import hashlib import json +import time +from base64 import urlsafe_b64encode from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -9432,3 +9435,229 @@ async def test_upstream_resource_sent_on_dcr_bridge_relay_authorize(): query = await _authorize_query(server) assert query["resource"] == ["https://mcp.example.com/mcp"] assert query["client_id"] == ["caller-client"] + + +def _s256(verifier: str) -> str: + return urlsafe_b64encode(hashlib.sha256(verifier.encode("ascii")).digest()).rstrip(b"=").decode("ascii") + + +_NATIVE_CLIENT_MASTER_KEY = "sk-test-salt-for-LIT-5874" + + +def _native_client_app(monkeypatch): + """The unauthenticated discoverable router served over TestClient with a signed UI session + cookie available, plus fakes for the two database-backed hooks the native-client flow calls.""" + import jwt + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ConsentTeam, MintedProxyCredential + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + monkeypatch.setenv("LITELLM_SALT_KEY", _NATIVE_CLIENT_MASTER_KEY) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", _NATIVE_CLIENT_MASTER_KEY, raising=False) + minted = [] + + async def fake_mint(user_id, team_id): + minted.append((user_id, team_id)) + return MintedProxyCredential(key=f"sk-cli-{len(minted)}", expires_in=3600, user_id=user_id, team_id=team_id) + + async def fake_lookup(user_id): + return (ConsentTeam(team_id="team-a", team_alias="Team A"), ConsentTeam(team_id="team-b")) + + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.mint_proxy_credential", fake_mint + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.lookup_consent_teams", fake_lookup + ) + global_mcp_server_manager.registry.clear() + app = FastAPI() + app.include_router(router) + client = TestClient(app) + session_cookie = jwt.encode( + {"user_id": "u1", "login_method": "username_password", "exp": int(time.time()) + 600}, + _NATIVE_CLIENT_MASTER_KEY, + algorithm="HS256", + ) + return client, session_cookie, minted + + +def _consent_flow_handle(page: str) -> str: + import re + + match = re.search(r'name="flow" value="([^"]+)"', page) + assert match is not None, page + return match.group(1) + + +def test_native_client_login_walks_discovery_consent_token_refresh_and_revoke(monkeypatch): + """The whole ``lite login --pkce`` server side over the real router: a Go CLI reads the versioned + discovery document, registers a loopback public client, the signed-in user consents to a team, + the code redeems for the ``lite login`` credential, the refresh token rotates, and revocation + kills it.""" + from http.cookies import SimpleCookie + from urllib.parse import parse_qs, urlparse + + client, session_cookie, minted = _native_client_app(monkeypatch) + redirect_uri = "http://127.0.0.1:51234/callback" + + discovery = client.get("/.well-known/litellm-cli-auth") + assert discovery.status_code == 200 + assert discovery.headers["cache-control"] == "no-store" + contract = discovery.json() + assert contract["contract_version"] == 1 + assert contract["resource"] == "http://testserver" + assert contract["code_challenge_methods_supported"] == ["S256"] + assert contract["token_endpoint_auth_methods_supported"] == ["none"] + for endpoint in ("authorization_endpoint", "token_endpoint", "registration_endpoint", "revocation_endpoint"): + assert contract[endpoint].startswith("http://testserver/") + + registered = client.post( + contract["registration_endpoint"], + json={ + "client_name": "litellm-cli", + "redirect_uris": [redirect_uri], + "grant_types": ["authorization_code", "refresh_token"], + "response_types": ["code"], + "token_endpoint_auth_method": "none", + }, + ) + assert registered.status_code == 201 + client_id = registered.json()["client_id"] + verifier = "v" * 43 + authorize_params = { + "response_type": "code", + "client_id": client_id, + "redirect_uri": redirect_uri, + "state": "cli-state", + "code_challenge": _s256(verifier), + "code_challenge_method": "S256", + "resource": contract["resource"], + } + + anonymous = client.get(contract["authorization_endpoint"], params=authorize_params, follow_redirects=False) + assert anonymous.status_code == 303 + login_target = urlparse(anonymous.headers["location"]) + assert login_target.path == "/sso/key/generate" + assert parse_qs(login_target.query)["return_to"][0].startswith("/authorize?") + + client.cookies.set("token", session_cookie) + consent = client.get(contract["authorization_endpoint"], params=authorize_params, follow_redirects=False) + assert consent.status_code == 200 + assert consent.headers["x-frame-options"] == "DENY" + assert consent.headers["cache-control"] == "no-store" + assert "http://127.0.0.1:51234" in consent.text + assert '' in consent.text + jar = SimpleCookie() + jar.load(consent.headers["set-cookie"]) + assert all(morsel["httponly"] for morsel in jar.values()) + + denied = client.post( + "/authorize/complete", + data={"flow": _consent_flow_handle(consent.text), "decision": "deny", "team_id": "team-a"}, + follow_redirects=False, + ) + assert denied.status_code == 303 + denied_query = parse_qs(urlparse(denied.headers["location"]).query) + assert denied.headers["location"].startswith(redirect_uri) + assert denied_query["error"] == ["access_denied"] + assert denied_query["state"] == ["cli-state"] + assert minted == [] + + consent_again = client.get(contract["authorization_endpoint"], params=authorize_params, follow_redirects=False) + approved = client.post( + "/authorize/complete", + data={"flow": _consent_flow_handle(consent_again.text), "decision": "approve", "team_id": "team-b"}, + follow_redirects=False, + ) + assert approved.status_code == 303 + assert approved.headers["location"].startswith(redirect_uri) + approved_query = parse_qs(urlparse(approved.headers["location"]).query) + assert approved_query["state"] == ["cli-state"] + code = approved_query["code"][0] + + token = client.post( + contract["token_endpoint"], + data={ + "grant_type": "authorization_code", + "code": code, + "redirect_uri": redirect_uri, + "client_id": client_id, + "code_verifier": verifier, + "resource": contract["resource"], + }, + ) + assert token.status_code == 200, token.text + assert token.headers["cache-control"] == "no-store" + body = token.json() + assert body["access_token"] == "sk-cli-1" + assert body["token_type"] == "Bearer" + assert body["expires_in"] == 3600 + assert body["user_id"] == "u1" + assert body["team_id"] == "team-b" + assert body["refresh_token"].startswith("llm_srefresh_") + assert minted == [("u1", "team-b")] + + refreshed = client.post( + contract["token_endpoint"], + data={ + "grant_type": "refresh_token", + "refresh_token": body["refresh_token"], + "client_id": client_id, + "resource": contract["resource"], + }, + ) + assert refreshed.status_code == 200, refreshed.text + assert refreshed.json()["access_token"] == "sk-cli-2" + assert refreshed.json()["team_id"] == "team-b" + assert refreshed.json()["refresh_token"] != body["refresh_token"] + assert minted == [("u1", "team-b"), ("u1", "team-b")] + + revoked = client.post( + contract["revocation_endpoint"], + data={"token": refreshed.json()["refresh_token"], "token_type_hint": "refresh_token", "client_id": client_id}, + ) + assert revoked.status_code == 200 + assert revoked.json() == {} + + after_revoke = client.post( + contract["token_endpoint"], + data={ + "grant_type": "refresh_token", + "refresh_token": refreshed.json()["refresh_token"], + "client_id": client_id, + "resource": contract["resource"], + }, + ) + assert after_revoke.status_code == 400 + assert after_revoke.json()["error"] == "invalid_grant" + + stranger = client.post( + contract["revocation_endpoint"], data={"token": "whatever", "client_id": "llm_dcrc_not_a_client"} + ) + assert stranger.status_code == 401 + assert stranger.json()["error"] == "invalid_client" + + +def test_native_client_authorize_without_the_proxy_resource_keeps_the_mcp_flow(monkeypatch): + """A registered client asking for the MCP resource (or no resource) never sees the consent + page, so existing MCP clients are untouched by the native-client arm.""" + client, session_cookie, minted = _native_client_app(monkeypatch) + registered = client.post("/register", json={"redirect_uris": ["http://127.0.0.1:51234/callback"]}) + client.cookies.set("token", session_cookie) + for resource in (None, "http://testserver/mcp"): + params = { + "response_type": "code", + "client_id": registered.json()["client_id"], + "redirect_uri": "http://127.0.0.1:51234/callback", + "state": "s", + "code_challenge": _s256("v" * 43), + "code_challenge_method": "S256", + **({"resource": resource} if resource else {}), + } + response = client.get("/authorize", params=params, follow_redirects=False) + assert 'name="decision"' not in response.text + assert "team-b" not in response.text + assert minted == [] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py index cc65970a180..4e665c552f3 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py @@ -1,8 +1,8 @@ """Tests for the aggregate gateway DCR flow (register, authorize, complete, token).""" import hashlib -import html import json +import re from base64 import urlsafe_b64encode from datetime import datetime, timedelta, timezone from http.cookies import SimpleCookie @@ -13,12 +13,13 @@ from starlette.requests import Request from litellm.caching.caching import DualCache from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( + _AUTH_CODE_DEBUG_KEY, CONNECT_FLOW_COOKIE_PREFIX, GATEWAY_AUTH_CODE_PREFIX, GATEWAY_AUTH_CODE_TTL_SECONDS, - GATEWAY_DCR_CLIENT_ID_PREFIX, MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS, - _AUTH_CODE_DEBUG_KEY, + ConsentTeam, + MintedProxyCredential, _GatewayAuthCode, _open_sealed, _seal, @@ -26,14 +27,21 @@ from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( aggregate_token, complete_connect_flow, is_gateway_dcr_client_id, + is_proxy_api_resource, + native_client_auth_contract, + native_client_authorize, open_gateway_dcr_client, register_aggregate_client, + revoke_refresh_token, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + SessionBearerAdmitted, + SessionRefreshOpened, + open_session_refresh_bearer, resolve_session_bearer, session_keys_from_master_key, - SessionBearerAdmitted, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import SESSION_REFRESH_PREFIX MASTER_KEY = "sk-gateway-dcr-flow-tests" REDIRECT_URI = "https://claude.ai/api/mcp/auth_callback" @@ -888,9 +896,14 @@ async def test_scoped_authorize_runs_connect_page_with_sealed_scope(): assert response.status_code == 303 assert "/ui/connect" in response.headers["location"] _, cookies = _flow_cookie_from(response) - assert _sealed_wire_json(next(iter(cookies.values())), "", "gateway_connect_flow")["resource_server_id"] == "github-id" + assert ( + _sealed_wire_json(next(iter(cookies.values())), "", "gateway_connect_flow")["resource_server_id"] == "github-id" + ) code = await _finish_connect_page(response) - assert _sealed_wire_json(code, GATEWAY_AUTH_CODE_PREFIX, "gateway_authorization_code")["resource_server_id"] == "github-id" + assert ( + _sealed_wire_json(code, GATEWAY_AUTH_CODE_PREFIX, "gateway_authorization_code")["resource_server_id"] + == "github-id" + ) token_response = await _redeem(code, client_id) assert token_response.status_code == 200 principal = _opened_principal(json.loads(token_response.body)) @@ -1039,3 +1052,542 @@ async def test_resource_resolution_is_identity_not_ip_filtered_access(): result = resolve_scoped_resource_server(_request(), SCOPED_RESOURCE) assert result is not None manager.get_mcp_server_by_name.assert_called_once_with("github") + + +LOOPBACK_REDIRECT_URI = "http://127.0.0.1:51234/callback" +PROXY_API_RESOURCE = "https://llm.example.com" +CONSENT_TEAMS = (ConsentTeam(team_id="team-a", team_alias="Team A"), ConsentTeam(team_id="team-b")) + + +class _Minter: + def __init__(self, result=None): + self.calls = [] + self.result = result + + async def __call__(self, user_id, team_id): + self.calls.append((user_id, team_id)) + if self.result is not None: + return self.result + return MintedProxyCredential(key=f"sk-cli-{user_id}", expires_in=3600, user_id=user_id, team_id=team_id) + + +class _ConsentTeams: + def __init__(self, result=CONSENT_TEAMS): + self.calls = [] + self.result = result + + async def __call__(self, user_id): + self.calls.append(user_id) + return self.result + + +async def _native_authorize(client_id, session_user_id="u1", lookup=None, **overrides): + arguments = { + "request": _request(query=f"resource={PROXY_API_RESOURCE}"), + "client_id": client_id, + "redirect_uri": LOOPBACK_REDIRECT_URI, + "state": "client-state-123", + "code_challenge": CODE_CHALLENGE, + "code_challenge_method": "S256", + "response_type": "code", + "session_user_id": session_user_id, + "lookup_consent_teams": lookup if lookup is not None else _ConsentTeams(), + } + return await native_client_authorize(**{**arguments, **overrides}) + + +def _consent_cookie_from(response) -> tuple: + match = re.search(r'name="flow" value="([^"]+)"', response.body.decode()) + assert match is not None + handle = match.group(1) + cookie = SimpleCookie() + cookie.load(response.headers["set-cookie"]) + name = f"{CONNECT_FLOW_COOKIE_PREFIX}{handle}" + return handle, {name: cookie[name].value} + + +async def _complete_consent(consent, cache=None, session_user_id="u1", **overrides): + handle, cookies = _consent_cookie_from(consent) + return await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id=session_user_id, + cache=cache or DualCache(), + **overrides, + ) + + +def _code_from(response) -> str: + return parse_qs(urlparse(response.headers["location"]).query)["code"][0] + + +async def _native_code(client_id, team_id="team-b", cache=None) -> str: + approved = await _complete_consent( + await _native_authorize(client_id), cache=cache, decision="approve", team_id=team_id + ) + assert approved.status_code == 303 + return _code_from(approved) + + +async def _redeem_native(code, client_id, minter, cache=None, resource=PROXY_API_RESOURCE, **overrides): + return await _redeem( + code, + client_id, + cache=cache, + redirect_uri=LOOPBACK_REDIRECT_URI, + resource=resource, + mint_proxy_credential=minter, + **overrides, + ) + + +async def _refresh_native(refresh_token, client_id, minter, cache, **overrides): + return await _redeem_native( + None, client_id, minter, cache=cache, grant_type="refresh_token", refresh_token=refresh_token, **overrides + ) + + +def _opened_refresh(refresh_token, client_id): + opened = open_session_refresh_bearer( + refresh_token, + session_keys_from_master_key(MASTER_KEY), + datetime.now(timezone.utc), + expected_client_id=client_id, + ) + assert isinstance(opened, SessionRefreshOpened) + return opened.principal + + +@pytest.mark.asyncio +async def test_native_authorize_renders_consent_page_and_sets_flow_cookie(): + """A native client (RFC 8707 resource = the proxy itself) gets the server-rendered consent + page instead of the MCP connect-page redirect: the flow handle rides only in the hidden + field, the sealed flow in an HttpOnly cookie, and the page can never be framed or cached.""" + client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"] + lookup = _ConsentTeams() + response = await _native_authorize(client_id, lookup=lookup) + assert response.status_code == 200 + assert response.headers["content-type"].startswith("text/html") + assert response.headers["cache-control"] == "no-store" + assert response.headers["x-frame-options"] == "DENY" + assert response.headers["content-security-policy"] == "frame-ancestors 'none'" + assert lookup.calls == ["u1"] + body = response.body.decode() + assert "http://127.0.0.1:51234" in body + assert "/callback" not in body + assert "u1" in body + assert '' in body + assert '' in body + assert 'action="https://llm.example.com/authorize/complete"' in body + handle, cookies = _consent_cookie_from(response) + assert "httponly" in response.headers["set-cookie"].lower() + flow = _sealed_wire_json(next(iter(cookies.values())), "", "gateway_connect_flow") + assert flow["audience"] == "proxy_api" + assert flow["client_id"] == client_id + assert flow["redirect_uri"] == LOOPBACK_REDIRECT_URI + assert flow["user_id"] == "u1" + assert "resource_server_id" not in flow + assert handle not in body.replace(f'value="{handle}"', "") + + +@pytest.mark.asyncio +async def test_native_authorize_without_session_redirects_to_login_before_any_lookup(): + client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"] + lookup = _ConsentTeams() + response = await _native_authorize(client_id, session_user_id=None, lookup=lookup) + assert response.status_code == 303 + location = response.headers["location"] + assert location.startswith("https://llm.example.com/sso/key/generate?return_to=") + assert "return_to=%2Fauthorize%3Fresource%3D" in location + assert lookup.calls == [] + assert "set-cookie" not in response.headers + + +@pytest.mark.asyncio +async def test_native_authorize_validation_failures_never_reach_consent(): + client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"] + lookup = _ConsentTeams() + for presented_client_id, overrides, expected_error in ( + ("llm_dcrc_bogus", {}, "invalid_client"), + (client_id, {"redirect_uri": "http://127.0.0.1:51235/callback"}, "invalid_request"), + (client_id, {"response_type": "token"}, "unsupported_response_type"), + (client_id, {"code_challenge": None}, "invalid_request"), + (client_id, {"code_challenge_method": "plain"}, "invalid_request"), + ): + response = await _native_authorize(presented_client_id, lookup=lookup, **overrides) + assert response.status_code == 400 + assert json.loads(response.body)["error"] == expected_error + assert "set-cookie" not in response.headers + assert lookup.calls == [] + + +@pytest.mark.asyncio +async def test_native_authorize_refuses_a_hosted_redirect_for_the_proxy_api(): + """Registration accepts any https redirect because MCP clients can be hosted, but a + proxy-API grant hands out the user's personal key, so it only ever goes back to loopback.""" + hosted = "https://evil.example/cb" + client_id = (await _register([hosted]))["client_id"] + lookup = _ConsentTeams() + response = await _native_authorize(client_id, redirect_uri=hosted, lookup=lookup) + assert response.status_code == 400 + assert json.loads(response.body) == { + "error": "invalid_request", + "error_description": "a proxy-API grant may only redirect to a loopback address", + } + assert "set-cookie" not in response.headers + assert lookup.calls == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "failure, status, error", + [ + ("unavailable", 503, "temporarily_unavailable"), + ("unresolvable", 500, "server_error"), + ("no_active_key", 403, "access_denied"), + ], +) +async def test_native_authorize_consent_lookup_failures_are_oauth_errors_without_a_flow(failure, status, error): + client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"] + response = await _native_authorize(client_id, lookup=_ConsentTeams(failure)) + assert response.status_code == status + assert json.loads(response.body)["error"] == error + assert "set-cookie" not in response.headers + + +@pytest.mark.asyncio +async def test_native_consent_escapes_untrusted_identifiers(): + client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"] + hostile = (ConsentTeam(team_id='t">