import uuid from datetime import datetime, timezone from sqlalchemy import DateTime, func from sqlalchemy.ext.asyncio import AsyncSession from sqlmodel import Field, SQLModel, select def _now() -> datetime: return datetime.now(timezone.utc) class EmailConfirmationTokens(SQLModel, table=True): __tablename__ = "email_confirmation_tokens" id: uuid.UUID = Field(default_factory=uuid.uuid4, primary_key=True) user_id: uuid.UUID = Field(index=True, foreign_key="users.id") email: str = Field(index=True) token_hash: str expires_at: datetime = Field(sa_type=DateTime(timezone=True)) is_used: bool = Field(default=False) confirmed_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_token_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_token_by_user(cls, session: AsyncSession, user_id: str): """Newest unused token for a user (expiry checked in Python by the service).""" uid = cls._as_uuid(user_id) if uid is None: return None statement = ( select(cls) .where(cls.user_id == uid, 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_token(cls, session: AsyncSession, fields: dict): row = cls(**fields) session.add(row) await session.commit() return await cls.get_token_by_id(session, row.id) @classmethod async def mark_confirmed(cls, session: AsyncSession, record_id: str): row = await cls.get_token_by_id(session, record_id) if not row: return None row.is_used = True row.confirmed_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): """Retire a token without confirming it — used when the mail send fails.""" row = await cls.get_token_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_tokens_for_user(cls, session: AsyncSession, user_id: str): uid = cls._as_uuid(user_id) if uid is None: return 0 statement = select(cls).where(cls.user_id == uid, 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) class Notifications(SQLModel, table=True): __tablename__ = "notifications" id: uuid.UUID = Field(default_factory=uuid.uuid4, primary_key=True) user_id: uuid.UUID = Field(index=True, foreign_key="users.id") kind: str title: str body: str | None = Field(default=None) link_path: str | None = Field(default=None) inbox_id: int | None = Field(default=None) job_post_id: uuid.UUID | None = Field(default=None) is_read: bool = Field(default=False, index=True) created_at: datetime = Field(default_factory=_now, sa_type=DateTime(timezone=True)) is_deleted: bool = Field(default=False) @staticmethod def _as_uuid(record_id) -> uuid.UUID | None: if record_id in (None, ""): return None try: return uuid.UUID(str(record_id)) except ValueError: return None @classmethod async def get_by_id(cls, session: AsyncSession, record_id, *, user_id=None): uid = cls._as_uuid(record_id) if uid is None: return None statement = select(cls).where(cls.id == uid, cls.is_deleted == False) # noqa: E712 if user_id is not None: statement = statement.where(cls.user_id == user_id) result = await session.execute(statement) return result.scalars().first() @classmethod async def fetch_notifications( cls, session: AsyncSession, *, user_id, unread_only: bool = False, top: int | None = None, skip: int = 0, ): statement = select(cls).where( cls.user_id == user_id, cls.is_deleted == False, # noqa: E712 cls.title != "ATS score ready", ) if unread_only: statement = statement.where(cls.is_read == False) # noqa: E712 count_statement = select(func.count()).select_from(statement.subquery()) total = (await session.execute(count_statement)).scalar_one() unread_statement = select(func.count()).select_from(cls).where( cls.user_id == user_id, cls.is_deleted == False, # noqa: E712 cls.is_read == False, # noqa: E712 cls.title != "ATS score ready", ) unread = (await session.execute(unread_statement)).scalar_one() statement = statement.order_by(cls.created_at.desc()) if skip: statement = statement.offset(skip) if top is not None: statement = statement.limit(top) result = await session.execute(statement) return list(result.scalars().all()), total, unread @classmethod async def insert_notification(cls, session: AsyncSession, fields: dict): row = cls(**fields) session.add(row) await session.commit() return await cls.get_by_id(session, row.id) @classmethod async def insert_many(cls, session: AsyncSession, payloads, *, commit: bool = True): """One row per payload. `commit=False` rides the caller's transaction.""" rows = [] for fields in payloads or []: uid = cls._as_uuid(fields.get("user_id")) if uid is None or not fields.get("kind") or not fields.get("title"): continue job_id = cls._as_uuid(fields.get("job_post_id")) if fields.get("job_post_id") else None row = cls( user_id=uid, kind=fields["kind"], title=fields["title"], body=fields.get("body"), link_path=fields.get("link_path"), inbox_id=fields.get("inbox_id"), job_post_id=job_id, ) session.add(row) rows.append(row) if commit and rows: await session.commit() return rows @classmethod async def mark_read(cls, session: AsyncSession, record_id, *, user_id): row = await cls.get_by_id(session, record_id, user_id=user_id) if not row: return None row.is_read = True session.add(row) await session.commit() await session.refresh(row) return row @classmethod async def mark_all_read(cls, session: AsyncSession, user_id): statement = select(cls).where( cls.user_id == user_id, cls.is_deleted == False, # noqa: E712 cls.is_read == False, # noqa: E712 ) result = await session.execute(statement) rows = list(result.scalars().all()) for row in rows: row.is_read = True session.add(row) await session.commit() return len(rows) @classmethod async def soft_delete_notification(cls, session: AsyncSession, record_id, *, user_id): row = await cls.get_by_id(session, record_id, user_id=user_id) if not row: return None row.is_deleted = True session.add(row) await session.commit() return row