Files
flowdeck/app/auth/providers/oidc_provider.py
T
bruno 937ecfc2e0
FlowDeck CI / lint (push) Canceled after 0s
FlowDeck CI / test (push) Canceled after 0s
FlowDeck CI / docker (push) Canceled after 0s
fix: A30 + A37 + A39 + A40 + A41 — fin du P2/P3 XS/S (v7.4.0)
- A30 — `require_scope()` câblé : 69 sites stricts de api_v2.py passent par la
  factory (Bearer + scope en 1 appel, contrôle manuel supprimé) ; sémantique
  alignée sur celle des handlers (pas de default "read" → 0 changement de
  comportement) ; 12 top-level morts supprimés (0 ref app ET tests) :
  unsync_block, find_referring, _b64url, strip_markdown, format_number,
  get_auto_property_value, get_next_unique_id, local_date_in_tz,
  verify_device_token, _get_dynamic_groups, _require_user_gitea,
  validate_upload_request
- A37 — CORS sans `*` : origines = app_base_url + allow_origin_regex
  (localhost/dev, origines d'extension pour le Web Clipper), méthodes et
  entêtes minutées, allow_credentials explicite + test test_cors_no_star
- A39 — htmx : décision « rien » documentée (32 attributs hx-* réels sur 6
  templates, conversion = refonte du view-switching sans test E2E)
- A40 — version d'assets à source unique : ENV.globals["asset_version"] lu au
  boot depuis le fichier VERSION ; littéraux `?v=` de base.html éliminés ;
  test test_asset_version_single_source
- A41 — app.css : 91 règles mortes purgées (-10 274 octets, 121 618 → 111 344),
  scan templates/JS/CSS/Python à 0 référence

suite **1031/1031** · `ruff check app tests` OK · OpenAPI 511 chemins / 7.4.0
docs (ROADMAP/CHANGELOG/WORKLOAD/VERSION) à jour
2026-10-01 09:30:37 -04:00

218 lines
7.7 KiB
Python

"""OIDC provider — authorization code flow with PKCE (v6.7.0).
Discovery (``.well-known/openid-configuration``) is cached for an hour, the
ID token signature is verified against the issuer JWKS via authlib's JOSE
implementation, and ``iss`` / ``aud`` / ``exp`` / ``nonce`` are checked here
explicitly so the rules are visible and unit-testable.
"""
from __future__ import annotations
import base64
import hashlib
import json
import logging
import secrets
import time
import warnings
import httpx
logger = logging.getLogger(__name__)
#: Default attribute mapping (design doc §3.2) — OIDC claim names.
DEFAULT_OIDC_MAPPING: dict[str, str] = {
"login": "sub",
"email": "email",
"full_name": "name",
"avatar_url": "picture",
"groups": "groups",
}
_DISCOVERY_TTL = 3600.0
_discovery_cache: dict[str, tuple[float, dict]] = {}
class OIDCError(Exception):
"""OIDC processing failure — ``message`` is user-facing."""
def pkce_pair() -> tuple[str, str]:
"""Return ``(code_verifier, code_challenge)`` for the S256 method."""
verifier = secrets.token_urlsafe(64)
digest = hashlib.sha256(verifier.encode("ascii")).digest()
challenge = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
return verifier, challenge
def _b64url_decode(data: str) -> bytes:
return base64.urlsafe_b64decode(data + "=" * (-len(data) % 4))
async def discover(issuer_url: str) -> dict:
"""Fetch (and cache) the issuer's OIDC discovery document."""
issuer = issuer_url.rstrip("/")
url = f"{issuer}/.well-known/openid-configuration"
now = time.time()
hit = _discovery_cache.get(issuer)
if hit and now - hit[0] < _DISCOVERY_TTL:
return hit[1]
try:
async with httpx.AsyncClient(timeout=15) as client:
r = await client.get(url)
r.raise_for_status()
doc = r.json()
except Exception as err:
raise OIDCError(f"OIDC discovery failed ({url}): {err}") from err
if not doc.get("authorization_endpoint") or not doc.get("token_endpoint"):
raise OIDCError("OIDC discovery document is missing authorization/token endpoints")
_discovery_cache[issuer] = (now, doc)
return doc
def build_authorize_url(
doc: dict,
*,
client_id: str,
redirect_uri: str,
scope: str,
state: str,
nonce: str,
code_challenge: str,
) -> str:
from urllib.parse import urlencode
params = {
"client_id": client_id,
"redirect_uri": redirect_uri,
"response_type": "code",
"scope": scope or "openid profile email",
"state": state,
"nonce": nonce,
"code_challenge": code_challenge,
"code_challenge_method": "S256",
}
sep = "&" if "?" in doc["authorization_endpoint"] else "?"
return doc["authorization_endpoint"] + sep + urlencode(params)
async def exchange_code(
doc: dict, *, client_id: str, client_secret: str, code: str, redirect_uri: str, code_verifier: str
) -> dict:
"""Exchange the authorization code for tokens (PKCE, confidential client)."""
data = {
"grant_type": "authorization_code",
"code": code,
"redirect_uri": redirect_uri,
"client_id": client_id,
"code_verifier": code_verifier,
}
auth = None
if client_secret:
auth = (client_id, client_secret)
try:
async with httpx.AsyncClient(timeout=15) as client:
r = await client.post(doc["token_endpoint"], data=data, auth=auth)
except Exception as err:
raise OIDCError(f"OIDC token request failed: {err}") from err
if r.status_code != 200:
raise OIDCError(f"OIDC token endpoint returned {r.status_code}: {r.text[:300]}")
try:
tokens = r.json()
except Exception as err:
raise OIDCError(f"OIDC token endpoint returned a non-JSON body: {err}") from err
if "error" in tokens:
raise OIDCError(f"OIDC error: {tokens.get('error')} {tokens.get('error_description', '')}".strip())
return tokens
async def fetch_userinfo(doc: dict, access_token: str) -> dict:
"""Best-effort userinfo fetch (groups often only live there)."""
endpoint = doc.get("userinfo_endpoint")
if not endpoint or not access_token:
return {}
try:
async with httpx.AsyncClient(timeout=15) as client:
r = await client.get(endpoint, headers={"Authorization": f"Bearer {access_token}"})
if r.status_code != 200:
return {}
data = r.json()
return data if isinstance(data, dict) else {}
except Exception as err: # userinfo is optional enrichment
logger.debug("userinfo fetch failed: %s", err)
return {}
def validate_id_token(
id_token: str, *, issuer: str, client_id: str, nonce: str, jwks: dict
) -> dict:
"""Verify the ID token signature and claims. Returns the claims dict."""
with warnings.catch_warnings():
warnings.simplefilter("ignore", DeprecationWarning)
from authlib.jose import JsonWebKey
from authlib.jose import jwt as jose_jwt
if isinstance(id_token, bytes):
# authlib's jose.jwt.encode() returns bytes; IdP token endpoints send
# str — accept both instead of crashing on ``bytes.count(".")``.
id_token = id_token.decode()
if not id_token or id_token.count(".") != 2:
raise OIDCError("Missing or malformed ID token")
try:
keyset = JsonWebKey.import_key_set(jwks)
except Exception as err:
raise OIDCError(f"Invalid issuer JWKS: {err}") from err
try:
# Pick the key matching the token header (kid) when several are offered.
header = json.loads(_b64url_decode(id_token.split(".")[0]))
kid = header.get("kid")
key = keyset.get_by_kid(kid) if kid and hasattr(keyset, "get_by_kid") else None
token_obj = jose_jwt.decode(id_token, key or keyset)
except Exception as err:
raise OIDCError(f"ID token signature verification failed: {err}") from err
claims = dict(token_obj) # authlib's JWTClaims is a dict subclass
now = int(time.time())
if claims.get("iss") != issuer.rstrip("/") and claims.get("iss") != issuer:
raise OIDCError(f"ID token issuer mismatch: {claims.get('iss')!r}")
aud = claims.get("aud")
aud_list = aud if isinstance(aud, list) else [aud]
if client_id not in aud_list:
raise OIDCError("ID token audience does not include this client")
exp = claims.get("exp")
if not isinstance(exp, int) or exp < now:
raise OIDCError("ID token expired")
iat = claims.get("iat")
if isinstance(iat, int) and iat > now + 300:
raise OIDCError("ID token issued in the future")
if nonce and claims.get("nonce") != nonce:
raise OIDCError("ID token nonce mismatch")
if not claims.get("sub"):
raise OIDCError("ID token has no subject")
return claims
def claims_to_identity(claims: dict, mapping: dict | None = None) -> dict:
"""Map OIDC claims onto the shared ``{login, email, full_name, avatar_url, groups}`` shape."""
mapping = mapping or DEFAULT_OIDC_MAPPING
identity: dict = {"_raw": claims}
for field in ("login", "email", "full_name", "avatar_url"):
source = mapping.get(field) or field
value = claims.get(source, "")
if isinstance(value, list):
value = value[0] if value else ""
identity[field] = str(value or "").strip()
groups = claims.get(mapping.get("groups", "groups"), [])
if isinstance(groups, str):
groups = [groups]
identity["groups"] = [str(g) for g in groups if g]
if not identity["email"]:
identity["email"] = claims.get("email", "") or ""
if not identity["full_name"]:
identity["full_name"] = claims.get("name", "") or identity["email"]
return identity