| 1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556 |
- from __future__ import annotations
- from datetime import datetime, timedelta, timezone
- 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),
- ) -> str:
- now = now or datetime.now(timezone.utc)
- claims = {
- "sub": str(user_id),
- "roles": sorted(set(roles)),
- "iat": int(now.timestamp()),
- "exp": int((now + lifetime).timestamp()),
- "jti": new_governance_uid(),
- }
- 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(timezone.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")
- return {key: claims[key] for key in ("sub", "roles", "iat", "exp", "jti")}
|