1
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user