Files
video-gen/video-gen-api/app/services/notification.py
T
2026-06-29 09:09:58 +08:00

368 lines
12 KiB
Python

import logging
from datetime import datetime, timezone, timedelta
from sqlalchemy import select, func, update
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.notification import Notification
from app.models.notification_read import NotificationRead
from app.utils.id_gen import generate_id
logger = logging.getLogger("videogen")
CST = timezone(timedelta(hours=8))
def _to_local_str(dt):
if dt is None:
return None
if dt.tzinfo and dt.tzinfo.utcoffset(None) == timedelta(0):
dt = dt.astimezone(CST)
return dt.replace(tzinfo=None).isoformat()
async def create_notification(
db: AsyncSession,
user_id: str | None,
title: str,
content: str,
notif_type: str = "system",
related_id: str | None = None,
push_ws: bool = True,
) -> Notification:
"""Create a notification. user_id=None means broadcast."""
notif = Notification(
id=generate_id(),
user_id=user_id,
title=title,
content=content,
type=notif_type,
related_id=related_id,
)
db.add(notif)
await db.flush()
# Push via WebSocket if available
if push_ws:
try:
from app.utils.redis import get_redis
redis = get_redis()
if redis:
import json
payload = json.dumps(
{
"event": "notification",
"data": {
"id": notif.id,
"title": notif.title,
"content": notif.content,
"type": notif.type,
"user_id": notif.user_id,
"created_at": _to_local_str(notif.created_at),
},
}
)
if notif.user_id:
await redis.publish(f"user:{notif.user_id}:notifications", payload)
else:
await redis.publish("broadcast:notifications", payload)
except Exception:
pass
return notif
async def get_notifications(
db: AsyncSession, user_id: str, page: int = 1, page_size: int = 20, is_read: bool = None
) -> tuple[list[dict], int]:
"""Get notifications for a user (personal + broadcast), with per-user read state.
Optimized query approach:
- Use NOT EXISTS for broadcast read status instead of LEFT JOIN + CASE
- Split conditions to enable index usage
"""
# Personal notifications query (direct is_read column)
personal_cond = Notification.user_id == user_id
if is_read is not None:
personal_cond = personal_cond & (Notification.is_read == is_read)
personal_query = (
select(
Notification.id,
Notification.user_id,
Notification.title,
Notification.content,
Notification.type,
Notification.is_read.label("user_is_read"),
Notification.related_id,
Notification.created_at,
Notification.updated_at,
)
.where(personal_cond)
)
# Broadcast notifications query (use NOT EXISTS for unread check)
broadcast_base = Notification.user_id.is_(None)
if is_read is True:
# Read broadcasts: EXISTS in notification_reads
broadcast_cond = broadcast_base & (
select(NotificationRead.id)
.where(
(NotificationRead.notification_id == Notification.id)
& (NotificationRead.user_id == user_id)
)
.exists()
)
elif is_read is False:
# Unread broadcasts: NOT EXISTS in notification_reads
broadcast_cond = broadcast_base & ~(
select(NotificationRead.id)
.where(
(NotificationRead.notification_id == Notification.id)
& (NotificationRead.user_id == user_id)
)
.exists()
)
else:
broadcast_cond = broadcast_base
# For broadcast notifications, user_is_read is True if exists in notification_reads
broadcast_read_flag = (
select(NotificationRead.id)
.where(
(NotificationRead.notification_id == Notification.id)
& (NotificationRead.user_id == user_id)
)
.exists()
.label("user_is_read")
)
broadcast_query = (
select(
Notification.id,
Notification.user_id,
Notification.title,
Notification.content,
Notification.type,
broadcast_read_flag,
Notification.related_id,
Notification.created_at,
Notification.updated_at,
)
.where(broadcast_cond)
)
# Count queries - optimized
personal_count_query = select(func.count(Notification.id)).where(personal_cond)
if is_read is True:
broadcast_count_cond = broadcast_base & (
select(NotificationRead.id)
.where(
(NotificationRead.notification_id == Notification.id)
& (NotificationRead.user_id == user_id)
)
.exists()
)
elif is_read is False:
broadcast_count_cond = broadcast_base & ~(
select(NotificationRead.id)
.where(
(NotificationRead.notification_id == Notification.id)
& (NotificationRead.user_id == user_id)
)
.exists()
)
else:
broadcast_count_cond = broadcast_base
broadcast_count_query = select(func.count(Notification.id)).where(broadcast_count_cond)
# Get counts
personal_count = (await db.execute(personal_count_query)).scalar() or 0
broadcast_count = (await db.execute(broadcast_count_query)).scalar() or 0
total = personal_count + broadcast_count
# Get paginated results: fetch enough from both sides to merge and sort
# Fetch up to page_size from each, then merge, sort, and slice in memory
# This is efficient because page sizes are small (20-50)
offset = (page - 1) * page_size
fetch_limit = offset + page_size
personal_items = []
if personal_count > 0:
p_result = await db.execute(
personal_query
.order_by(Notification.created_at.desc())
.limit(fetch_limit)
)
personal_items = p_result.all()
broadcast_items = []
if broadcast_count > 0:
b_result = await db.execute(
broadcast_query
.order_by(Notification.created_at.desc())
.limit(fetch_limit)
)
broadcast_items = b_result.all()
# Merge and sort by created_at descending
all_items = list(personal_items) + list(broadcast_items)
all_items.sort(key=lambda x: x.created_at, reverse=True)
# Apply pagination
paginated = all_items[offset:offset + page_size]
items = []
for row in paginated:
items.append({
"id": row.id,
"user_id": row.user_id,
"title": row.title,
"content": row.content,
"type": row.type,
"is_read": bool(row.user_is_read),
"related_id": row.related_id,
"created_at": row.created_at,
})
return items, total
async def mark_read(db: AsyncSession, notification_id: str, user_id: str) -> None:
"""Mark a notification as read for a specific user."""
result = await db.execute(
select(Notification).where(Notification.id == notification_id).limit(1)
)
notif = result.scalar_one_or_none()
if not notif:
return
if notif.user_id is not None:
# Personal notification: update the row directly
if notif.user_id == user_id:
notif.is_read = True
await db.flush()
else:
# Broadcast notification: use notification_reads table
existing = await db.execute(
select(NotificationRead).where(
NotificationRead.notification_id == notification_id,
NotificationRead.user_id == user_id,
)
.limit(1)
)
if not existing.scalar_one_or_none():
db.add(NotificationRead(
id=generate_id(),
notification_id=notification_id,
user_id=user_id,
))
await db.flush()
async def mark_all_read(db: AsyncSession, user_id: str) -> None:
"""Mark all notifications as read for a user."""
# Mark personal notifications
await db.execute(
update(Notification)
.where(
Notification.user_id == user_id,
Notification.is_read == False,
)
.values(is_read=True)
)
# Mark all broadcast notifications
broadcasts = await db.execute(
select(Notification.id)
.where(Notification.user_id.is_(None))
)
broadcast_ids = [row[0] for row in broadcasts.all()]
if broadcast_ids:
existing_reads = await db.execute(
select(NotificationRead.notification_id).where(
NotificationRead.notification_id.in_(broadcast_ids),
NotificationRead.user_id == user_id,
)
)
already_read = {row[0] for row in existing_reads.all()}
for bid in broadcast_ids:
if bid not in already_read:
db.add(NotificationRead(
id=generate_id(),
notification_id=bid,
user_id=user_id,
))
await db.flush()
async def get_unread_count(db: AsyncSession, user_id: str) -> int:
"""Get unread notification count for a user. Optimized with EXISTS."""
# Count personal unread (uses composite index)
personal_count = (
await db.execute(
select(func.count(Notification.id)).where(
Notification.user_id == user_id,
Notification.is_read == False,
)
)
).scalar() or 0
# Count broadcast unread using NOT EXISTS (efficient, can use index on notification_reads)
broadcast_count = (
await db.execute(
select(func.count(Notification.id)).where(
Notification.user_id.is_(None),
~(
select(NotificationRead.id)
.where(
(NotificationRead.notification_id == Notification.id)
& (NotificationRead.user_id == user_id)
)
.exists()
)
)
)
).scalar() or 0
return personal_count + broadcast_count
async def get_notification_read_users(
db: AsyncSession, notification_id: str, page: int = 1, page_size: int = 50
) -> tuple[list[dict], int]:
"""Get list of users who have read a specific notification."""
from app.models.user import User
query = (
select(NotificationRead, User.username)
.join(User, NotificationRead.user_id == User.id)
.where(NotificationRead.notification_id == notification_id)
.order_by(NotificationRead.created_at.desc())
)
count_query = (
select(func.count(NotificationRead.id))
.where(NotificationRead.notification_id == notification_id)
)
total = (await db.execute(count_query)).scalar() or 0
result = await db.execute(query.offset((page - 1) * page_size).limit(page_size))
rows = result.all()
items = []
for read_record, username in rows:
items.append({
"user_id": read_record.user_id,
"username": username,
"read_at": read_record.created_at,
})
return items, total