- `app/services/http_client.py` : `async with shared_client(timeout=15) as client:` remplace les 49 créations `async with httpx.AsyncClient(` de 14 fichiers (gitea ×21, providers oidc/oauth ×11, calendar ×4, automations ×3…) — le pool de connexions est réutilisé au lieu d'être recréé à chaque appel. __aexit__ no-op (le client partagé ne se ferme pas à la sortie). - Cache par (boucle d'event, kwargs) en WeakKeyDictionary : un AsyncClient n'est JAMAIS partagé entre deux loops (piège des tests « Event loop is closed ») — une boucle par test = client propre collecté avec la boucle. Clé = kwargs triés, repr() pour les valeurs non hashables (`headers=` dict → TypeError rattrapé par la suite). - Laissés délibérément : github_adapter (transport MockTransport injecté), webhook_outbound (client « own_client » fermé par la fonction). - Tests : `test_http_client_shared_and_loop_scoped` (réutilisation mêmes kwargs / cloisonné kwargs / cloisonné loop) ; le stub des webhooks patche aussi la fabrique `http_client.httpx` + purge du cache (avant : webhook_outbound.httpx patché mais la fabrique partagée créait un vrai client → réseau réel dans les tests). suite **1091/1091** · ruff OK · docs à jour
218 lines
7.7 KiB
Python
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
|
|
|
|
from app.services.http_client import shared_client
|
|
|
|
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 shared_client(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 shared_client(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 shared_client(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
|