from __future__ import annotations from datetime import UTC, datetime, timedelta import jwt from app.core.common.identifiers import new_governance_uid class TokenError(ValueError): pass def issue_access_token( *, user_id: str, roles: list[str], secret: str, now: datetime | None = None, lifetime: timedelta = timedelta(minutes=30), session_uid: str | None = None, token_version: int | None = None, identity_source: str | None = None, ) -> str: now = now or datetime.now(UTC) claims = { "sub": str(user_id), "roles": sorted(set(roles)), "iat": int(now.timestamp()), "exp": int((now + lifetime).timestamp()), "jti": new_governance_uid(), } if session_uid is not None: claims["sid"] = str(session_uid) if token_version is not None: claims["token_version"] = int(token_version) if identity_source is not None: claims["identity_source"] = identity_source return jwt.encode(claims, secret, algorithm="HS256") def decode_access_token( token: str, *, secret: str, now: datetime | None = None, ) -> dict: try: claims = jwt.decode( token, secret, algorithms=["HS256"], options={"verify_exp": False, "require": ["sub", "roles", "iat", "exp", "jti"]}, ) except jwt.PyJWTError as exc: raise TokenError("invalid access token") from exc current = int((now or datetime.now(UTC)).timestamp()) if int(claims["exp"]) <= current: raise TokenError("access token expired") if int(claims["iat"]) > current + 30: raise TokenError("access token issued in the future") if not isinstance(claims.get("roles"), list): raise TokenError("invalid role claims") allowed = ("sub", "roles", "iat", "exp", "jti", "sid", "token_version", "identity_source") return {key: claims[key] for key in allowed if key in claims}