mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
6fcdea03b0
commit
2c691d3820
20 changed files with 3209 additions and 143 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
471
litellm/proxy/client/cli/commands/pkce_login.py
Normal file
471
litellm/proxy/client/cli/commands/pkce_login.py
Normal 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)
|
||||
|
|
@ -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>'
|
||||
)
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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 "<b>Team</b>" in body
|
||||
assert "<img src=x onerror="x">" 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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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
|
||||
|
|
|
|||
553
tests/test_litellm/proxy/client/cli/test_pkce_login.py
Normal file
553
tests/test_litellm/proxy/client/cli/test_pkce_login.py
Normal 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",
|
||||
}
|
||||
|
|
@ -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 "<b>x</b>" in page
|
||||
assert 'value="h" onmouseover="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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue