107 lines
3.5 KiB
Python
107 lines
3.5 KiB
Python
import uuid
|
|
from datetime import datetime, timezone
|
|
|
|
from sqlalchemy import DateTime
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
from sqlmodel import Field, SQLModel, select
|
|
|
|
|
|
def _now() -> datetime:
|
|
return datetime.now(timezone.utc)
|
|
|
|
|
|
class PasswordResetCodes(SQLModel, table=True):
|
|
__tablename__ = "password_reset_codes"
|
|
|
|
id: uuid.UUID = Field(default_factory=uuid.uuid4, primary_key=True)
|
|
email: str = Field(index=True)
|
|
code_hash: str
|
|
expires_at: datetime = Field(sa_type=DateTime(timezone=True))
|
|
attempts: int = Field(default=0)
|
|
is_used: bool = Field(default=False)
|
|
verified_at: datetime | None = Field(default=None, sa_type=DateTime(timezone=True))
|
|
created_at: datetime = Field(default_factory=_now, sa_type=DateTime(timezone=True))
|
|
updated_at: datetime = Field(default_factory=_now, sa_type=DateTime(timezone=True))
|
|
|
|
@staticmethod
|
|
def _as_uuid(record_id: str) -> uuid.UUID | None:
|
|
try:
|
|
return uuid.UUID(str(record_id))
|
|
except ValueError:
|
|
return None
|
|
|
|
@classmethod
|
|
async def get_code_by_id(cls, session: AsyncSession, record_id: str):
|
|
uid = cls._as_uuid(record_id)
|
|
if uid is None:
|
|
return None
|
|
result = await session.execute(select(cls).where(cls.id == uid))
|
|
return result.scalars().first()
|
|
|
|
@classmethod
|
|
async def get_active_code_by_email(cls, session: AsyncSession, email: str):
|
|
"""Newest unused code for email (expiry checked in Python by the service)."""
|
|
statement = (
|
|
select(cls)
|
|
.where(cls.email == email, cls.is_used == False) # noqa: E712
|
|
.order_by(cls.created_at.desc())
|
|
)
|
|
result = await session.execute(statement)
|
|
return result.scalars().first()
|
|
|
|
@classmethod
|
|
async def insert_code(cls, session: AsyncSession, fields: dict):
|
|
row = cls(**fields)
|
|
session.add(row)
|
|
await session.commit()
|
|
return await cls.get_code_by_id(session, row.id)
|
|
|
|
@classmethod
|
|
async def increment_attempts(cls, session: AsyncSession, record_id: str):
|
|
row = await cls.get_code_by_id(session, record_id)
|
|
if not row:
|
|
return None
|
|
row.attempts = (row.attempts or 0) + 1
|
|
row.updated_at = _now()
|
|
session.add(row)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return row
|
|
|
|
@classmethod
|
|
async def mark_verified(cls, session: AsyncSession, record_id: str):
|
|
row = await cls.get_code_by_id(session, record_id)
|
|
if not row:
|
|
return None
|
|
row.verified_at = _now()
|
|
row.updated_at = _now()
|
|
session.add(row)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return row
|
|
|
|
@classmethod
|
|
async def mark_used(cls, session: AsyncSession, record_id: str):
|
|
row = await cls.get_code_by_id(session, record_id)
|
|
if not row:
|
|
return None
|
|
row.is_used = True
|
|
row.updated_at = _now()
|
|
session.add(row)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return row
|
|
|
|
@classmethod
|
|
async def invalidate_codes_for_email(cls, session: AsyncSession, email: str):
|
|
statement = select(cls).where(cls.email == email, cls.is_used == False) # noqa: E712
|
|
result = await session.execute(statement)
|
|
rows = result.scalars().all()
|
|
now = _now()
|
|
for row in rows:
|
|
row.is_used = True
|
|
row.updated_at = now
|
|
session.add(row)
|
|
await session.commit()
|
|
return len(rows)
|