98 lines
3.4 KiB
Python
98 lines
3.4 KiB
Python
"""SQLAlchemy-backed notification repository."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from datetime import UTC, datetime
|
|
from typing import Any
|
|
|
|
from sqlalchemy import func, select, update
|
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
|
|
|
from deerflow.persistence.notifications.base import NotificationStore
|
|
from deerflow.persistence.notifications.model import NotificationRow
|
|
|
|
|
|
def _row_to_dict(row: NotificationRow) -> dict[str, Any]:
|
|
data = row.to_dict()
|
|
value = data.get("created_at")
|
|
if isinstance(value, datetime):
|
|
data["created_at"] = value.isoformat()
|
|
return data
|
|
|
|
|
|
class NotificationRepository(NotificationStore):
|
|
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
|
|
self._sf = session_factory
|
|
|
|
async def create(self, data: dict[str, Any]) -> dict[str, Any]:
|
|
row = NotificationRow(
|
|
id=data["id"],
|
|
user_id=data["user_id"],
|
|
type=data["type"],
|
|
title=data.get("title") or "",
|
|
body=data.get("body") or "",
|
|
payload=data.get("payload") or {},
|
|
read=False,
|
|
created_at=datetime.now(UTC),
|
|
)
|
|
async with self._sf() as session:
|
|
session.add(row)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return _row_to_dict(row)
|
|
|
|
async def list_for_user(
|
|
self,
|
|
user_id: str,
|
|
*,
|
|
unread_only: bool = False,
|
|
limit: int = 50,
|
|
) -> list[dict[str, Any]]:
|
|
stmt = select(NotificationRow).where(NotificationRow.user_id == user_id)
|
|
if unread_only:
|
|
stmt = stmt.where(NotificationRow.read.is_(False))
|
|
stmt = stmt.order_by(NotificationRow.created_at.desc()).limit(limit)
|
|
async with self._sf() as session:
|
|
result = await session.execute(stmt)
|
|
return [_row_to_dict(row) for row in result.scalars()]
|
|
|
|
async def unread_count(self, user_id: str) -> int:
|
|
stmt = (
|
|
select(func.count())
|
|
.select_from(NotificationRow)
|
|
.where(NotificationRow.user_id == user_id, NotificationRow.read.is_(False))
|
|
)
|
|
async with self._sf() as session:
|
|
result = await session.execute(stmt)
|
|
return int(result.scalar_one() or 0)
|
|
|
|
async def mark_read(self, notification_id: str, user_id: str) -> bool:
|
|
async with self._sf() as session:
|
|
row = await session.get(NotificationRow, notification_id)
|
|
if row is None or row.user_id != user_id:
|
|
return False
|
|
if not row.read:
|
|
row.read = True
|
|
await session.commit()
|
|
return True
|
|
|
|
async def mark_all_read(self, user_id: str) -> int:
|
|
stmt = (
|
|
update(NotificationRow)
|
|
.where(NotificationRow.user_id == user_id, NotificationRow.read.is_(False))
|
|
.values(read=True)
|
|
)
|
|
async with self._sf() as session:
|
|
result = await session.execute(stmt)
|
|
await session.commit()
|
|
return int(result.rowcount or 0)
|
|
|
|
async def delete(self, notification_id: str, user_id: str) -> bool:
|
|
async with self._sf() as session:
|
|
row = await session.get(NotificationRow, notification_id)
|
|
if row is None or row.user_id != user_id:
|
|
return False
|
|
await session.delete(row)
|
|
await session.commit()
|
|
return True
|