This commit is contained in:
2026-05-25 17:08:18 +08:00
parent df501f6151
commit f1259b71b5
9178 changed files with 1626125 additions and 0 deletions
+142
View File
@@ -0,0 +1,142 @@
import json
import logging
from datetime import datetime, timezone, timedelta
from fastapi import APIRouter, Depends, Query, WebSocket, WebSocketDisconnect
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_db, get_current_user
from app.models.user import User
from app.models.notification import Notification
from app.schemas.notification import NotificationOut, UnreadCountOut
from app.services.auth import decode_access_token
from app.services.notification import (
get_notifications,
mark_read,
mark_all_read,
get_unread_count,
create_notification,
)
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()
router = APIRouter(prefix="/notifications", tags=["notifications"])
# Active WebSocket connections: user_id -> set of WebSockets
_active_connections: dict[str, set[WebSocket]] = {}
async def push_notification_to_user(user_id: str, notification: Notification) -> None:
"""Push a notification to all active WebSocket connections for a user."""
connections = _active_connections.get(user_id, set())
if not connections:
return
message = json.dumps({
"id": notification.id,
"title": notification.title,
"content": notification.content,
"type": notification.type,
"is_read": notification.is_read,
"created_at": _to_local_str(notification.created_at),
}, ensure_ascii=False)
dead: list[WebSocket] = []
for ws in connections:
try:
await ws.send_text(message)
except Exception:
dead.append(ws)
for ws in dead:
connections.discard(ws)
@router.websocket("/ws")
async def notifications_ws(websocket: WebSocket):
"""Authenticated WebSocket for real-time notifications."""
await websocket.accept()
# Authenticate via token query param or first message
token = websocket.query_params.get("token")
if not token:
try:
data = await websocket.receive_text()
msg = json.loads(data)
token = msg.get("token")
except Exception:
await websocket.close(code=4001, reason="Authentication required")
return
user_id = decode_access_token(token) if token else None
if not user_id or user_id.startswith("captcha:"):
await websocket.close(code=4001, reason="Invalid token")
return
# Register connection
if user_id not in _active_connections:
_active_connections[user_id] = set()
_active_connections[user_id].add(websocket)
try:
# Keep connection alive, listen for pings
while True:
data = await websocket.receive_text()
# Echo back for heartbeat
if data == "ping":
await websocket.send_text("pong")
except WebSocketDisconnect:
pass
except Exception:
pass
finally:
conns = _active_connections.get(user_id)
if conns:
conns.discard(websocket)
if not conns:
_active_connections.pop(user_id, None)
@router.get("", response_model=list[NotificationOut])
async def list_notifications(
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=100),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
items, _ = await get_notifications(db, current_user.id, page, page_size)
return items
@router.put("/{notification_id}/read")
async def read_notification(
notification_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
await mark_read(db, notification_id, current_user.id)
return {"message": "ok"}
@router.put("/read-all")
async def read_all(
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
await mark_all_read(db, current_user.id)
return {"message": "ok"}
@router.get("/unread-count", response_model=UnreadCountOut)
async def unread_count(
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
count = await get_unread_count(db, current_user.id)
return UnreadCountOut(count=count)