from datetime import datetime, UTC, timedelta import abc from src.config.email import EmailAddress from src.config.parse import assert_key_of_type, ParseError from typing import Any, TypeVar from dataclasses import dataclass, asdict from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey import jwt from src.services.users.data import User from src.config.auth import Auth as AuthConfig JWT = str T = TypeVar('T', bound='Claim') @dataclass class Claim(abc.ABC): sub: int exp: datetime iat: datetime @classmethod @abc.abstractmethod def from_dict(cls, claims: dict[str, Any]) -> T: pass @classmethod @abc.abstractmethod def from_user(cls, user: User, expire_in_secs: int) -> T: pass class ExpiredTokenError(Exception): def __init__(self): super().__init__('token was expired') @dataclass class BaseClaims(Claim): @classmethod def from_user(cls, user: User, expire_in_secs: int) -> 'BaseClaims': now = datetime.now(tz=UTC) expires_at = now + timedelta(seconds=expire_in_secs) return BaseClaims(sub=user.id, exp=expires_at, iat=now) @classmethod def from_dict(cls, claims: dict[str, Any]) -> 'BaseClaims': assert_key_of_type(claims, 'sub', int) assert_key_of_type(claims, 'exp', int) assert_key_of_type(claims, 'iat', int) exp = datetime.fromtimestamp(claims['exp'], UTC) now = datetime.now(UTC) if exp < now: raise ExpiredTokenError() return BaseClaims( sub=claims['sub'], exp=exp, iat=datetime.fromtimestamp(claims['iat'], UTC), ) @dataclass class EmailConfirmationClaim(BaseClaims): email: EmailAddress @classmethod def from_user(cls, user: User, expire_in_secs) -> 'EmailConfirmationClaim': base = BaseClaims.from_user(user, expire_in_secs) return EmailConfirmationClaim(**asdict(base), email=user.email) @classmethod def from_dict(cls, claims: dict[str, Any]) -> 'EmailConfirmationClaim': base = BaseClaims.from_dict(claims) assert_key_of_type(claims, 'email', str) return EmailConfirmationClaim(**asdict(base), email=claims['email']) class AuthService: def __init__(self, config: AuthConfig): self._private_key = Ed25519PrivateKey.from_private_bytes(config.ed25519_private_key.expose_secret()) self._public_key = self._private_key.public_key() def mint_claim_from_user(self, claim: type[Claim], user: User, expires_in_secs=600) -> JWT: claims = claim.from_user(user, expires_in_secs) return jwt.encode(asdict(claims), self._private_key, algorithm='EdDSA') def validate_token[T: Claim](self, claim: type[T], token: JWT) -> T | None: try: claims = jwt.decode(token, key=self._public_key, algorithms=['EdDSA']) except jwt.InvalidTokenError: return None try: return claim.from_dict(claims) except ParseError: return None