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