tokens.py 1.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566
  1. from __future__ import annotations
  2. from datetime import UTC, datetime, timedelta
  3. import jwt
  4. from app.core.common.identifiers import new_governance_uid
  5. class TokenError(ValueError):
  6. pass
  7. def issue_access_token(
  8. *,
  9. user_id: str,
  10. roles: list[str],
  11. secret: str,
  12. now: datetime | None = None,
  13. lifetime: timedelta = timedelta(minutes=30),
  14. session_uid: str | None = None,
  15. token_version: int | None = None,
  16. identity_source: str | None = None,
  17. ) -> str:
  18. now = now or datetime.now(UTC)
  19. claims = {
  20. "sub": str(user_id),
  21. "roles": sorted(set(roles)),
  22. "iat": int(now.timestamp()),
  23. "exp": int((now + lifetime).timestamp()),
  24. "jti": new_governance_uid(),
  25. }
  26. if session_uid is not None:
  27. claims["sid"] = str(session_uid)
  28. if token_version is not None:
  29. claims["token_version"] = int(token_version)
  30. if identity_source is not None:
  31. claims["identity_source"] = identity_source
  32. return jwt.encode(claims, secret, algorithm="HS256")
  33. def decode_access_token(
  34. token: str,
  35. *,
  36. secret: str,
  37. now: datetime | None = None,
  38. ) -> dict:
  39. try:
  40. claims = jwt.decode(
  41. token,
  42. secret,
  43. algorithms=["HS256"],
  44. options={"verify_exp": False, "require": ["sub", "roles", "iat", "exp", "jti"]},
  45. )
  46. except jwt.PyJWTError as exc:
  47. raise TokenError("invalid access token") from exc
  48. current = int((now or datetime.now(UTC)).timestamp())
  49. if int(claims["exp"]) <= current:
  50. raise TokenError("access token expired")
  51. if int(claims["iat"]) > current + 30:
  52. raise TokenError("access token issued in the future")
  53. if not isinstance(claims.get("roles"), list):
  54. raise TokenError("invalid role claims")
  55. allowed = ("sub", "roles", "iat", "exp", "jti", "sid", "token_version", "identity_source")
  56. return {key: claims[key] for key in allowed if key in claims}