HR-ATS-Portal/backend/notifications/models.py

237 lines
8.3 KiB
Python

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
)
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
)
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