111 lines
3.9 KiB
Python
111 lines
3.9 KiB
Python
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()
|
|
|
|
@property
|
|
def has_keys(self) -> bool:
|
|
return bool(self._keys)
|
|
|
|
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
|