"""OIDC Authorization Code + PKCE relying-party support.""" from __future__ import annotations from dataclasses import dataclass import base64 import hashlib import json import secrets import time from typing import Any, Mapping from urllib.parse import urlencode from urllib.request import Request, urlopen @dataclass class PendingLogin: verifier: str created_at: float @dataclass class BrowserSession: claims: Mapping[str, Any] expires_at: float csrf_token: str = "" class OIDCClient: """Minimal confidential-state/public-client OIDC adapter.""" def __init__( self, *, issuer: str, client_id: str, redirect_uri: str, audience: str, backend_url: str | None = None, session_ttl: int = 3600, ) -> None: self.issuer = issuer.rstrip("/") self.client_id = client_id self.redirect_uri = redirect_uri self.audience = audience self.backend_url = (backend_url or issuer).rstrip("/") self.session_ttl = session_ttl self.pending: dict[str, PendingLogin] = {} self.sessions: dict[str, BrowserSession] = {} def begin(self, *, tenant_hint: str | None = None) -> str: state = secrets.token_urlsafe(32) verifier = secrets.token_urlsafe(64) challenge = _b64(hashlib.sha256(verifier.encode("ascii")).digest()) self.pending[state] = PendingLogin(verifier=verifier, created_at=time.time()) self._prune() parameters = { 'response_type': 'code', 'client_id': self.client_id, 'redirect_uri': self.redirect_uri, 'scope': 'openid profile email groups', 'state': state, 'code_challenge': challenge, 'code_challenge_method': 'S256', } if tenant_hint: parameters["tenant_hint"] = tenant_hint return f"{self.issuer}/authorize?{urlencode(parameters)}" def complete(self, *, code: str, state: str) -> str: pending = self.pending.pop(state, None) if pending is None or time.time() - pending.created_at > 600: raise ValueError("invalid or expired OIDC state") form = urlencode( { "grant_type": "authorization_code", "client_id": self.client_id, "redirect_uri": self.redirect_uri, "code": code, "code_verifier": pending.verifier, } ).encode() request = Request( f"{self.backend_url}/token", data=form, headers={"Content-Type": "application/x-www-form-urlencoded"}, method="POST", ) with urlopen(request, timeout=10) as response: tokens = json.loads(response.read()) token = str(tokens.get("id_token") or tokens.get("access_token") or "") claims = self._verify(token) session_id = secrets.token_urlsafe(32) expiry = min(float(claims.get("exp", time.time() + self.session_ttl)), time.time() + self.session_ttl) self.sessions[session_id] = BrowserSession( claims=claims, expires_at=expiry, csrf_token=secrets.token_urlsafe(32), ) self._prune() return session_id def claims(self, session_id: str) -> Mapping[str, Any] | None: session = self.sessions.get(session_id) if session is None or session.expires_at <= time.time(): self.sessions.pop(session_id, None) return None return session.claims def logout(self, session_id: str) -> None: self.sessions.pop(session_id, None) def csrf_token(self, session_id: str) -> str | None: session = self.sessions.get(session_id) if session is None or session.expires_at <= time.time(): self.sessions.pop(session_id, None) return None return session.csrf_token def _verify(self, token: str) -> Mapping[str, Any]: if not token: raise ValueError("OIDC token response is missing a token") try: import jwt except ImportError as exc: # pragma: no cover raise RuntimeError("install user-engine[oidc] for OIDC login") from exc jwks = jwt.PyJWKClient(f"{self.backend_url}/jwks") key = jwks.get_signing_key_from_jwt(token) claims = jwt.decode( token, key.key, algorithms=["RS256"], audience=self.audience, issuer=self.issuer, options={"require": ["exp", "iat", "iss", "sub", "aud"]}, ) return dict(claims) def _prune(self) -> None: now = time.time() self.pending = { key: value for key, value in self.pending.items() if now - value.created_at <= 600 } self.sessions = { key: value for key, value in self.sessions.items() if value.expires_at > now } def cookie_value(header: str, name: str) -> str | None: for part in header.split(";"): key, separator, value = part.strip().partition("=") if separator and key == name: return value return None def _b64(value: bytes) -> str: return base64.urlsafe_b64encode(value).decode("ascii").rstrip("=")