| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566 |
- 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}
|