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)