deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/persistence/notifications/sql.py
2026-09-07 18:24:55 +08:00

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