133 lines
4.1 KiB
Python
133 lines
4.1 KiB
Python
from __future__ import annotations
|
|
|
|
from datetime import UTC, datetime
|
|
|
|
from sqlalchemy import delete, select
|
|
|
|
from app.core.database import session_scope
|
|
from app.modules.auth.models import EmailVerificationToken, PasswordResetToken, RefreshToken
|
|
|
|
|
|
def _detach(db, instance):
|
|
db.refresh(instance)
|
|
db.expunge(instance)
|
|
return instance
|
|
|
|
|
|
def create_refresh_token(
|
|
user_id: str,
|
|
token_hash: str,
|
|
family_id: str,
|
|
expires_at: datetime,
|
|
) -> RefreshToken:
|
|
with session_scope() as db:
|
|
token = RefreshToken(
|
|
user_id=user_id,
|
|
token_hash=token_hash,
|
|
family_id=family_id,
|
|
expires_at=expires_at,
|
|
)
|
|
db.add(token)
|
|
db.flush()
|
|
return _detach(db, token)
|
|
|
|
|
|
def get_refresh_token(token_hash: str) -> RefreshToken | None:
|
|
with session_scope() as db:
|
|
token = db.scalar(select(RefreshToken).where(RefreshToken.token_hash == token_hash))
|
|
if not token:
|
|
return None
|
|
return _detach(db, token)
|
|
|
|
|
|
def revoke_refresh_token(token_hash: str) -> RefreshToken | None:
|
|
with session_scope() as db:
|
|
token = db.scalar(select(RefreshToken).where(RefreshToken.token_hash == token_hash))
|
|
if not token:
|
|
return None
|
|
token.revoked_at = datetime.now(UTC)
|
|
db.flush()
|
|
return _detach(db, token)
|
|
|
|
|
|
def revoke_user_families(user_id: str) -> None:
|
|
with session_scope() as db:
|
|
tokens = db.scalars(
|
|
select(RefreshToken).where(
|
|
RefreshToken.user_id == user_id,
|
|
RefreshToken.revoked_at.is_(None),
|
|
)
|
|
).all()
|
|
now = datetime.now(UTC)
|
|
for token in tokens:
|
|
token.revoked_at = now
|
|
|
|
|
|
def revoke_family_tokens(family_id: str) -> None:
|
|
with session_scope() as db:
|
|
tokens = db.scalars(
|
|
select(RefreshToken).where(
|
|
RefreshToken.family_id == family_id,
|
|
RefreshToken.revoked_at.is_(None),
|
|
)
|
|
).all()
|
|
now = datetime.now(UTC)
|
|
for token in tokens:
|
|
token.revoked_at = now
|
|
|
|
|
|
def create_email_verification_token(user_id: str, token_hash: str, expires_at: datetime) -> None:
|
|
with session_scope() as db:
|
|
db.execute(
|
|
delete(EmailVerificationToken).where(
|
|
EmailVerificationToken.user_id == user_id,
|
|
EmailVerificationToken.used_at.is_(None),
|
|
)
|
|
)
|
|
db.add(EmailVerificationToken(user_id=user_id, token_hash=token_hash, expires_at=expires_at))
|
|
|
|
|
|
def get_email_verification_token(token_hash: str) -> EmailVerificationToken | None:
|
|
with session_scope() as db:
|
|
token = db.scalar(
|
|
select(EmailVerificationToken).where(EmailVerificationToken.token_hash == token_hash)
|
|
)
|
|
if not token:
|
|
return None
|
|
return _detach(db, token)
|
|
|
|
|
|
def mark_email_verification_token_used(token_hash: str) -> None:
|
|
with session_scope() as db:
|
|
token = db.scalar(
|
|
select(EmailVerificationToken).where(EmailVerificationToken.token_hash == token_hash)
|
|
)
|
|
if token:
|
|
token.used_at = datetime.now(UTC)
|
|
|
|
|
|
def create_password_reset_token(user_id: str, token_hash: str, expires_at: datetime) -> None:
|
|
with session_scope() as db:
|
|
db.execute(
|
|
delete(PasswordResetToken).where(
|
|
PasswordResetToken.user_id == user_id,
|
|
PasswordResetToken.used_at.is_(None),
|
|
)
|
|
)
|
|
db.add(PasswordResetToken(user_id=user_id, token_hash=token_hash, expires_at=expires_at))
|
|
|
|
|
|
def get_password_reset_token(token_hash: str) -> PasswordResetToken | None:
|
|
with session_scope() as db:
|
|
token = db.scalar(select(PasswordResetToken).where(PasswordResetToken.token_hash == token_hash))
|
|
if not token:
|
|
return None
|
|
return _detach(db, token)
|
|
|
|
|
|
def mark_password_reset_token_used(token_hash: str) -> None:
|
|
with session_scope() as db:
|
|
token = db.scalar(select(PasswordResetToken).where(PasswordResetToken.token_hash == token_hash))
|
|
if token:
|
|
token.used_at = datetime.now(UTC)
|