user-engine/src/user_engine/oidc.py

159 lines
5.2 KiB
Python
Raw Normal View History

2026-07-28 00:06:21 +02:00
"""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 = ""
2026-07-28 00:06:21 +02:00
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,
2026-07-28 00:06:21 +02:00
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("/")
2026-07-28 00:06:21 +02:00
self.session_ttl = session_ttl
self.pending: dict[str, PendingLogin] = {}
self.sessions: dict[str, BrowserSession] = {}
def begin(self, *, tenant_hint: str | None = None) -> str:
2026-07-28 00:06:21 +02:00
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 = {
2026-07-28 00:06:21 +02:00
'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)}"
2026-07-28 00:06:21 +02:00
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",
2026-07-28 00:06:21 +02:00
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),
)
2026-07-28 00:06:21 +02:00
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
2026-07-28 00:06:21 +02:00
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")
2026-07-28 00:06:21 +02:00
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("=")