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=<proxy origin> 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
This commit is contained in:
mateo-berri 2026-08-20 04:10:10 -07:00
parent 6fcdea03b0
commit 2c691d3820
20 changed files with 3209 additions and 143 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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"""<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<meta name="referrer" content="no-referrer">
<title>Authorize CLI access - LiteLLM</title>
<style>
body {{
font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, Oxygen, Ubuntu, Cantarell, sans-serif;
background-color: #f8fafc;
margin: 0;
padding: 20px;
display: flex;
justify-content: center;
align-items: center;
min-height: 100vh;
color: #1e293b;
}}
.container {{
background-color: #fff;
padding: 40px;
border-radius: 8px;
box-shadow: 0 2px 8px rgba(0, 0, 0, 0.1);
width: 450px;
max-width: 100%;
}}
h1 {{ margin: 0 0 16px; font-size: 24px; font-weight: 600; }}
p {{ margin: 0 0 12px; line-height: 1.5; }}
code {{ background: #f1f5f9; padding: 2px 6px; border-radius: 4px; }}
label {{ display: block; margin: 16px 0 6px; font-weight: 600; }}
select {{ width: 100%; padding: 8px; border: 1px solid #cbd5e1; border-radius: 6px; font-size: 14px; }}
.actions {{ display: flex; gap: 12px; margin-top: 24px; }}
button {{ flex: 1; padding: 10px; border-radius: 6px; font-size: 15px; cursor: pointer; border: 1px solid #cbd5e1; }}
.approve {{ background: #2563eb; color: #fff; border-color: #2563eb; }}
.deny {{ background: #fff; color: #1e293b; }}
</style>
</head>
<body>
<div class="container">
<h1>Authorize CLI access</h1>
<p>A command-line client at <code>{escape(client_origin)}</code> wants to call LiteLLM as <strong>{escape(user_id)}</strong>.</p>
<p>Approving issues it a personal credential that expires within {CLI_JWT_EXPIRATION_HOURS} hours. <code>lite logout</code> stops it from being renewed. Only approve if you started this sign-in yourself.</p>
<form method="post" action="{escape(complete_url)}">
<input type="hidden" name="flow" value="{escape(flow_handle)}">
{_team_field(teams)}
<div class="actions">
<button type="submit" name="decision" value="deny" class="deny">Deny</button>
<button type="submit" name="decision" value="approve" class="approve">Approve</button>
</div>
</form>
</div>
</body>
</html>
"""
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'<input type="hidden" name="team_id" value="{escape(team_id)}">'
f"<p>Requests are attributed to team <strong>{escape(team_label)}</strong>.</p>"
)
options: Final = "".join(
f'<option value="{escape(team_id)}">{escape(team_label)}</option>' for team_id, team_label in teams
)
return (
f'<label for="team_id">Attribute requests to team</label><select id="team_id" name="team_id">{options}</select>'
)

View file

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

View file

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

View file

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

View file

@ -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 '<option value="team-b">team-b</option>' 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 == []

View file

@ -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 "<strong>u1</strong>" in body
assert '<option value="team-a">Team A</option>' in body
assert '<option value="team-b">team-b</option>' 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"><script>', team_alias="<b>Team</b>"), ConsentTeam(team_id="team-b"))
response = await _native_authorize(
client_id, session_user_id='<img src=x onerror="x">', lookup=_ConsentTeams(hostile)
)
body = response.body.decode()
assert "<script>" not in body
assert "<b>Team</b>" not in body
assert "<img" not in body
assert "&lt;b&gt;Team&lt;/b&gt;" in body
assert "&lt;img src=x onerror=&quot;x&quot;&gt;" in body
@pytest.mark.asyncio
async def test_native_deny_redirects_with_access_denied_and_burns_the_flow():
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
cache = DualCache()
consent = await _native_authorize(client_id)
denied = await _complete_consent(consent, cache=cache, decision="deny", team_id="team-a")
assert denied.status_code == 303
location = urlparse(denied.headers["location"])
assert f"{location.scheme}://{location.netloc}{location.path}" == LOOPBACK_REDIRECT_URI
assert parse_qs(location.query) == {"error": ["access_denied"], "state": ["client-state-123"]}
assert "max-age=0" in denied.headers["set-cookie"].lower()
retried = await _complete_consent(consent, cache=cache, decision="approve", team_id="team-a")
assert retried.status_code == 400
@pytest.mark.asyncio
async def test_native_invalid_decision_is_rejected_before_the_flow_is_consumed():
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
cache = DualCache()
consent = await _native_authorize(client_id)
bad = await _complete_consent(consent, cache=cache, decision="maybe", team_id="team-a")
assert bad.status_code == 400
assert json.loads(bad.body)["error"] == "invalid_request"
assert "set-cookie" not in bad.headers
approved = await _complete_consent(consent, cache=cache, decision="approve", team_id="team-a")
assert approved.status_code == 303
assert "code" in parse_qs(urlparse(approved.headers["location"]).query)
@pytest.mark.asyncio
async def test_native_complete_by_another_session_user_is_refused():
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
consent = await _native_authorize(client_id)
hijacked = await _complete_consent(consent, session_user_id="attacker", decision="approve", team_id="team-a")
assert hijacked.status_code == 403
@pytest.mark.asyncio
async def test_native_approve_seals_audience_and_chosen_team_into_the_code():
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
approved = await _complete_consent(await _native_authorize(client_id), decision="approve", team_id="team-b")
assert approved.status_code == 303
assert "max-age=0" in approved.headers["set-cookie"].lower()
location = urlparse(approved.headers["location"])
assert location.netloc == "127.0.0.1:51234"
assert location.path == "/callback"
query = parse_qs(location.query)
assert query["state"] == ["client-state-123"]
wire = _sealed_wire_json(query["code"][0], GATEWAY_AUTH_CODE_PREFIX, _AUTH_CODE_DEBUG_KEY)
assert wire["audience"] == "proxy_api"
assert wire["team_id"] == "team-b"
assert wire["user_id"] == "u1"
@pytest.mark.asyncio
@pytest.mark.parametrize("team_id", [None, ""])
async def test_native_approve_without_a_team_mints_an_unscoped_credential(team_id):
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
consent = await _native_authorize(client_id, lookup=_ConsentTeams(()))
assert 'name="team_id"' not in consent.body.decode()
approved = await _complete_consent(consent, decision="approve", team_id=team_id)
minter = _Minter()
response = await _redeem_native(_code_from(approved), client_id, minter)
assert response.status_code == 200
assert minter.calls == [("u1", None)]
payload = json.loads(response.body)
assert payload["team_id"] is None
assert _opened_refresh(payload["refresh_token"], client_id).team_id is None
@pytest.mark.asyncio
async def test_native_code_redeems_for_the_proxy_credential_and_a_rotating_refresh_token():
"""The whole native walk on one set of artifacts: code -> CLI credential + refresh, code
replay refused, refresh rotates and keeps the team, the old refresh is dead."""
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
cache = DualCache()
minter = _Minter()
code = await _native_code(client_id, team_id="team-b", cache=cache)
response = await _redeem_native(code, client_id, minter, cache=cache)
assert response.status_code == 200
assert response.headers["cache-control"] == "no-store"
payload = json.loads(response.body)
assert minter.calls == [("u1", "team-b")]
assert payload["access_token"] == "sk-cli-u1"
assert payload["token_type"] == "Bearer"
assert payload["expires_in"] == 3600
assert payload["user_id"] == "u1"
assert payload["team_id"] == "team-b"
assert payload["refresh_token"].startswith(SESSION_REFRESH_PREFIX)
principal = _opened_refresh(payload["refresh_token"], client_id)
assert principal.audience == "proxy_api"
assert principal.team_id == "team-b"
assert principal.user_id == "u1"
assert principal.client_id == client_id
replayed = await _redeem_native(code, client_id, minter, cache=cache)
assert replayed.status_code == 400
assert json.loads(replayed.body)["error"] == "invalid_grant"
assert "already used" in json.loads(replayed.body)["error_description"]
refreshed = await _refresh_native(payload["refresh_token"], client_id, minter, cache)
assert refreshed.status_code == 200
rotated = json.loads(refreshed.body)
assert minter.calls[-1] == ("u1", "team-b")
assert rotated["access_token"] == "sk-cli-u1"
assert rotated["team_id"] == "team-b"
assert rotated["refresh_token"] != payload["refresh_token"]
assert _opened_refresh(rotated["refresh_token"], client_id).team_id == "team-b"
stale = await _refresh_native(payload["refresh_token"], client_id, minter, cache)
assert stale.status_code == 400
assert "already used" in json.loads(stale.body)["error_description"]
assert "access_token" not in json.loads(stale.body)
@pytest.mark.asyncio
async def test_native_refresh_is_bound_to_the_issuing_client():
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
other = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
cache = DualCache()
payload = json.loads(
(await _redeem_native(await _native_code(client_id, cache=cache), client_id, _Minter(), cache=cache)).body
)
minter = _Minter()
stolen = await _refresh_native(payload["refresh_token"], other, minter, cache)
assert stolen.status_code == 400
assert json.loads(stolen.body)["error"] == "invalid_grant"
assert minter.calls == []
@pytest.mark.asyncio
async def test_native_code_without_a_minter_is_refused_server_side():
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
code = await _native_code(client_id)
response = await _redeem(code, client_id, redirect_uri=LOOPBACK_REDIRECT_URI, resource=PROXY_API_RESOURCE)
assert response.status_code == 500
assert json.loads(response.body)["error"] == "server_error"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"failure, status, error",
[
("not_a_member", 400, "invalid_grant"),
("no_active_key", 400, "invalid_grant"),
("unavailable", 503, "temporarily_unavailable"),
("unresolvable", 500, "server_error"),
],
)
async def test_native_mint_failure_maps_to_an_oauth_error_and_keeps_the_code_usable(failure, status, error):
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
cache = DualCache()
code = await _native_code(client_id, cache=cache)
failed = await _redeem_native(code, client_id, _Minter(failure), cache=cache)
assert failed.status_code == status
assert json.loads(failed.body)["error"] == error
assert "refresh_token" not in json.loads(failed.body)
retried = await _redeem_native(code, client_id, _Minter(), cache=cache)
assert retried.status_code == 200
@pytest.mark.asyncio
async def test_native_refresh_mint_failure_does_not_burn_the_refresh_token():
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
cache = DualCache()
payload = json.loads(
(await _redeem_native(await _native_code(client_id, cache=cache), client_id, _Minter(), cache=cache)).body
)
failed = await _refresh_native(payload["refresh_token"], client_id, _Minter("unavailable"), cache)
assert failed.status_code == 503
retried = await _refresh_native(payload["refresh_token"], client_id, _Minter(), cache)
assert retried.status_code == 200
@pytest.mark.asyncio
@pytest.mark.parametrize(
"resource", [PROXY_API_RESOURCE, "https://llm.example.com/", "HTTPS://LLM.EXAMPLE.COM:443", None]
)
async def test_native_code_accepts_the_proxy_resource_in_any_equivalent_spelling(resource):
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
response = await _redeem_native(await _native_code(client_id), client_id, _Minter(), resource=resource)
assert response.status_code == 200
@pytest.mark.asyncio
@pytest.mark.parametrize(
"resource", ["https://llm.example.com/mcp", "https://other.example.com", "http://llm.example.com", "not a url"]
)
async def test_native_code_refuses_a_foreign_resource_without_burning_it(resource):
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
cache = DualCache()
code = await _native_code(client_id, cache=cache)
minter = _Minter()
refused = await _redeem_native(code, client_id, minter, cache=cache, resource=resource)
assert refused.status_code == 400
assert json.loads(refused.body)["error"] == "invalid_target"
assert minter.calls == []
retried = await _redeem_native(code, client_id, minter, cache=cache)
assert retried.status_code == 200
@pytest.mark.asyncio
async def test_mcp_code_ignores_the_minter_and_still_yields_a_session_pair():
client_id = (await _register([REDIRECT_URI]))["client_id"]
code = await _finish_connect_page(_authorize(client_id, session_user_id="u1"))
minter = _Minter()
response = await _redeem(code, client_id, mint_proxy_credential=minter, resource=PROXY_API_RESOURCE)
assert response.status_code == 200
payload = json.loads(response.body)
assert minter.calls == []
principal = _opened_principal(payload)
assert principal.audience is None
assert principal.team_id is None
assert "user_id" not in payload
@pytest.mark.asyncio
async def test_mcp_wire_formats_carry_no_native_client_fields():
client_id = (await _register([REDIRECT_URI]))["client_id"]
handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1"))
flow_wire = _sealed_wire_json(next(iter(cookies.values())), "", "gateway_connect_flow")
assert "audience" not in flow_wire
assert "team_id" not in flow_wire
completed = await complete_connect_flow(
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="u1",
cache=DualCache(),
team_id="team-a",
decision="approve",
)
code_wire = _sealed_wire_json(_code_from(completed), GATEWAY_AUTH_CODE_PREFIX, _AUTH_CODE_DEBUG_KEY)
assert "audience" not in code_wire
assert "team_id" not in code_wire
@pytest.mark.asyncio
async def test_revoke_burns_the_refresh_token_and_answers_200_for_dead_or_unknown_tokens():
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
cache = DualCache()
payload = json.loads(
(await _redeem_native(await _native_code(client_id, cache=cache), client_id, _Minter(), cache=cache)).body
)
revoked = await revoke_refresh_token(
token=payload["refresh_token"], client_id=client_id, master_key=MASTER_KEY, cache=cache
)
assert revoked.status_code == 200
assert json.loads(revoked.body) == {}
assert revoked.headers["cache-control"] == "no-store"
refreshed = await _refresh_native(payload["refresh_token"], client_id, _Minter(), cache)
assert refreshed.status_code == 400
assert json.loads(refreshed.body)["error"] == "invalid_grant"
again = await revoke_refresh_token(
token=payload["refresh_token"], client_id=client_id, master_key=MASTER_KEY, cache=cache
)
assert again.status_code == 200
garbage = await revoke_refresh_token(token="nonsense", client_id=client_id, master_key=MASTER_KEY, cache=cache)
assert garbage.status_code == 200
@pytest.mark.asyncio
async def test_revoke_from_another_client_leaves_the_token_usable():
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
other = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
cache = DualCache()
payload = json.loads(
(await _redeem_native(await _native_code(client_id, cache=cache), client_id, _Minter(), cache=cache)).body
)
revoked = await revoke_refresh_token(
token=payload["refresh_token"], client_id=other, master_key=MASTER_KEY, cache=cache
)
assert revoked.status_code == 200
refreshed = await _refresh_native(payload["refresh_token"], client_id, _Minter(), cache)
assert refreshed.status_code == 200
@pytest.mark.asyncio
async def test_revoke_refuses_unknown_clients_and_a_missing_master_key():
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
for bogus in ("llm_dcrc_bogus", "anything"):
response = await revoke_refresh_token(token="x", client_id=bogus, master_key=MASTER_KEY, cache=DualCache())
assert response.status_code == 401
assert json.loads(response.body)["error"] == "invalid_client"
no_key = await revoke_refresh_token(token="x", client_id=client_id, master_key=None, cache=DualCache())
assert no_key.status_code == 500
assert json.loads(no_key.body)["error"] == "server_error"
def test_native_client_auth_contract_points_every_endpoint_at_this_proxy():
assert json.loads(json.dumps(native_client_auth_contract(_request("/.well-known/litellm-cli-auth")))) == {
"contract_version": 1,
"issuer": "https://llm.example.com",
"authorization_endpoint": "https://llm.example.com/authorize",
"token_endpoint": "https://llm.example.com/token",
"registration_endpoint": "https://llm.example.com/register",
"revocation_endpoint": "https://llm.example.com/revoke",
"resource": "https://llm.example.com",
"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"],
}
@pytest.mark.parametrize(
"resource, expected",
[
("https://llm.example.com", True),
("https://llm.example.com/", True),
("HTTPS://LLM.EXAMPLE.COM:443", True),
("https://llm.example.com/mcp", False),
("http://llm.example.com", False),
("https://other.example.com", False),
("llm.example.com", False),
("", False),
(None, False),
],
)
def test_is_proxy_api_resource_matches_only_this_proxy(resource, expected):
assert is_proxy_api_resource(_request(), resource) is expected

View file

@ -0,0 +1,167 @@
"""Tests for minting the ``lite login`` credential from a consented native-client grant."""
from unittest.mock import ANY, AsyncMock
import pytest
from litellm.constants import CLI_JWT_EXPIRATION_HOURS
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ConsentTeam, MintedProxyCredential
from litellm.proxy._experimental.mcp_server.proxy_api_credentials import lookup_consent_teams, mint_proxy_credential
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
from litellm.proxy.management_endpoints.ui_sso import CliSsoTeamDetail
_LOAD_USER = "litellm.proxy._experimental.mcp_server.proxy_api_credentials.load_active_user_by_id"
_FETCH_TEAMS = "litellm.proxy._experimental.mcp_server.proxy_api_credentials.fetch_cli_sso_team_details"
_PRISMA = "litellm.proxy.proxy_server.prisma_client"
TEAM_DETAILS = (
CliSsoTeamDetail(team_id="team-a", team_alias="Team A", team_models=("gpt-5.4-mini",)),
CliSsoTeamDetail(team_id="team-b", team_models=(), team_model_aliases={"fast": "gpt-5.4-mini"}),
)
def _user(**overrides) -> LiteLLM_UserTable:
return LiteLLM_UserTable(
**{
"user_id": "u1",
"user_role": "internal_user",
"teams": ["team-a", "team-b"],
"models": ["gpt-5.4"],
**overrides,
}
)
def _decoded(minted: MintedProxyCredential):
key_object = ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(minted.key)
assert key_object is not None
return key_object
@pytest.fixture(autouse=True)
def _salt_key(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-proxy-api-credentials-tests")
@pytest.fixture
def fetch_teams(monkeypatch):
monkeypatch.setattr(_PRISMA, object(), raising=False)
fetch = AsyncMock(return_value=TEAM_DETAILS)
monkeypatch.setattr(_FETCH_TEAMS, fetch)
return fetch
@pytest.fixture
def load_user(monkeypatch):
load = AsyncMock(return_value=_user())
monkeypatch.setattr(_LOAD_USER, load)
return load
@pytest.mark.asyncio
@pytest.mark.parametrize("failure", ["unavailable", "unresolvable", "no_active_key"])
async def test_mint_passes_user_lookup_failures_through(failure, load_user, fetch_teams):
load_user.return_value = failure
assert await mint_proxy_credential("u1", "team-a") == failure
fetch_teams.assert_not_awaited()
@pytest.mark.asyncio
async def test_mint_refuses_a_user_without_a_role(load_user, fetch_teams):
load_user.return_value = _user(user_role=None)
assert await mint_proxy_credential("u1", None) == "no_active_key"
fetch_teams.assert_not_awaited()
@pytest.mark.asyncio
async def test_mint_without_a_chosen_team_never_picks_one_for_a_team_member(load_user, fetch_teams):
"""The consent page is the only place a team gets chosen, so a grant sealed without one
stays unscoped on every redemption and refresh instead of drifting onto the first team."""
minted = await mint_proxy_credential("u1", None)
assert isinstance(minted, MintedProxyCredential)
assert minted.user_id == "u1"
assert minted.team_id is None
assert minted.expires_in == CLI_JWT_EXPIRATION_HOURS * 3600
load_user.assert_awaited_once_with("u1")
fetch_teams.assert_not_awaited()
decoded = _decoded(minted)
assert decoded.user_id == "u1"
assert decoded.team_id is None
assert decoded.team_alias is None
assert decoded.models == ["gpt-5.4"]
assert decoded.is_session_token is True
@pytest.mark.asyncio
async def test_mint_honors_the_consented_team(load_user, fetch_teams):
minted = await mint_proxy_credential("u1", "team-b")
assert isinstance(minted, MintedProxyCredential)
assert minted.team_id == "team-b"
decoded = _decoded(minted)
assert decoded.team_id == "team-b"
assert decoded.team_alias is None
assert decoded.team_models == []
assert decoded.team_model_aliases == {"fast": "gpt-5.4-mini"}
@pytest.mark.asyncio
async def test_mint_refuses_a_team_the_user_is_not_on(load_user, fetch_teams):
assert await mint_proxy_credential("u1", "team-c") == "not_a_member"
fetch_teams.assert_not_awaited()
@pytest.mark.asyncio
async def test_mint_refuses_when_the_teams_grants_are_unknown(load_user, fetch_teams):
fetch_teams.return_value = TEAM_DETAILS[:1]
assert await mint_proxy_credential("u1", "team-b") == "not_a_member"
@pytest.mark.asyncio
async def test_mint_reports_unavailable_when_the_team_lookup_fails(load_user, fetch_teams):
fetch_teams.return_value = None
assert await mint_proxy_credential("u1", "team-a") == "unavailable"
@pytest.mark.asyncio
async def test_mint_reports_unavailable_without_a_database(load_user, monkeypatch):
monkeypatch.setattr(_PRISMA, None, raising=False)
assert await mint_proxy_credential("u1", "team-a") == "unavailable"
@pytest.mark.asyncio
async def test_mint_for_a_teamless_user_is_unscoped(load_user, fetch_teams):
load_user.return_value = _user(teams=[])
minted = await mint_proxy_credential("u1", None)
assert isinstance(minted, MintedProxyCredential)
assert minted.team_id is None
fetch_teams.assert_not_awaited()
decoded = _decoded(minted)
assert decoded.team_id is None
assert decoded.models == ["gpt-5.4"]
@pytest.mark.asyncio
async def test_lookup_consent_teams_lists_the_users_teams_with_aliases(load_user, fetch_teams):
teams = await lookup_consent_teams("u1")
assert teams == (ConsentTeam(team_id="team-a", team_alias="Team A"), ConsentTeam(team_id="team-b"))
fetch_teams.assert_awaited_once_with(ANY, ["team-a", "team-b"])
@pytest.mark.asyncio
async def test_lookup_consent_teams_drops_details_without_a_team_id(load_user, fetch_teams):
fetch_teams.return_value = (CliSsoTeamDetail(team_models=()), *TEAM_DETAILS[1:])
assert await lookup_consent_teams("u1") == (ConsentTeam(team_id="team-b"),)
@pytest.mark.asyncio
@pytest.mark.parametrize("failure", ["unavailable", "unresolvable", "no_active_key"])
async def test_lookup_consent_teams_passes_user_lookup_failures_through(failure, load_user, fetch_teams):
load_user.return_value = failure
assert await lookup_consent_teams("u1") == failure
fetch_teams.assert_not_awaited()
@pytest.mark.asyncio
async def test_lookup_consent_teams_reports_unavailable_when_details_cannot_load(load_user, fetch_teams):
fetch_teams.return_value = None
assert await lookup_consent_teams("u1") == "unavailable"

View file

@ -190,19 +190,15 @@ class TestTokenUtilities:
result = get_token_file_path()
assert result == "/home/user/.litellm/token.json"
mock_mkdir.assert_called_once_with(exist_ok=True)
mock_mkdir.assert_not_called()
def test_get_token_file_path_creates_directory(self):
"""Test that get_token_file_path creates the config directory"""
with (
patch("pathlib.Path.home") as mock_home,
patch("pathlib.Path.mkdir") as mock_mkdir,
):
mock_home.return_value = Path("/home/user")
get_token_file_path()
mock_mkdir.assert_called_once_with(exist_ok=True)
def test_reading_the_token_never_creates_the_config_directory(self, tmp_path):
"""Every `lite` invocation reads the token; only saving one may touch ~/.litellm"""
with patch("pathlib.Path.home", return_value=tmp_path):
assert load_token() is None
assert not (tmp_path / ".litellm").exists()
save_token({"key": "sk-test"})
assert load_token() == {"key": "sk-test"}
def test_save_token(self, tmp_path):
"""Test saving token data to file"""
@ -309,7 +305,7 @@ class TestTokenUtilities:
token_data = {"key": "test-api-key-123", "user_id": "test-user"}
with patch(
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
"litellm.proxy.client.cli.commands.auth.load_token",
return_value=token_data,
):
result = get_stored_api_key()
@ -318,7 +314,7 @@ class TestTokenUtilities:
def test_get_stored_api_key_no_token(self):
"""Test getting stored API key when no token exists"""
with patch(
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
"litellm.proxy.client.cli.commands.auth.load_token",
return_value=None,
):
result = get_stored_api_key()
@ -329,7 +325,7 @@ class TestTokenUtilities:
token_data = {"user_id": "test-user"}
with patch(
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
"litellm.proxy.client.cli.commands.auth.load_token",
return_value=token_data,
):
result = get_stored_api_key()
@ -339,7 +335,7 @@ class TestTokenUtilities:
"""Stored key is returned when expected_base_url matches stored origin"""
token_data = {"key": "sk-prod", "base_url": "https://real-proxy.com"}
with patch(
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
"litellm.proxy.client.cli.commands.auth.load_token",
return_value=token_data,
):
assert get_stored_api_key(expected_base_url="https://real-proxy.com") == "sk-prod"
@ -348,7 +344,7 @@ class TestTokenUtilities:
"""Trailing slash on expected_base_url is normalised before comparison"""
token_data = {"key": "sk-prod", "base_url": "https://real-proxy.com"}
with patch(
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
"litellm.proxy.client.cli.commands.auth.load_token",
return_value=token_data,
):
assert get_stored_api_key(expected_base_url="https://real-proxy.com/") == "sk-prod"
@ -357,7 +353,7 @@ class TestTokenUtilities:
"""Stored key is NOT returned when expected_base_url differs from stored origin"""
token_data = {"key": "sk-prod", "base_url": "https://real-proxy.com"}
with patch(
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
"litellm.proxy.client.cli.commands.auth.load_token",
return_value=token_data,
):
assert get_stored_api_key(expected_base_url="https://evil.com") is None
@ -366,7 +362,7 @@ class TestTokenUtilities:
"""Old tokens without a base_url field are rejected when origin check is requested"""
token_data = {"key": "sk-old-token"}
with patch(
"litellm.litellm_core_utils.cli_token_utils.load_cli_token",
"litellm.proxy.client.cli.commands.auth.load_token",
return_value=token_data,
):
assert get_stored_api_key(expected_base_url="https://real-proxy.com") is None
@ -1110,3 +1106,292 @@ class TestLoginConfigClaude:
assert "could not configure Claude Code" in result.output
assert "invalid JSON" in result.output
assert "Authentication failed" not in result.output
class _FakeHttpResponse:
def __init__(self, status_code, payload):
self.status_code = status_code
self._payload = payload
self.text = json.dumps(payload)
self.content = self.text.encode()
def json(self):
return self._payload
class _FakeSession:
"""Stands in for ``requests.Session`` so the CLI's refresh and revoke calls can be observed."""
instances = []
def __init__(self):
self.posts = []
self.response = _FakeHttpResponse(200, {})
_FakeSession.instances.append(self)
def post(self, url, *, data=None, json=None, timeout):
self.posts.append((url, data))
return self.response
def get(self, url, *, timeout):
raise AssertionError(f"unexpected GET {url}")
PKCE_BASE_URL = "https://llm.example.com"
PKCE_TOKEN_RESPONSE = {
"access_token": "sk-cli-rotated",
"token_type": "Bearer",
"expires_in": 3600,
"refresh_token": "llm_srefresh_rotated",
"user_id": "u1",
"team_id": "team-b",
}
def _pkce_record(**overrides):
return {
"base_url": PKCE_BASE_URL,
"key": "sk-cli-old",
"user_id": "u1",
"user_email": "unknown",
"user_role": "cli",
"auth_header_name": "Authorization",
"jwt_token": "",
"timestamp": time.time(),
"expires_at": time.time() + 30,
"refresh_token": "llm_srefresh_old",
"client_id": "llm_dcrc_abc",
"token_endpoint": f"{PKCE_BASE_URL}/token",
"revocation_endpoint": f"{PKCE_BASE_URL}/revoke",
"resource": PKCE_BASE_URL,
"team_id": "team-b",
**overrides,
}
def _pkce_credential():
from litellm.proxy.client.cli.commands.pkce_login import PkceCredential
return PkceCredential(
access_token="sk-cli-fresh",
refresh_token="llm_srefresh_fresh",
expires_at=time.time() + 3600,
client_id="llm_dcrc_abc",
token_endpoint=f"{PKCE_BASE_URL}/token",
revocation_endpoint=f"{PKCE_BASE_URL}/revoke",
resource=PKCE_BASE_URL,
user_id="u1",
team_id="team-b",
)
class TestPkceLoginCommand:
"""``lite login --pkce`` swaps the proxy-mediated SSO poll for the browser PKCE flow."""
def setup_method(self):
self.runner = CliRunner()
_FakeSession.instances.clear()
def test_pkce_login_saves_the_refreshable_record_and_skips_the_sso_poll(self):
with (
patch("litellm.proxy.client.cli.commands.auth.run_pkce_login", return_value=_pkce_credential()) as run,
patch("litellm.proxy.client.cli.commands.auth._start_cli_sso_flow") as sso_start,
patch("litellm.proxy.client.cli.commands.auth.save_token") as save,
patch("litellm.proxy.client.cli.interface.show_commands"),
):
result = self.runner.invoke(login, ["--pkce"], obj={"base_url": f"{PKCE_BASE_URL}/"})
assert result.exit_code == 0, result.output
assert "Login successful!" in result.output
assert "JWT Token: sk-cli-fresh..." in result.output
sso_start.assert_not_called()
assert run.call_args.args[0] == f"{PKCE_BASE_URL}/"
saved = save.call_args.args[0]
assert saved["base_url"] == PKCE_BASE_URL
assert saved["key"] == "sk-cli-fresh"
assert saved["refresh_token"] == "llm_srefresh_fresh"
assert saved["client_id"] == "llm_dcrc_abc"
assert saved["token_endpoint"] == f"{PKCE_BASE_URL}/token"
assert saved["revocation_endpoint"] == f"{PKCE_BASE_URL}/revoke"
assert saved["resource"] == PKCE_BASE_URL
assert saved["user_id"] == "u1"
assert saved["team_id"] == "team-b"
def test_pkce_login_failure_is_reported_and_nothing_is_saved(self):
from litellm.proxy.client.cli.commands.pkce_login import PkceFailure
with (
patch(
"litellm.proxy.client.cli.commands.auth.run_pkce_login",
return_value=PkceFailure("sign-in was not approved (access_denied): no details"),
),
patch("litellm.proxy.client.cli.commands.auth.save_token") as save,
):
result = self.runner.invoke(login, ["--pkce"], obj={"base_url": PKCE_BASE_URL})
assert result.exit_code == 0
assert "Authentication failed: sign-in was not approved (access_denied): no details" in result.output
save.assert_not_called()
def test_login_without_the_flag_never_touches_the_pkce_flow(self):
with (
patch("litellm.proxy.client.cli.commands.auth.run_pkce_login") as run,
patch("litellm.proxy.client.cli.commands.auth._start_cli_sso_flow", side_effect=KeyboardInterrupt),
):
result = self.runner.invoke(login, obj={"base_url": PKCE_BASE_URL})
assert "cancelled" in result.output
run.assert_not_called()
class TestPkceLogoutCommand:
def setup_method(self):
self.runner = CliRunner()
_FakeSession.instances.clear()
def test_logout_revokes_the_refresh_token_before_clearing(self):
with (
patch("litellm.proxy.client.cli.commands.auth.load_token", return_value=_pkce_record()),
patch("litellm.proxy.client.cli.commands.auth.clear_token") as clear,
patch("litellm.proxy.client.cli.commands.auth.requests.Session", _FakeSession),
):
result = self.runner.invoke(logout)
assert result.exit_code == 0
assert result.output == "Logged out successfully. Authentication token cleared.\n"
clear.assert_called_once()
assert _FakeSession.instances[0].posts == [
(
f"{PKCE_BASE_URL}/revoke",
{"token": "llm_srefresh_old", "token_type_hint": "refresh_token", "client_id": "llm_dcrc_abc"},
)
]
def test_logout_still_clears_when_revocation_fails(self):
class _FailingSession(_FakeSession):
def __init__(self):
super().__init__()
self.response = _FakeHttpResponse(503, {"error": "temporarily_unavailable"})
with (
patch("litellm.proxy.client.cli.commands.auth.load_token", return_value=_pkce_record()),
patch("litellm.proxy.client.cli.commands.auth.clear_token") as clear,
patch("litellm.proxy.client.cli.commands.auth.requests.Session", _FailingSession),
):
result = self.runner.invoke(logout)
assert result.exit_code == 0
assert "Could not revoke the refresh token on the proxy (revocation failed with 503" in result.output
assert "Logged out successfully" in result.output
clear.assert_called_once()
def test_logout_of_a_classic_token_makes_no_request(self):
with (
patch("litellm.proxy.client.cli.commands.auth.load_token", return_value={"key": "sk-classic"}),
patch("litellm.proxy.client.cli.commands.auth.clear_token") as clear,
patch("litellm.proxy.client.cli.commands.auth.requests.Session", _FakeSession),
):
result = self.runner.invoke(logout)
assert result.output == "Logged out successfully. Authentication token cleared.\n"
clear.assert_called_once()
assert _FakeSession.instances[0].posts == []
class TestPkcePrintToken:
"""``lite print-token`` is Claude Code's apiKeyHelper, so a near-expiry PKCE key must be
refreshed silently and stdout must carry nothing but the key."""
def setup_method(self):
self.runner = CliRunner()
_FakeSession.instances.clear()
def test_print_token_refreshes_a_near_expiry_key_and_saves_the_rotation(self):
class _RefreshingSession(_FakeSession):
def __init__(self):
super().__init__()
self.response = _FakeHttpResponse(200, PKCE_TOKEN_RESPONSE)
with (
patch("litellm.proxy.client.cli.commands.auth.load_token", return_value=_pkce_record()),
patch("litellm.proxy.client.cli.commands.auth.save_token") as save,
patch("litellm.proxy.client.cli.commands.auth.requests.Session", _RefreshingSession),
):
result = self.runner.invoke(print_token, obj={})
assert result.exit_code == 0, result.output
assert result.stdout == "sk-cli-rotated\n"
assert _FakeSession.instances[0].posts[0][0] == f"{PKCE_BASE_URL}/token"
assert _FakeSession.instances[0].posts[0][1]["refresh_token"] == "llm_srefresh_old"
saved = save.call_args.args[0]
assert saved["key"] == "sk-cli-rotated"
assert saved["refresh_token"] == "llm_srefresh_rotated"
def test_print_token_prints_a_fresh_pkce_key_without_a_request(self):
with (
patch(
"litellm.proxy.client.cli.commands.auth.load_token",
return_value=_pkce_record(expires_at=time.time() + 3600),
),
patch("litellm.proxy.client.cli.commands.auth.requests.Session", _FakeSession),
):
result = self.runner.invoke(print_token, obj={})
assert result.stdout == "sk-cli-old\n"
assert _FakeSession.instances[0].posts == []
def test_print_token_fails_when_the_key_expired_and_refresh_is_refused(self):
class _RefusingSession(_FakeSession):
def __init__(self):
super().__init__()
self.response = _FakeHttpResponse(400, {"error": "invalid_grant"})
with (
patch(
"litellm.proxy.client.cli.commands.auth.load_token",
return_value=_pkce_record(expires_at=time.time() - 1),
),
patch("litellm.proxy.client.cli.commands.auth.save_token") as save,
patch("litellm.proxy.client.cli.commands.auth.requests.Session", _RefusingSession),
):
result = self.runner.invoke(print_token, obj={})
assert result.exit_code == 1
assert result.stdout == ""
assert "Token expired. Run 'lite login' again." in result.output
save.assert_not_called()
def test_print_token_for_an_expired_classic_token_makes_no_request(self):
with (
patch(
"litellm.proxy.client.cli.commands.auth.load_token",
return_value={"key": "sk-classic", "timestamp": time.time() - (CLI_JWT_EXPIRATION_HOURS + 1) * 3600},
),
patch("litellm.proxy.client.cli.commands.auth.requests.Session", _FakeSession),
):
result = self.runner.invoke(print_token, obj={})
assert result.exit_code == 1
assert "Token expired" in result.output
assert _FakeSession.instances == []
class TestGetStoredApiKeyRefresh:
def test_get_stored_api_key_refreshes_a_near_expiry_pkce_key(self):
_FakeSession.instances.clear()
class _RefreshingSession(_FakeSession):
def __init__(self):
super().__init__()
self.response = _FakeHttpResponse(200, PKCE_TOKEN_RESPONSE)
with (
patch("litellm.proxy.client.cli.commands.auth.load_token", return_value=_pkce_record()),
patch("litellm.proxy.client.cli.commands.auth.save_token") as save,
patch("litellm.proxy.client.cli.commands.auth.requests.Session", _RefreshingSession),
):
assert get_stored_api_key(PKCE_BASE_URL) == "sk-cli-rotated"
assert get_stored_api_key("https://other.example.com") is None
assert save.call_count == 1
assert len(_FakeSession.instances) == 1

View file

@ -0,0 +1,553 @@
"""Tests for the ``lite login --pkce`` browser flow against a fake proxy and a real loopback listener."""
import hashlib
import json
import re
import socket
import threading
import urllib.error
import urllib.request
from base64 import urlsafe_b64encode
from urllib.parse import parse_qs, urlparse
import pytest
import requests
from litellm.proxy.client.cli.commands.pkce_login import (
CallbackCode,
CallbackDenied,
CliAuthContract,
LoopbackServer,
PkceCredential,
PkceFailure,
_error_detail,
authorize_url,
discover_cli_auth,
fresh_api_key,
pkce_pair,
pkce_token_record,
redeem_code,
refresh_credential,
register_client,
revoke_credential,
revoke_stored_credential,
run_pkce_login,
)
BASE = "https://llm.example.com"
DISCOVERY_URL = f"{BASE}/.well-known/litellm-cli-auth"
CONTRACT_DOC = {
"contract_version": 1,
"issuer": BASE,
"authorization_endpoint": f"{BASE}/authorize",
"token_endpoint": f"{BASE}/token",
"registration_endpoint": f"{BASE}/register",
"revocation_endpoint": f"{BASE}/revoke",
"resource": BASE,
"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"],
}
CONTRACT = CliAuthContract.model_validate(CONTRACT_DOC)
TOKEN_JSON = {
"access_token": "sk-cli-new",
"token_type": "Bearer",
"expires_in": 3600,
"refresh_token": "llm_srefresh_new",
"user_id": "u1",
"team_id": "team-b",
}
STORED = {
"base_url": BASE,
"key": "sk-cli-old",
"expires_at": 1_000_000.0,
"refresh_token": "llm_srefresh_old",
"client_id": "llm_dcrc_abc",
"token_endpoint": f"{BASE}/token",
"revocation_endpoint": f"{BASE}/revoke",
"resource": BASE,
}
class _FakeResponse:
def __init__(self, status_code, payload=None, text=""):
self.status_code = status_code
self._payload = payload
self.text = json.dumps(payload) if payload is not None else text
self.content = self.text.encode()
def json(self):
if self._payload is None:
raise ValueError("not json")
return self._payload
def _wire(body):
return None if body is None else json.loads(json.dumps(body))
class _FakeHttp:
def __init__(self, routes=None):
self.routes = routes or {}
self.calls = []
def get(self, url, *, timeout):
self.calls.append(("GET", url, None, None))
return self._route("GET", url, None, None)
def post(self, url, *, data=None, json=None, timeout):
self.calls.append(("POST", url, data, _wire(json)))
return self._route("POST", url, data, json)
def _route(self, method, url, data, json):
handler = self.routes[(method, url)]
if isinstance(handler, Exception):
raise handler
return handler(data, json) if callable(handler) else handler
def _get(url):
try:
with urllib.request.urlopen(url, timeout=5) as response:
return response.status, response.read().decode()
except urllib.error.HTTPError as error:
return error.code, error.read().decode()
def _fresh(token_data, save, http, reload=lambda: None, **options):
return fresh_api_key(token_data, save, http, reload=reload, **options)
def _challenge_for(verifier: str) -> str:
return urlsafe_b64encode(hashlib.sha256(verifier.encode("ascii")).digest()).rstrip(b"=").decode("ascii")
def test_loopback_settles_only_on_a_code_for_the_expected_state():
with LoopbackServer("expected-state") as server:
result = {}
thread = threading.Thread(target=lambda: result.setdefault("outcome", server.wait(10)))
thread.start()
base = f"http://127.0.0.1:{server.server_address[1]}"
assert server.redirect_uri == f"{base}/callback"
assert _get(f"{base}/nope")[0] == 404
status, body = _get(f"{base}/callback?state=other&code=stolen")
assert status == 400
assert "still waiting" in body
status, body = _get(f"{base}/callback?state=expected-state")
assert status == 400
assert "no authorization code" in body
assert thread.is_alive()
status, body = _get(f"{base}/callback?state=expected-state&code=the-code")
assert status == 200
assert "Signed in to LiteLLM" in body
thread.join(5)
assert not thread.is_alive()
assert result["outcome"] == CallbackCode(code="the-code")
def test_loopback_reports_a_denied_sign_in():
with LoopbackServer("expected-state") as server:
result = {}
thread = threading.Thread(target=lambda: result.setdefault("outcome", server.wait(10)))
thread.start()
base = f"http://127.0.0.1:{server.server_address[1]}"
status, body = _get(f"{base}/callback?state=expected-state&error=access_denied&error_description=nope")
assert status == 200
assert "not approved" in body
thread.join(5)
assert result["outcome"] == CallbackDenied(error="access_denied", description="nope")
def test_loopback_drops_an_idle_connection_instead_of_waiting_on_it():
with LoopbackServer("expected-state", connection_timeout_seconds=0.2) as server:
result = {}
thread = threading.Thread(target=lambda: result.setdefault("outcome", server.wait(10)))
thread.start()
idle = socket.create_connection(server.server_address)
try:
status, body = _get(f"http://127.0.0.1:{server.server_address[1]}/callback?state=expected-state&code=c1")
finally:
idle.close()
assert status == 200
assert "Signed in to LiteLLM" in body
thread.join(5)
assert not thread.is_alive()
assert result["outcome"] == CallbackCode(code="c1")
def test_loopback_wait_times_out_on_its_clock():
ticks = iter([0.0, 0.5, 1.5])
with LoopbackServer("expected-state") as server:
server.timeout = 0.01
outcome = server.wait(1.0, clock=lambda: next(ticks))
assert outcome == PkceFailure("timed out waiting for the browser sign-in to finish")
def test_discover_reads_the_contract_from_the_well_known_path():
http = _FakeHttp({("GET", DISCOVERY_URL): _FakeResponse(200, CONTRACT_DOC)})
assert discover_cli_auth(f"{BASE}/", http) == CONTRACT
@pytest.mark.parametrize(
"response, expected",
[
(_FakeResponse(404, text="not found"), "does not support `lite login --pkce`"),
(_FakeResponse(200, {**CONTRACT_DOC, "contract_version": 2}), "unsupported discovery document"),
(_FakeResponse(200, {"issuer": BASE}), "unsupported discovery document"),
(_FakeResponse(200, text="<html>"), "unsupported discovery document"),
(_FakeResponse(200, {**CONTRACT_DOC, "code_challenge_methods_supported": ["plain"]}), "PKCE S256"),
(requests.ConnectionError("refused"), "could not reach"),
],
)
def test_discover_failures_name_the_cause(response, expected):
result = discover_cli_auth(BASE, _FakeHttp({("GET", DISCOVERY_URL): response}))
assert isinstance(result, PkceFailure)
assert expected in result.reason
def test_register_sends_a_public_loopback_client_and_returns_its_id():
http = _FakeHttp({("POST", f"{BASE}/register"): _FakeResponse(201, {"client_id": "llm_dcrc_abc"})})
assert register_client(CONTRACT, "http://127.0.0.1:5/callback", http) == "llm_dcrc_abc"
assert http.calls == [
(
"POST",
f"{BASE}/register",
None,
{
"client_name": "litellm-cli",
"redirect_uris": ["http://127.0.0.1:5/callback"],
"grant_types": ["authorization_code", "refresh_token"],
"response_types": ["code"],
"token_endpoint_auth_method": "none",
},
)
]
@pytest.mark.parametrize(
"response, expected",
[
(
_FakeResponse(400, {"error": "invalid_redirect_uri", "error_description": "loopback only"}),
"400: loopback only",
),
(_FakeResponse(201, {}), "unexpected body"),
(requests.ConnectionError("refused"), "registration failed"),
],
)
def test_register_failures_name_the_cause(response, expected):
result = register_client(
CONTRACT, "http://127.0.0.1:5/callback", _FakeHttp({("POST", f"{BASE}/register"): response})
)
assert isinstance(result, PkceFailure)
assert expected in result.reason
def test_pkce_pair_is_a_high_entropy_s256_pair():
verifier, challenge = pkce_pair()
assert 43 <= len(verifier) <= 128
assert challenge == _challenge_for(verifier)
assert pkce_pair()[0] != verifier
def test_authorize_url_carries_every_oauth_parameter_and_the_proxy_resource():
url = authorize_url(CONTRACT, "llm_dcrc_abc", "http://127.0.0.1:5/callback", "state-1", "challenge-1")
parsed = urlparse(url)
assert f"{parsed.scheme}://{parsed.netloc}{parsed.path}" == f"{BASE}/authorize"
assert parse_qs(parsed.query) == {
"response_type": ["code"],
"client_id": ["llm_dcrc_abc"],
"redirect_uri": ["http://127.0.0.1:5/callback"],
"state": ["state-1"],
"code_challenge": ["challenge-1"],
"code_challenge_method": ["S256"],
"resource": [BASE],
}
def test_redeem_code_posts_the_verifier_and_binds_the_credential_to_the_contract():
http = _FakeHttp({("POST", f"{BASE}/token"): _FakeResponse(200, TOKEN_JSON)})
credential = redeem_code(
CONTRACT, "llm_dcrc_abc", "http://127.0.0.1:5/callback", "the-code", "the-verifier", http, now=lambda: 100.0
)
assert credential == PkceCredential(
access_token="sk-cli-new",
refresh_token="llm_srefresh_new",
expires_at=3700.0,
client_id="llm_dcrc_abc",
token_endpoint=f"{BASE}/token",
revocation_endpoint=f"{BASE}/revoke",
resource=BASE,
user_id="u1",
team_id="team-b",
)
assert http.calls[0][2] == {
"grant_type": "authorization_code",
"code": "the-code",
"redirect_uri": "http://127.0.0.1:5/callback",
"client_id": "llm_dcrc_abc",
"code_verifier": "the-verifier",
"resource": BASE,
}
def test_refresh_posts_the_refresh_grant_for_the_same_client_and_resource():
http = _FakeHttp({("POST", f"{BASE}/token"): _FakeResponse(200, TOKEN_JSON)})
credential = refresh_credential(
f"{BASE}/token", f"{BASE}/revoke", BASE, "llm_dcrc_abc", "llm_srefresh_old", http, now=lambda: 5.0
)
assert isinstance(credential, PkceCredential)
assert credential.expires_at == 3605.0
assert credential.refresh_token == "llm_srefresh_new"
assert http.calls[0][2] == {
"grant_type": "refresh_token",
"refresh_token": "llm_srefresh_old",
"client_id": "llm_dcrc_abc",
"resource": BASE,
}
@pytest.mark.parametrize(
"response, expected",
[
(_FakeResponse(400, {"error": "invalid_grant", "error_description": "already used"}), "400: already used"),
(_FakeResponse(200, {**TOKEN_JSON, "refresh_token": ""}), "unexpected body"),
(_FakeResponse(200, {**TOKEN_JSON, "expires_in": 0}), "unexpected body"),
(_FakeResponse(200, text="<html>"), "unexpected body"),
(requests.ConnectionError("refused"), "token request failed"),
],
)
def test_token_request_failures_name_the_cause(response, expected):
result = redeem_code(CONTRACT, "c", "r", "code", "v", _FakeHttp({("POST", f"{BASE}/token"): response}))
assert isinstance(result, PkceFailure)
assert expected in result.reason
def test_revoke_posts_an_rfc7009_request_for_the_refresh_token():
http = _FakeHttp({("POST", f"{BASE}/revoke"): _FakeResponse(200, {})})
assert revoke_credential(f"{BASE}/revoke", "llm_dcrc_abc", "llm_srefresh_old", http) is None
assert http.calls == [
(
"POST",
f"{BASE}/revoke",
{"token": "llm_srefresh_old", "token_type_hint": "refresh_token", "client_id": "llm_dcrc_abc"},
None,
)
]
@pytest.mark.parametrize(
"response, expected",
[
(_FakeResponse(401, {"error": "invalid_client"}), "401: invalid_client"),
(requests.ConnectionError("refused"), "revocation request failed"),
],
)
def test_revoke_failures_name_the_cause(response, expected):
result = revoke_credential(f"{BASE}/revoke", "llm_dcrc_abc", "t", _FakeHttp({("POST", f"{BASE}/revoke"): response}))
assert isinstance(result, PkceFailure)
assert expected in result.reason
@pytest.mark.parametrize(
"response, expected",
[
(_FakeResponse(400, {"error": "invalid_grant", "error_description": "the code expired"}), "the code expired"),
(_FakeResponse(400, {"error": "invalid_grant"}), "invalid_grant"),
(_FakeResponse(404, {"detail": "Not Found"}), "Not Found"),
(_FakeResponse(502, text="<html>bad gateway</html>"), "<html>bad gateway</html>"),
(_FakeResponse(500, text="x" * 500), "x" * 200),
],
)
def test_error_detail_prefers_the_oauth_description(response, expected):
assert _error_detail(response) == expected
def _blocking_browser(seen, suffix):
seen["done"] = threading.Event()
def open_browser(url):
query = parse_qs(urlparse(url).query)
seen["authorize"] = query
seen["callback"] = _get(f"{query['redirect_uri'][0]}?state={query['state'][0]}&{suffix}")
seen["done"].set()
return open_browser
def test_run_pkce_login_end_to_end_against_a_fake_proxy():
seen = {}
def register(data, json):
seen["register"] = json
return _FakeResponse(201, {"client_id": "llm_dcrc_abc"})
def token(data, json):
seen["token"] = data
return _FakeResponse(200, TOKEN_JSON)
http = _FakeHttp(
{
("GET", DISCOVERY_URL): _FakeResponse(200, CONTRACT_DOC),
("POST", f"{BASE}/register"): register,
("POST", f"{BASE}/token"): token,
}
)
echoed = []
credential = run_pkce_login(
BASE, http, open_browser=_blocking_browser(seen, "code=the-code"), echo=echoed.append, timeout_seconds=10
)
assert isinstance(credential, PkceCredential)
assert credential.access_token == "sk-cli-new"
assert credential.refresh_token == "llm_srefresh_new"
assert credential.client_id == "llm_dcrc_abc"
assert credential.team_id == "team-b"
redirect_uri = seen["register"]["redirect_uris"][0]
assert re.fullmatch(r"http://127\.0\.0\.1:\d+/callback", redirect_uri)
assert seen["authorize"]["redirect_uri"] == [redirect_uri]
assert seen["authorize"]["client_id"] == ["llm_dcrc_abc"]
assert seen["authorize"]["code_challenge_method"] == ["S256"]
assert seen["authorize"]["resource"] == [BASE]
form = seen["token"]
assert form["code"] == "the-code"
assert form["redirect_uri"] == redirect_uri
assert form["client_id"] == "llm_dcrc_abc"
assert _challenge_for(form["code_verifier"]) == seen["authorize"]["code_challenge"][0]
assert echoed[0].startswith(f"Opening browser to: {BASE}/authorize?")
assert echoed[1] == "Approve the sign-in in your browser. Waiting..."
assert seen["done"].wait(5)
assert seen["callback"][0] == 200
def test_run_pkce_login_reports_a_denied_sign_in_without_touching_the_token_endpoint():
seen = {}
http = _FakeHttp(
{
("GET", DISCOVERY_URL): _FakeResponse(200, CONTRACT_DOC),
("POST", f"{BASE}/register"): _FakeResponse(201, {"client_id": "llm_dcrc_abc"}),
}
)
result = run_pkce_login(
BASE,
http,
open_browser=_blocking_browser(seen, "error=access_denied&error_description=nope"),
echo=lambda _: None,
timeout_seconds=10,
)
assert result == PkceFailure("sign-in was not approved (access_denied): nope")
assert [call[1] for call in http.calls] == [DISCOVERY_URL, f"{BASE}/register"]
def test_run_pkce_login_stops_before_the_browser_when_registration_fails():
opened = []
http = _FakeHttp(
{
("GET", DISCOVERY_URL): _FakeResponse(200, CONTRACT_DOC),
("POST", f"{BASE}/register"): _FakeResponse(400, {"error": "invalid_client_metadata"}),
}
)
result = run_pkce_login(BASE, http, open_browser=opened.append, echo=lambda _: None, timeout_seconds=1)
assert isinstance(result, PkceFailure)
assert "invalid_client_metadata" in result.reason
assert opened == []
def test_pkce_token_record_keeps_every_refresh_input_next_to_the_key():
credential = PkceCredential(
access_token="sk-cli-new",
refresh_token="llm_srefresh_new",
expires_at=3700.0,
client_id="llm_dcrc_abc",
token_endpoint=f"{BASE}/token",
revocation_endpoint=f"{BASE}/revoke",
resource=BASE,
user_id=None,
team_id=None,
)
record = pkce_token_record(f"{BASE}/", credential)
assert record["base_url"] == BASE
assert record["key"] == "sk-cli-new"
assert record["user_id"] == "cli-user"
assert record["expires_at"] == 3700.0
assert record["refresh_token"] == "llm_srefresh_new"
assert record["client_id"] == "llm_dcrc_abc"
assert record["token_endpoint"] == f"{BASE}/token"
assert record["revocation_endpoint"] == f"{BASE}/revoke"
assert record["resource"] == BASE
assert record["team_id"] is None
assert record["auth_header_name"] == "Authorization"
def test_fresh_api_key_returns_a_classic_or_still_fresh_key_without_network():
http = _FakeHttp()
saved = []
assert _fresh({"key": "sk-classic"}, saved.append, http) == "sk-classic"
assert _fresh(STORED, saved.append, http, now=lambda: 999_000.0) == "sk-cli-old"
assert _fresh({}, saved.append, http) is None
assert _fresh({"key": ""}, saved.append, http) is None
assert http.calls == []
assert saved == []
def test_fresh_api_key_refreshes_near_expiry_and_saves_before_returning():
http = _FakeHttp({("POST", f"{BASE}/token"): _FakeResponse(200, TOKEN_JSON)})
saved = []
assert _fresh(STORED, saved.append, http, now=lambda: 999_950.0) == "sk-cli-new"
assert len(saved) == 1
assert saved[0]["key"] == "sk-cli-new"
assert saved[0]["refresh_token"] == "llm_srefresh_new"
assert saved[0]["expires_at"] == 999_950.0 + 3600
assert saved[0]["base_url"] == BASE
assert saved[0]["team_id"] == "team-b"
assert http.calls[0][2]["refresh_token"] == "llm_srefresh_old"
def test_fresh_api_key_never_hands_out_a_rotated_key_it_could_not_save():
http = _FakeHttp({("POST", f"{BASE}/token"): _FakeResponse(200, TOKEN_JSON)})
def save(_record):
raise OSError("disk full")
with pytest.raises(OSError):
_fresh(STORED, save, http, now=lambda: 999_950.0)
def test_fresh_api_key_falls_back_to_the_old_key_only_while_it_is_still_valid():
failing = _FakeHttp({("POST", f"{BASE}/token"): _FakeResponse(503, {"error": "temporarily_unavailable"})})
saved = []
assert _fresh(STORED, saved.append, failing, now=lambda: 999_950.0) == "sk-cli-old"
assert _fresh(STORED, saved.append, failing, now=lambda: 1_000_001.0) is None
assert saved == []
def test_fresh_api_key_uses_a_sibling_rotation_when_its_own_refresh_loses_the_race():
rotated = {**STORED, "key": "sk-cli-sibling", "refresh_token": "llm_srefresh_sibling", "expires_at": 1_003_600.0}
http = _FakeHttp({("POST", f"{BASE}/token"): _FakeResponse(400, {"error": "invalid_grant"})})
saved = []
assert _fresh(STORED, saved.append, http, reload=lambda: rotated, now=lambda: 999_950.0) == "sk-cli-sibling"
assert _fresh(STORED, saved.append, http, reload=lambda: rotated, now=lambda: 1_000_001.0) == "sk-cli-sibling"
assert _fresh(STORED, saved.append, http, reload=lambda: STORED, now=lambda: 1_000_001.0) is None
assert _fresh(STORED, saved.append, http, reload=lambda: None, now=lambda: 1_000_001.0) is None
assert saved == []
def test_fresh_api_key_without_refresh_inputs_expires_like_a_classic_key():
http = _FakeHttp()
no_refresh = {key: value for key, value in STORED.items() if key != "refresh_token"}
assert _fresh(no_refresh, lambda _: None, http, now=lambda: 999_950.0) == "sk-cli-old"
assert _fresh(no_refresh, lambda _: None, http, now=lambda: 1_000_001.0) is None
assert http.calls == []
def test_revoke_stored_credential_revokes_only_pkce_records():
http = _FakeHttp({("POST", f"{BASE}/revoke"): _FakeResponse(200, {})})
assert revoke_stored_credential({"key": "sk-classic"}, http) is None
assert http.calls == []
assert revoke_stored_credential(STORED, http) is None
assert http.calls[0][2] == {
"token": "llm_srefresh_old",
"token_type_hint": "refresh_token",
"client_id": "llm_dcrc_abc",
}

View file

@ -0,0 +1,70 @@
import os
import sys
sys.path.insert(0, os.path.abspath("../../../"))
from litellm.constants import CLI_JWT_EXPIRATION_HOURS
from litellm.proxy.common_utils.html_forms.native_client_consent import render_native_client_consent_page
def _render(teams=(), **overrides) -> str:
arguments = {
"client_origin": "http://127.0.0.1:51234",
"user_id": "u1",
"teams": teams,
"flow_handle": "handle-123",
"complete_url": "https://llm.example.com/authorize/complete",
}
return render_native_client_consent_page(**{**arguments, **overrides})
def test_consent_page_posts_the_flow_handle_and_both_decisions_to_the_complete_url():
page = _render()
assert '<meta name="referrer" content="no-referrer">' in page
assert '<form method="post" action="https://llm.example.com/authorize/complete">' in page
assert '<input type="hidden" name="flow" value="handle-123">' in page
assert '<button type="submit" name="decision" value="deny"' in page
assert '<button type="submit" name="decision" value="approve"' in page
assert "<code>http://127.0.0.1:51234</code>" in page
assert "<strong>u1</strong>" in page
assert 'name="team_id"' not in page
def test_consent_page_pins_a_single_team_without_a_chooser():
page = _render(teams=(("team-a", "Team A"),))
assert '<input type="hidden" name="team_id" value="team-a">' in page
assert "<strong>Team A</strong>" in page
assert "<select" not in page
def test_consent_page_offers_a_chooser_for_several_teams():
page = _render(teams=(("team-a", "Team A"), ("team-b", "team-b")))
assert '<select id="team_id" name="team_id">' in page
assert '<option value="team-a">Team A</option>' in page
assert '<option value="team-b">team-b</option>' in page
assert 'type="hidden" name="team_id"' not in page
def test_consent_page_escapes_every_untrusted_value():
page = _render(
teams=(('t"><script>', "<b>x</b>"),),
client_origin="http://127.0.0.1:1/<svg>",
user_id='<img src=x onerror="y">',
flow_handle='h" onmouseover="z',
complete_url="https://llm.example.com/authorize/complete?x=<y>",
)
assert "<script>" not in page
assert "<b>x</b>" not in page
assert "<svg>" not in page
assert "<img" not in page
assert 'onmouseover="z' not in page
assert "?x=<y>" not in page
assert "&lt;b&gt;x&lt;/b&gt;" in page
assert 'value="h&quot; onmouseover=&quot;z"' in page
def test_consent_page_promises_only_what_logout_can_deliver():
page = _render()
assert f"expires within {CLI_JWT_EXPIRATION_HOURS} hours" in page
assert "<code>lite logout</code> stops it from being renewed" in page
assert "revoked" not in page

View file

@ -3266,7 +3266,7 @@ class TestCLIKeyRegenerationFlow:
too, otherwise alias lookup at request time is a substring match on a string.
"""
from litellm.proxy.management_endpoints.ui_sso import (
_fetch_cli_sso_team_details,
fetch_cli_sso_team_details,
)
team_row = MagicMock()
@ -3285,7 +3285,7 @@ class TestCLIKeyRegenerationFlow:
prisma_client = MagicMock()
prisma_client.db.litellm_teamtable.find_many = find_many
details = await _fetch_cli_sso_team_details(
details = await fetch_cli_sso_team_details(
prisma_client=prisma_client, teams=["team-a"]
)
@ -3308,7 +3308,7 @@ class TestCLIKeyRegenerationFlow:
real empty answer means the team rows are genuinely gone.
"""
from litellm.proxy.management_endpoints.ui_sso import (
_fetch_cli_sso_team_details,
fetch_cli_sso_team_details,
)
failing_client = MagicMock()
@ -3316,7 +3316,7 @@ class TestCLIKeyRegenerationFlow:
side_effect=Exception("connection reset")
)
assert (
await _fetch_cli_sso_team_details(
await fetch_cli_sso_team_details(
prisma_client=failing_client, teams=["team-a"]
)
is None
@ -3325,7 +3325,7 @@ class TestCLIKeyRegenerationFlow:
empty_client = MagicMock()
empty_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
assert (
await _fetch_cli_sso_team_details(
await fetch_cli_sso_team_details(
prisma_client=empty_client, teams=["team-a"]
)
== ()
@ -8124,7 +8124,7 @@ async def test_cli_completion_persists_assertion_under_db_user_id():
AsyncMock(return_value=user_info),
),
patch(
"litellm.proxy.management_endpoints.ui_sso._fetch_cli_sso_team_details",
"litellm.proxy.management_endpoints.ui_sso.fetch_cli_sso_team_details",
AsyncMock(return_value=[]),
),
patch(
@ -8202,11 +8202,11 @@ async def test_cli_completion_drops_teams_whose_rows_no_longer_exist():
would be refused with no way for the user to recover.
"""
from litellm.proxy.management_endpoints.ui_sso import (
_CliSsoTeamDetail,
CliSsoTeamDetail,
_complete_cli_sso_callback_session,
)
live_detail = _CliSsoTeamDetail(
live_detail = CliSsoTeamDetail(
team_id="team-live", team_alias="Live", team_models=("gpt-4.1",)
)
flow = {}
@ -8216,7 +8216,7 @@ async def test_cli_completion_drops_teams_whose_rows_no_longer_exist():
AsyncMock(return_value=_cli_callback_user_info(["team-live", "team-deleted"])),
),
patch(
"litellm.proxy.management_endpoints.ui_sso._fetch_cli_sso_team_details",
"litellm.proxy.management_endpoints.ui_sso.fetch_cli_sso_team_details",
AsyncMock(return_value=(live_detail,)),
),
patch(
@ -8254,7 +8254,7 @@ async def test_cli_completion_fails_the_login_when_team_lookup_fails():
AsyncMock(return_value=_cli_callback_user_info(["team-live"])),
),
patch(
"litellm.proxy.management_endpoints.ui_sso._fetch_cli_sso_team_details",
"litellm.proxy.management_endpoints.ui_sso.fetch_cli_sso_team_details",
AsyncMock(return_value=None),
),
patch(