1. materials.py - get_material_groups和get_material_items补全Depends(get_current_user) - 移除TYPE_FILTER="1=1"死代码及4处f-string SQL拼接 - 全部SQL改为纯参数化text()查询 2. notifications.py - list_notifications废弃user_id查询参数(越权漏洞) - user_id强制从JWT Token解析,防止篡改参数偷看他人通知 - mark_notification_read补全鉴权
110 lines
3.6 KiB
Python
110 lines
3.6 KiB
Python
"""通知 API 端点 — 获取列表、标记已读"""
|
||
from __future__ import annotations
|
||
import uuid
|
||
from fastapi import APIRouter, Depends, Query
|
||
from sqlalchemy import select, func
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
from sqlalchemy.orm import selectinload
|
||
|
||
from app.core.database import get_db
|
||
from app.models.notification import Notification
|
||
from app.models.task import Task
|
||
from app.models.product import Product
|
||
from app.schemas.notification import NotificationResponse, NotificationListResponse
|
||
from app.services.auth_service import get_current_user
|
||
|
||
router = APIRouter(prefix="/notifications", tags=["消息通知"])
|
||
|
||
|
||
@router.get("/", response_model=NotificationListResponse)
|
||
async def list_notifications(
|
||
skip: int = Query(0, ge=0),
|
||
limit: int = Query(20, ge=1, le=100),
|
||
db: AsyncSession = Depends(get_db),
|
||
current_user: dict = Depends(get_current_user),
|
||
):
|
||
"""
|
||
获取当前用户的通知列表(按时间倒序)。
|
||
|
||
安全:user_id 强制从 JWT Token 解析,不接受查询参数,
|
||
杜绝通过篡改 user_id 参数越权查看他人通知。
|
||
"""
|
||
user_id: str = current_user.get("username", "") or current_user.get("sub", "")
|
||
|
||
# 总数
|
||
count_stmt = select(func.count()).select_from(Notification).where(
|
||
Notification.user_id == user_id
|
||
)
|
||
total_result = await db.execute(count_stmt)
|
||
total = total_result.scalar() or 0
|
||
|
||
# 未读数
|
||
unread_stmt = select(func.count()).select_from(Notification).where(
|
||
Notification.user_id == user_id,
|
||
Notification.is_read.is_(False),
|
||
)
|
||
unread_result = await db.execute(unread_stmt)
|
||
unread_count = unread_result.scalar() or 0
|
||
|
||
# 列表
|
||
stmt = (
|
||
select(Notification)
|
||
.where(Notification.user_id == user_id)
|
||
.order_by(Notification.created_at.desc())
|
||
.offset(skip)
|
||
.limit(limit)
|
||
)
|
||
result = await db.execute(stmt)
|
||
notifications = result.scalars().all()
|
||
|
||
# 🚀 批量查询关联的 product_serial_number
|
||
task_ids = [n.task_id for n in notifications if n.task_id]
|
||
serial_map: dict[uuid.UUID, str] = {}
|
||
if task_ids:
|
||
task_result = await db.execute(
|
||
select(Task.id, Product.serial_number)
|
||
.join(Product, Task.product_id == Product.id)
|
||
.where(Task.id.in_(task_ids))
|
||
)
|
||
for row in task_result:
|
||
serial_map[row[0]] = row[1]
|
||
|
||
# 组装响应
|
||
response_list: list[NotificationResponse] = []
|
||
for n in notifications:
|
||
resp = NotificationResponse.model_validate(n)
|
||
if n.task_id and n.task_id in serial_map:
|
||
resp.product_serial_number = serial_map[n.task_id]
|
||
response_list.append(resp)
|
||
|
||
return NotificationListResponse(
|
||
notifications=response_list,
|
||
total=total,
|
||
unread_count=unread_count,
|
||
)
|
||
|
||
|
||
@router.put("/{notification_id}/read", response_model=NotificationResponse)
|
||
async def mark_notification_read(
|
||
notification_id: str,
|
||
db: AsyncSession = Depends(get_db),
|
||
current_user: dict = Depends(get_current_user),
|
||
):
|
||
"""标记单条通知为已读"""
|
||
nid = uuid.UUID(notification_id)
|
||
result = await db.execute(
|
||
select(Notification).where(Notification.id == nid)
|
||
)
|
||
notification = result.scalar_one_or_none()
|
||
if not notification:
|
||
from fastapi import HTTPException, status
|
||
raise HTTPException(
|
||
status_code=status.HTTP_404_NOT_FOUND,
|
||
detail=f"通知不存在: {notification_id}",
|
||
)
|
||
|
||
notification.is_read = True
|
||
await db.commit()
|
||
await db.refresh(notification)
|
||
return NotificationResponse.model_validate(notification)
|