import asyncio import time from dataclasses import dataclass from typing import Any import httpx import jwt import phonenumbers from jwt import ExpiredSignatureError, InvalidTokenError, PyJWK from app.settings import Settings class AuthError(Exception): def __init__(self, code: str = "unauthorized") -> None: self.code = code @dataclass(frozen=True, slots=True) class Principal: subject: str phone_number: str | None claims: dict[str, Any] class JWKSValidator: def __init__(self, settings: Settings, http: httpx.AsyncClient) -> None: self.settings = settings self.http = http self._keys: dict[str, PyJWK] = {} self._loaded_at = 0.0 self._lock = asyncio.Lock() async def refresh(self) -> None: async with self._lock: discovery_url = ( f"{str(self.settings.keycloak_internal_url).rstrip('/')}" f"/realms/{self.settings.keycloak_realm}/.well-known/openid-configuration" ) discovery = (await self.http.get(discovery_url, timeout=3)).raise_for_status().json() jwks_uri = discovery["jwks_uri"] public_prefix = str(self.settings.keycloak_public_url).rstrip("/") internal_prefix = str(self.settings.keycloak_internal_url).rstrip("/") if jwks_uri.startswith(public_prefix): jwks_uri = internal_prefix + jwks_uri[len(public_prefix) :] payload = (await self.http.get(jwks_uri, timeout=3)).raise_for_status().json() self._keys = { key["kid"]: PyJWK.from_dict(key) for key in payload.get("keys", []) if key.get("kid") and key.get("kty") == "RSA" } self._loaded_at = time.monotonic() async def validate(self, token: str) -> Principal: try: header = jwt.get_unverified_header(token) except InvalidTokenError as exc: raise AuthError() from exc if header.get("alg") != "RS256" or not header.get("kid"): raise AuthError() kid = str(header["kid"]) stale = time.monotonic() - self._loaded_at > self.settings.jwks_cache_ttl_seconds if stale or kid not in self._keys: try: await self.refresh() except (httpx.HTTPError, KeyError, ValueError): if ( kid not in self._keys or time.monotonic() - self._loaded_at > self.settings.jwks_stale_grace_seconds ): raise AuthError() from None key = self._keys.get(kid) if key is None: raise AuthError() try: claims = jwt.decode( token, key.key, algorithms=["RS256"], audience=self.settings.keycloak_audience, issuer=self.settings.issuer, options={"require": ["exp", "sub"]}, leeway=30, ) except ExpiredSignatureError as exc: raise AuthError("token_expired") from exc except InvalidTokenError as exc: raise AuthError() from exc subject = claims.get("sub") if not isinstance(subject, str) or not subject: raise AuthError() return Principal(subject, canonical_phone(claims), claims) def canonical_phone(claims: dict[str, Any]) -> str | None: for name in ("phone_number", "preferred_username"): value = claims.get(name) if not isinstance(value, str): continue try: number = phonenumbers.parse(value, None) except phonenumbers.NumberParseException: continue if phonenumbers.is_valid_number(number) and value.startswith("+"): return phonenumbers.format_number(number, phonenumbers.PhoneNumberFormat.E164) return None