"""IAM v0.3 access-token verification. No username-based authority.""" from __future__ import annotations import asyncio import ipaddress import time from dataclasses import dataclass from urllib.parse import urlsplit import httpx import jwt class AccessFailure(Exception): def __init__(self, status: int, code: str): self.status, self.code = status, code super().__init__(code) @dataclass(frozen=True) class Actor: issuer: str subject: str tenant: str principal_type: str assurance: str authenticated_at: int expires_at: int def require_https(url: str) -> None: parsed = urlsplit(url) hostname = (parsed.hostname or "").rstrip(".").lower() try: local = ipaddress.ip_address(hostname).is_loopback except ValueError: local = hostname == "localhost" or hostname.endswith(".localhost") if (parsed.scheme != "https" or not parsed.hostname or parsed.username or parsed.password or parsed.fragment or parsed.query or local): raise ValueError("a non-local HTTPS trust endpoint is required") class OIDCVerifier: """Discover keys at an explicitly trusted issuer; bounded, rotation-aware cache. Supports RFC 9068 at+jwt tokens or the admitted KeyCape Bearer payload type. ID tokens without either access-token marker are rejected. """ def __init__(self, *, issuer: str, audience: str, client: httpx.AsyncClient, key_ttl: int = 60, max_token_age: int = 300): require_https(issuer) if not audience or not 1 <= key_ttl <= 300 or not 1 <= max_token_age <= 300: raise ValueError("audience and bounded key/token lifetimes are required") self.issuer, self.audience, self.client = issuer, audience, client self.key_ttl, self.max_token_age = key_ttl, max_token_age self._keys: dict = {} self._loaded = 0.0 self._lock = asyncio.Lock() async def _refresh(self) -> None: try: response = await self.client.get( self.issuer.rstrip("/") + "/.well-known/openid-configuration", timeout=3, follow_redirects=False, ) response.raise_for_status() discovery = response.json() if discovery["issuer"] != self.issuer: raise ValueError("issuer mismatch") require_https(discovery["jwks_uri"]) response = await self.client.get(discovery["jwks_uri"], timeout=3, follow_redirects=False) response.raise_for_status() keys = {} for value in response.json()["keys"]: if value.get("kty") != "RSA" or value.get("use", "sig") != "sig": continue if value.get("alg", "RS256") != "RS256": continue if "verify" not in value.get("key_ops", ["verify"]): continue kid = value["kid"] if not isinstance(kid, str) or not kid or kid in keys: raise ValueError("invalid key IDs") key = jwt.PyJWK.from_dict(value, algorithm="RS256").key if key.key_size < 2048: raise ValueError("weak issuer key") keys[kid] = key if not keys: raise ValueError("no signing keys") self._keys, self._loaded = keys, time.monotonic() except (httpx.HTTPError, ValueError, KeyError, TypeError, jwt.PyJWTError) as exc: raise AccessFailure(503, "identity_unavailable") from exc async def authenticate(self, token: str) -> Actor: try: header = jwt.get_unverified_header(token) if header.get("alg") != "RS256" or not isinstance(header.get("kid"), str): raise ValueError("unsupported token") async with self._lock: # At most one unknown-key refresh per second, to bound random-kid traffic. age = time.monotonic() - self._loaded if age >= self.key_ttl or (header["kid"] not in self._keys and age >= 1): await self._refresh() key = self._keys.get(header["kid"]) if key is None: raise ValueError("unknown key") claims = jwt.decode(token, key, algorithms=["RS256"], issuer=self.issuer, audience=self.audience, leeway=0, options={"require": ["iss", "sub", "aud", "exp", "iat", "tenant", "principal_type", "groups", "roles", "assurance"]}) if header.get("typ") != "at+jwt" and claims.get("typ") != "Bearer": raise ValueError("not an access token") if claims.get("typ", "Bearer") != "Bearer" or claims.get("environment") in { "local", "development", "test", }: raise ValueError("unsupported token profile") for name in ("sub", "tenant"): if not isinstance(claims[name], str) or not claims[name]: raise ValueError("invalid identity") for name in ("groups", "roles"): if not isinstance(claims[name], list) or any( not isinstance(item, str) for item in claims[name] ): raise ValueError("invalid IAM array") scope = claims.get("scope", claims.get("scp")) if not isinstance(scope, (str, list)) or ( isinstance(scope, list) and any(not isinstance(item, str) for item in scope) ): raise ValueError("invalid scope") assurance = claims["assurance"] if (not isinstance(assurance, dict) or assurance.get("level") not in {"aal1", "aal2", "aal3"} or type(assurance.get("mfa")) is not bool or not isinstance(assurance.get("source"), str) or not assurance["source"] or not isinstance(assurance.get("methods"), list) or not all(isinstance(x, str) for x in assurance["methods"])): raise ValueError("invalid assurance") for value in (claims["iat"], claims["exp"], assurance.get("at")): if type(value) is not int: raise ValueError("integer timestamps required") if "nbf" in claims and type(claims["nbf"]) is not int: raise ValueError("integer not-before required") if assurance["level"] in {"aal2", "aal3"} and not assurance["mfa"]: raise ValueError("missing MFA evidence") now = time.time() if (not claims["iat"] <= now < claims["exp"] or claims["exp"] - claims["iat"] > self.max_token_age or not 0 <= now - assurance["at"] <= self.max_token_age): raise ValueError("stale identity or assurance") principal = claims["principal_type"] if principal not in {"human", "service", "agent"}: raise ValueError("invalid principal type") if principal == "agent": agent = claims.get("agent", {}) if not agent.get("id") or agent.get("mode") != "autonomous": # Delegation needs a separately admitted actor/workload contract. raise ValueError("unsupported delegation") return Actor(self.issuer, claims["sub"], claims["tenant"], principal, assurance["level"], assurance["at"], claims["exp"]) except AccessFailure: raise except (jwt.PyJWTError, ValueError, KeyError, TypeError, AttributeError) as exc: raise AccessFailure(401, "invalid_access_token") from exc