Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/client/cli/commands/pkce_login.py: 0%
310 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""Browser sign-in for ``lite login --pkce``: OAuth 2.1 authorization code + PKCE S256
2against the proxy's own authorization server, as a public client on a loopback redirect.
3The proxy publishes everything this needs at ``/.well-known/litellm-cli-auth``, so a CLI
4in any other language can run the same steps from that document alone."""
6from __future__ import annotations
8import hashlib
9import secrets
10import socket
11import threading
12import time
13import webbrowser
14from base64 import urlsafe_b64encode
15from collections.abc import Callable, Mapping, Sequence
16from dataclasses import dataclass
17from http.server import BaseHTTPRequestHandler, HTTPServer
18from types import MappingProxyType
19from typing import TYPE_CHECKING, Final, Literal, Protocol
20from urllib.parse import parse_qs, urlencode, urlparse
22import requests
23from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
24from typing_extensions import ReadOnly, TypedDict
26from litellm.litellm_core_utils.cli_token_utils import CLI_TOKEN_FRESHNESS_BUFFER_SECONDS
28if TYPE_CHECKING:
29 from .auth import CliTokenData
31CLI_AUTH_DISCOVERY_PATH: Final = "/.well-known/litellm-cli-auth"
32CALLBACK_PATH: Final = "/callback"
33LOGIN_TIMEOUT_SECONDS: Final = 300
34_HTTP_TIMEOUT_SECONDS: Final = 15
35_CLIENT_NAME: Final = "litellm-cli"
38class CliAuthContract(BaseModel):
39 model_config = ConfigDict(frozen=True)
41 contract_version: Literal[1]
42 issuer: str
43 authorization_endpoint: str
44 token_endpoint: str
45 registration_endpoint: str
46 revocation_endpoint: str
47 resource: str
48 code_challenge_methods_supported: tuple[str, ...]
51class _RegisteredClient(BaseModel):
52 model_config = ConfigDict(frozen=True)
54 client_id: str = Field(min_length=1)
57class _TokenResponse(BaseModel):
58 model_config = ConfigDict(frozen=True)
60 access_token: str = Field(min_length=1)
61 expires_in: int = Field(gt=0)
62 refresh_token: str = Field(min_length=1)
63 user_id: str | None = None
64 team_id: str | None = None
67@dataclass(frozen=True, slots=True)
68class PkceFailure:
69 reason: str
72@dataclass(frozen=True, slots=True)
73class RevocationUnavailable:
74 reason: str
77@dataclass(frozen=True, slots=True)
78class PkceCredential:
79 access_token: str
80 refresh_token: str
81 expires_at: float
82 client_id: str
83 token_endpoint: str
84 revocation_endpoint: str
85 resource: str
86 user_id: str | None
87 team_id: str | None
90@dataclass(frozen=True, slots=True)
91class CallbackCode:
92 code: str
95@dataclass(frozen=True, slots=True)
96class CallbackDenied:
97 error: str
98 description: str | None
101CallbackOutcome = CallbackCode | CallbackDenied
104class Http(Protocol):
105 def get(self, url: str, *, timeout: float) -> requests.Response: ...
107 def post(
108 self,
109 url: str,
110 *,
111 data: Mapping[str, str] | None = None,
112 json: Mapping[str, object] | None = None,
113 timeout: float,
114 allow_redirects: bool,
115 ) -> requests.Response: ...
118class LoopbackServer(HTTPServer):
119 """The OS-assigned loopback listener the browser is sent back to. Only the response
120 carrying the pending sign-in's ``state`` settles it; anything else (a stray request, a
121 stale tab, an attacker poking the port) gets a 400 and the wait continues. A connection
122 that opens and then sends nothing is dropped after ``connection_timeout_seconds`` so it
123 cannot hold the single-threaded wait past its deadline."""
125 def __init__(self, expected_state: str, connection_timeout_seconds: float = 5) -> None:
126 super().__init__(("127.0.0.1", 0), _CallbackHandler)
127 self.expected_state: Final = expected_state
128 self.connection_timeout_seconds: Final = connection_timeout_seconds
129 self.outcome: CallbackOutcome | None = None
130 self.timeout = 1
132 @property
133 def redirect_uri(self) -> str:
134 return f"http://127.0.0.1:{self.server_address[1]}{CALLBACK_PATH}"
136 def get_request(self) -> tuple[socket.socket, object]:
137 accepted: Final[tuple[socket.socket, object]] = super().get_request()
138 accepted[0].settimeout(self.connection_timeout_seconds)
139 return accepted
141 def wait(
142 self, timeout_seconds: float, clock: Callable[[], float] = time.monotonic
143 ) -> CallbackOutcome | PkceFailure:
144 deadline: Final = clock() + timeout_seconds
145 while self.outcome is None:
146 if clock() >= deadline:
147 return PkceFailure("timed out waiting for the browser sign-in to finish")
148 self.handle_request()
149 return self.outcome
152class _CallbackHandler(BaseHTTPRequestHandler):
153 server: LoopbackServer # pyright: ignore[reportIncompatibleVariableOverride] # only ever constructed by LoopbackServer
155 def do_GET(self) -> None:
156 parsed: Final = urlparse(self.path)
157 if parsed.path != CALLBACK_PATH:
158 self._respond(404, "Not found.")
159 return
160 params: Final = parse_qs(parsed.query)
161 if _first(params, "state") != self.server.expected_state:
162 self._respond(400, "This response does not belong to the pending sign-in; still waiting.")
163 return
164 error: Final = _first(params, "error")
165 if error is not None:
166 self.server.outcome = CallbackDenied(error=error, description=_first(params, "error_description"))
167 self._respond(200, "Sign-in was not approved. You can close this window.")
168 return
169 code: Final = _first(params, "code")
170 if code is None:
171 self._respond(400, "The sign-in response carried no authorization code; still waiting.")
172 return
173 self.server.outcome = CallbackCode(code=code)
174 self._respond(200, "Signed in to LiteLLM. You can close this window and return to the terminal.")
176 def log_message(self, format: str, *args: object) -> None:
177 return
179 def _respond(self, status: int, text: str) -> None:
180 body: Final = text.encode("utf-8")
181 self.send_response(status)
182 self.send_header("Content-Type", "text/plain; charset=utf-8")
183 self.send_header("Content-Length", str(len(body)))
184 self.send_header("Cache-Control", "no-store")
185 self.end_headers()
186 self.wfile.write(body)
189def _first(params: Mapping[str, Sequence[str]], key: str) -> str | None:
190 values: Final = params.get(key)
191 return values[0] if values else None
194def discover_cli_auth(base_url: str, http: Http) -> CliAuthContract | PkceFailure:
195 url: Final = f"{base_url.rstrip('/')}{CLI_AUTH_DISCOVERY_PATH}"
196 try:
197 response: Final = http.get(url, timeout=_HTTP_TIMEOUT_SECONDS)
198 except requests.RequestException as exc:
199 return PkceFailure(f"could not reach {url}: {exc}")
200 if response.status_code != 200:
201 return PkceFailure(
202 f"{url} answered {response.status_code}; this proxy version does not support `lite login --pkce`"
203 )
204 try:
205 contract: Final = CliAuthContract.model_validate(response.json())
206 except (ValueError, ValidationError) as exc:
207 return PkceFailure(f"{url} returned an unsupported discovery document: {exc}")
208 if "S256" not in contract.code_challenge_methods_supported:
209 return PkceFailure("the proxy does not support PKCE S256")
210 if _canonical_url(contract.issuer) != _canonical_url(base_url):
211 return PkceFailure(f"{url} is issued for {contract.issuer}, not {base_url}; pass that address as --base-url")
212 foreign: Final = _endpoints_outside(contract, _origin(base_url))
213 if foreign:
214 return PkceFailure(
215 f"{url} names endpoints outside {base_url} ({', '.join(foreign)}); refusing to send credentials there"
216 )
217 return contract
220def _endpoints_outside(contract: CliAuthContract, origin: str | None) -> tuple[str, ...]:
221 endpoints: Final = (
222 contract.authorization_endpoint,
223 contract.token_endpoint,
224 contract.registration_endpoint,
225 contract.revocation_endpoint,
226 contract.resource,
227 )
228 return tuple(endpoint for endpoint in endpoints if origin is None or _origin(endpoint) != origin)
231def _origin(url: str) -> str | None:
232 """``scheme://host:port`` with the default port made explicit, so the same server spelled
233 two ways (``https://llm.example.com`` and ``https://LLM.example.com:443/``) compares equal
234 and two different servers never do."""
235 parsed: Final = urlparse(url)
236 try:
237 port: Final = parsed.port
238 except ValueError:
239 return None
240 if parsed.scheme not in ("http", "https") or not parsed.hostname:
241 return None
242 host: Final = f"[{parsed.hostname}]" if ":" in parsed.hostname else parsed.hostname
243 return f"{parsed.scheme}://{host}:{port or (443 if parsed.scheme == 'https' else 80)}"
246def _canonical_url(url: str) -> str | None:
247 """The origin plus the path with its trailing slash dropped: the RFC 8414 section 3.3
248 identity check, so a document can only ever be accepted for the proxy it was fetched from."""
249 origin: Final = _origin(url)
250 return None if origin is None else f"{origin}{urlparse(url).path.rstrip('/')}"
253class _ClientRegistration(TypedDict):
254 client_name: ReadOnly[str]
255 redirect_uris: ReadOnly[tuple[str, ...]]
256 grant_types: ReadOnly[tuple[str, ...]]
257 response_types: ReadOnly[tuple[str, ...]]
258 token_endpoint_auth_method: ReadOnly[Literal["none"]]
261def _form(**fields: str) -> Mapping[str, str]:
262 return MappingProxyType(fields)
265def _refused_redirect(request_name: str, response: requests.Response) -> PkceFailure | None:
266 """Every POST to the proxy is sent with ``allow_redirects=False``: a 307 or 308 would make
267 ``requests`` replay the form, code and verifier or refresh token included, wherever ``Location``
268 points, past the origin check discovery passed."""
269 if not 300 <= response.status_code < 400:
270 return None
271 return PkceFailure(
272 f"{request_name} redirected to {response.headers.get('Location', 'another address')}; refusing to follow it"
273 )
276def register_client(contract: CliAuthContract, redirect_uri: str, http: Http) -> str | PkceFailure:
277 registration: Final[_ClientRegistration] = {
278 "client_name": _CLIENT_NAME,
279 "redirect_uris": (redirect_uri,),
280 "grant_types": ("authorization_code", "refresh_token"),
281 "response_types": ("code",),
282 "token_endpoint_auth_method": "none",
283 }
284 try:
285 response: Final = http.post(
286 contract.registration_endpoint, json=registration, timeout=_HTTP_TIMEOUT_SECONDS, allow_redirects=False
287 )
288 except requests.RequestException as exc:
289 return PkceFailure(f"client registration failed: {exc}")
290 redirected: Final = _refused_redirect("client registration", response)
291 if redirected is not None:
292 return redirected
293 if response.status_code not in (200, 201):
294 return PkceFailure(f"client registration failed with {response.status_code}: {_error_detail(response)}")
295 try:
296 return _RegisteredClient.model_validate(response.json()).client_id
297 except (ValueError, ValidationError) as exc:
298 return PkceFailure(f"client registration returned an unexpected body: {exc}")
301def pkce_pair() -> tuple[str, str]:
302 verifier: Final = secrets.token_urlsafe(64)
303 digest: Final = hashlib.sha256(verifier.encode("ascii")).digest()
304 return verifier, urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
307def authorize_url(contract: CliAuthContract, client_id: str, redirect_uri: str, state: str, code_challenge: str) -> str:
308 query: Final = urlencode(
309 _form(
310 response_type="code",
311 client_id=client_id,
312 redirect_uri=redirect_uri,
313 state=state,
314 code_challenge=code_challenge,
315 code_challenge_method="S256",
316 resource=contract.resource,
317 )
318 )
319 return f"{contract.authorization_endpoint}?{query}"
322def redeem_code(
323 contract: CliAuthContract,
324 client_id: str,
325 redirect_uri: str,
326 code: str,
327 code_verifier: str,
328 http: Http,
329 now: Callable[[], float] = time.time,
330) -> PkceCredential | PkceFailure:
331 return _token_request(
332 token_endpoint=contract.token_endpoint,
333 revocation_endpoint=contract.revocation_endpoint,
334 resource=contract.resource,
335 client_id=client_id,
336 form=_form(
337 grant_type="authorization_code",
338 code=code,
339 redirect_uri=redirect_uri,
340 client_id=client_id,
341 code_verifier=code_verifier,
342 resource=contract.resource,
343 ),
344 http=http,
345 now=now,
346 )
349def refresh_credential(
350 token_endpoint: str,
351 revocation_endpoint: str,
352 resource: str,
353 client_id: str,
354 refresh_token: str,
355 http: Http,
356 now: Callable[[], float] = time.time,
357) -> PkceCredential | PkceFailure:
358 return _token_request(
359 token_endpoint=token_endpoint,
360 revocation_endpoint=revocation_endpoint,
361 resource=resource,
362 client_id=client_id,
363 form=_form(grant_type="refresh_token", refresh_token=refresh_token, client_id=client_id, resource=resource),
364 http=http,
365 now=now,
366 )
369def _token_request(
370 token_endpoint: str,
371 revocation_endpoint: str,
372 resource: str,
373 client_id: str,
374 form: Mapping[str, str],
375 http: Http,
376 now: Callable[[], float],
377) -> PkceCredential | PkceFailure:
378 try:
379 response: Final = http.post(token_endpoint, data=form, timeout=_HTTP_TIMEOUT_SECONDS, allow_redirects=False)
380 except requests.RequestException as exc:
381 return PkceFailure(f"token request failed: {exc}")
382 redirected: Final = _refused_redirect("token request", response)
383 if redirected is not None:
384 return redirected
385 if response.status_code != 200:
386 return PkceFailure(f"token request failed with {response.status_code}: {_error_detail(response)}")
387 try:
388 token: Final = _TokenResponse.model_validate(response.json())
389 except (ValueError, ValidationError) as exc:
390 return PkceFailure(f"token endpoint returned an unexpected body: {exc}")
391 return PkceCredential(
392 access_token=token.access_token,
393 refresh_token=token.refresh_token,
394 expires_at=now() + token.expires_in,
395 client_id=client_id,
396 token_endpoint=token_endpoint,
397 revocation_endpoint=revocation_endpoint,
398 resource=resource,
399 user_id=token.user_id,
400 team_id=token.team_id,
401 )
404def revoke_credential(
405 revocation_endpoint: str, client_id: str, refresh_token: str, http: Http
406) -> PkceFailure | RevocationUnavailable | None:
407 try:
408 response: Final = http.post(
409 revocation_endpoint,
410 data=_form(token=refresh_token, token_type_hint="refresh_token", client_id=client_id),
411 timeout=_HTTP_TIMEOUT_SECONDS,
412 allow_redirects=False,
413 )
414 except requests.RequestException as exc:
415 return PkceFailure(f"revocation request failed: {exc}")
416 redirected: Final = _refused_redirect("revocation request", response)
417 if redirected is not None:
418 return redirected
419 if response.status_code == 503:
420 return RevocationUnavailable(f"revocation failed with 503: {_error_detail(response)}")
421 if response.status_code != 200:
422 return PkceFailure(f"revocation failed with {response.status_code}: {_error_detail(response)}")
423 return None
426_ERROR_BODY: Final = TypeAdapter(Mapping[str, object])
429def _error_detail(response: requests.Response) -> str:
430 try:
431 body: Final = _ERROR_BODY.validate_json(response.content)
432 except ValidationError:
433 return response.text[:200]
434 return str(body.get("error_description") or body.get("error") or body.get("detail") or body)[:200]
437def run_pkce_login(
438 base_url: str,
439 http: Http,
440 open_browser: Callable[[str], object] = webbrowser.open,
441 echo: Callable[[str], None] = print,
442 timeout_seconds: float = LOGIN_TIMEOUT_SECONDS,
443) -> PkceCredential | PkceFailure:
444 contract: Final = discover_cli_auth(base_url, http)
445 if isinstance(contract, PkceFailure):
446 return contract
447 state: Final = secrets.token_urlsafe(32)
448 verifier, challenge = pkce_pair()
449 with LoopbackServer(state) as server:
450 client_id: Final = register_client(contract, server.redirect_uri, http)
451 if isinstance(client_id, PkceFailure):
452 return client_id
453 url: Final = authorize_url(contract, client_id, server.redirect_uri, state, challenge)
454 echo(f"Opening browser to: {url}")
455 echo("Approve the sign-in in your browser. Waiting...")
456 threading.Thread(target=open_browser, args=(url,), name="lite-login-browser", daemon=True).start()
457 outcome: Final = server.wait(timeout_seconds)
458 match outcome:
459 case PkceFailure():
460 return outcome
461 case CallbackDenied():
462 return PkceFailure(f"sign-in was not approved ({outcome.error}): {outcome.description or 'no details'}")
463 case CallbackCode():
464 return redeem_code(contract, client_id, server.redirect_uri, outcome.code, verifier, http)
467def pkce_token_record(base_url: str, credential: PkceCredential) -> CliTokenData:
468 record: Final[CliTokenData] = {
469 "base_url": base_url.rstrip("/"),
470 "key": credential.access_token,
471 "user_id": credential.user_id or "cli-user",
472 "user_email": "unknown",
473 "user_role": "cli",
474 "auth_header_name": "Authorization",
475 "jwt_token": "",
476 "timestamp": time.time(),
477 "expires_at": credential.expires_at,
478 "refresh_token": credential.refresh_token,
479 "client_id": credential.client_id,
480 "token_endpoint": credential.token_endpoint,
481 "revocation_endpoint": credential.revocation_endpoint,
482 "resource": credential.resource,
483 "team_id": credential.team_id,
484 }
485 return record
488def _ignore_warning(_message: str) -> None:
489 return None
492def fresh_api_key(
493 token_data: Mapping[str, object],
494 save: Callable[[CliTokenData], None],
495 http: Http,
496 *,
497 reload: Callable[[], Mapping[str, object] | None],
498 now: Callable[[], float] = time.time,
499 warn: Callable[[str], None] = _ignore_warning,
500) -> str | None:
501 """The stored key, refreshed first when it is about to expire and a refresh token is
502 on file. The refresh fires at the same moment ``is_cli_token_fresh`` stops calling the
503 key fresh, so a command that checks freshness and then asks for the key never disagrees
504 with itself. The rotated pair is saved before the new key is returned, so a crash after
505 this point never strands the CLI with a burned refresh token. A refresh that fails
506 reads the record again, because a sibling ``lite`` process may have rotated the pair
507 first, in which case the key it saved for this same proxy is the live one; when no sibling
508 did, the reason the proxy gave goes to ``warn`` so a revoked or refused refresh token is
509 never a silent failure. A record without ``expires_at`` (the classic ``lite login``
510 credential) is returned as stored."""
511 key: Final = token_data.get("key")
512 if not isinstance(key, str) or not key:
513 return None
514 expires_at: Final = token_data.get("expires_at")
515 if not isinstance(expires_at, (int, float)):
516 return key
517 if now() < expires_at - CLI_TOKEN_FRESHNESS_BUFFER_SECONDS:
518 return key
519 still_valid: Final = key if now() < expires_at else None
520 refresh_inputs: Final = _refresh_inputs(token_data)
521 if refresh_inputs is None:
522 return still_valid
523 refreshed: Final = refresh_credential(*refresh_inputs, http=http, now=now)
524 if isinstance(refreshed, PkceFailure):
525 sibling_key: Final = _key_rotated_by_a_sibling(reload(), token_data, now())
526 if sibling_key is None:
527 warn(f"Could not renew the key: {refreshed.reason}")
528 return sibling_key or still_valid
529 base_url: Final = token_data.get("base_url")
530 save(pkce_token_record(base_url if isinstance(base_url, str) else "", refreshed))
531 return refreshed.access_token
534_CREDENTIAL_IDENTITY_FIELDS: Final = ("base_url", "token_endpoint", "resource", "user_id", "team_id")
537def _key_rotated_by_a_sibling(
538 record: Mapping[str, object] | None, token_data: Mapping[str, object], now: float
539) -> str | None:
540 """The key a sibling process saved, but only when it continues this very credential:
541 same proxy, same token endpoint, same resource, same user and team, and not yet expired.
542 A concurrent ``lite login`` against a different proxy, or as someone else on this one,
543 replaces the same file, and its key must never be sent as this credential."""
544 if record is None or record.get("refresh_token") == token_data.get("refresh_token"):
545 return None
546 if any(record.get(field) != token_data.get(field) for field in _CREDENTIAL_IDENTITY_FIELDS):
547 return None
548 expires_at: Final = record.get("expires_at")
549 if not isinstance(expires_at, (int, float)) or now >= expires_at:
550 return None
551 key: Final = record.get("key")
552 return key if isinstance(key, str) and key else None
555def _refresh_inputs(token_data: Mapping[str, object]) -> tuple[str, str, str, str, str] | None:
556 values: Final = tuple(
557 token_data.get(field)
558 for field in ("token_endpoint", "revocation_endpoint", "resource", "client_id", "refresh_token")
559 )
560 if not all(isinstance(value, str) and value for value in values):
561 return None
562 token_endpoint, revocation_endpoint, resource, client_id, refresh_token = values
563 return str(token_endpoint), str(revocation_endpoint), str(resource), str(client_id), str(refresh_token)
566def revoke_stored_credential(
567 token_data: Mapping[str, object], http: Http
568) -> PkceFailure | RevocationUnavailable | None:
569 refresh_inputs: Final = _refresh_inputs(token_data)
570 if refresh_inputs is None:
571 return None
572 _, revocation_endpoint, _, client_id, refresh_token = refresh_inputs
573 return revoke_credential(revocation_endpoint, client_id, refresh_token, http)