feat: 双Token认证(Access 2h/Refresh 7d) + 通知系统(转交/驳回自动推送)

This commit is contained in:
2026-08-07 11:44:04 +08:00
parent dc97a7385f
commit b71c5a2d07
16 changed files with 426 additions and 41 deletions

View File

@ -0,0 +1,32 @@
"""add_notifications
Revision ID: b8c9d0e1f2a3
Revises: a7b8c9d0e1f2
Create Date: 2026-08-07
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = "b8c9d0e1f2a3"
down_revision: Union[str, None] = "a7b8c9d0e1f2"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.create_table(
"notifications",
sa.Column("id", sa.UUID(), primary_key=True, server_default=sa.text("gen_random_uuid()")),
sa.Column("user_id", sa.String(64), nullable=False, index=True, comment="接收人ID(逻辑外键→老系统)"),
sa.Column("title", sa.String(200), nullable=False, comment="通知标题"),
sa.Column("content", sa.Text(), nullable=False, comment="通知内容详情"),
sa.Column("type", sa.String(20), nullable=False, comment="通知类型: TRANSFER(转交派发) | REJECT(驳回)"),
sa.Column("task_id", sa.UUID(), sa.ForeignKey("tasks.id", ondelete="SET NULL"), nullable=True, comment="关联任务ID"),
sa.Column("is_read", sa.Boolean(), server_default=sa.text("false"), comment="是否已读"),
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), comment="创建时间"),
)
def downgrade() -> None:
op.drop_table("notifications")

View File

@ -1,20 +1,32 @@
"""认证 API — 对接 MOM sys_user""" """认证 API — 对接 MOM sys_user + 双 Token 刷新"""
from fastapi import APIRouter, Depends from fastapi import APIRouter, Depends
from app.schemas.user import LoginRequest, LoginResponse, UserResponse from app.schemas.user import (
from app.services.auth_service import login, get_current_user LoginRequest,
LoginResponse,
RefreshRequest,
RefreshResponse,
UserResponse,
)
from app.services.auth_service import login, refresh_access_token, get_current_user
router = APIRouter(prefix="/auth", tags=["认证"]) router = APIRouter(prefix="/auth", tags=["认证"])
@router.post("/login", response_model=LoginResponse) @router.post("/login", response_model=LoginResponse)
def login_endpoint(data: LoginRequest): def login_endpoint(data: LoginRequest):
"""登录 — 验证 MOM sys_user 表,返回 JWT""" """登录 — 验证 MOM sys_user 表,返回 Access + Refresh 双 Token"""
return login(data.username, data.password) return login(data.username, data.password)
@router.post("/refresh", response_model=RefreshResponse)
def refresh_endpoint(data: RefreshRequest):
"""刷新 Access Token — 使用 Refresh Token 换取新的 Access Token"""
return refresh_access_token(data.refresh_token)
@router.get("/me", response_model=UserResponse) @router.get("/me", response_model=UserResponse)
def get_me(current_user: dict = Depends(get_current_user)): def get_me(current_user: dict = Depends(get_current_user)):
"""获取当前用户信息(从 JWT 解析)""" """获取当前用户信息(从 Access Token 解析)"""
return UserResponse( return UserResponse(
id=current_user["sub"], id=current_user["sub"],
username=current_user.get("username", ""), username=current_user.get("username", ""),

View File

@ -0,0 +1,79 @@
"""通知 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 app.core.database import get_db
from app.models.notification import Notification
from app.schemas.notification import NotificationResponse, NotificationListResponse
router = APIRouter(prefix="/notifications", tags=["消息通知"])
@router.get("/", response_model=NotificationListResponse)
async def list_notifications(
user_id: str = Query(..., description="当前用户ID"),
skip: int = Query(0, ge=0),
limit: int = Query(20, ge=1, le=100),
db: AsyncSession = Depends(get_db),
):
"""获取当前用户的通知列表(按时间倒序)"""
# 总数
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()
return NotificationListResponse(
notifications=[
NotificationResponse.model_validate(n) for n in notifications
],
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),
):
"""标记单条通知为已读"""
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)

View File

@ -62,11 +62,15 @@ async def scan_product(
@router.get("/", response_model=list[ProductResponse]) @router.get("/", response_model=list[ProductResponse])
async def list_products( async def list_products(
skip: int = Query(0, ge=0, description="跳过条数"), skip: int = Query(0, ge=0, description="跳过条数"),
limit: int = Query(50, ge=1, le=200, description="返回条数"), limit: int = Query(50, ge=1, le=1000, description="返回条数"),
keyword: str | None = Query(None, description="多维搜索: 产品身份证/订单号/规格型号"),
status: str | None = Query(None, description="产品状态筛选: PENDING/WIP/COMPLETED/ARCHIVED"),
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
): ):
"""获取产品列表""" """获取产品列表 — 支持 keyword 搜索 + 状态筛选"""
return await product_service.get_all_products(db, skip=skip, limit=limit) return await product_service.get_all_products(
db, skip=skip, limit=limit, keyword=keyword, status_filter=status,
)
@router.get("/{product_id}", response_model=ProductResponse) @router.get("/{product_id}", response_model=ProductResponse)

View File

@ -10,6 +10,7 @@ from app.api.v1.endpoints.materials import router as materials_router
from app.api.v1.endpoints.users import router as users_router from app.api.v1.endpoints.users import router as users_router
from app.api.v1.endpoints.upload import router as upload_router from app.api.v1.endpoints.upload import router as upload_router
from app.api.v1.endpoints.records import router as records_router from app.api.v1.endpoints.records import router as records_router
from app.api.v1.endpoints.notifications import router as notifications_router
api_router = APIRouter() api_router = APIRouter()
@ -23,3 +24,4 @@ api_router.include_router(materials_router)
api_router.include_router(users_router) api_router.include_router(users_router)
api_router.include_router(upload_router) api_router.include_router(upload_router)
api_router.include_router(records_router) api_router.include_router(records_router)
api_router.include_router(notifications_router)

View File

@ -9,7 +9,8 @@ class Settings(BaseSettings):
# ---- JWT ---- # ---- JWT ----
SECRET_KEY: str = "change-me-in-production" SECRET_KEY: str = "change-me-in-production"
ACCESS_TOKEN_EXPIRE_MINUTES: int = 30 ACCESS_TOKEN_EXPIRE_MINUTES: int = 120 # Access Token: 2 小时
REFRESH_TOKEN_EXPIRE_DAYS: int = 7 # Refresh Token: 7 天
# ---- 调试 ---- # ---- 调试 ----
DEBUG: bool = True DEBUG: bool = True

View File

@ -1,4 +1,4 @@
"""安全模块 — JWT Token 生成与验证""" """安全模块 — JWT Token 生成与验证(双 Token 架构)"""
from datetime import timedelta from datetime import timedelta
from jose import JWTError, jwt from jose import JWTError, jwt
from passlib.context import CryptContext from passlib.context import CryptContext
@ -9,17 +9,36 @@ pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
ALGORITHM = "HS256" ALGORITHM = "HS256"
# Token 类型声明
TOKEN_TYPE_ACCESS = "access"
TOKEN_TYPE_REFRESH = "refresh"
def create_access_token(data: dict, expires_delta: timedelta | None = None) -> str: def create_access_token(data: dict, expires_delta: timedelta | None = None) -> str:
"""生成 JWT Access Token""" """生成 JWT Access Token(2 小时有效)"""
to_encode = data.copy() to_encode = data.copy()
expire = get_beijing_time() + ( expire = get_beijing_time() + (
expires_delta or timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) expires_delta or timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES)
) )
to_encode.update({"exp": expire}) to_encode.update({"exp": expire, "type": TOKEN_TYPE_ACCESS})
return jwt.encode(to_encode, settings.SECRET_KEY, algorithm=ALGORITHM) return jwt.encode(to_encode, settings.SECRET_KEY, algorithm=ALGORITHM)
def create_refresh_token(data: dict, expires_delta: timedelta | None = None) -> str:
"""生成 JWT Refresh Token(7 天有效,仅用于刷新 Access Token)"""
to_encode = data.copy()
expire = get_beijing_time() + (
expires_delta or timedelta(days=settings.REFRESH_TOKEN_EXPIRE_DAYS)
)
to_encode.update({"exp": expire, "type": TOKEN_TYPE_REFRESH})
return jwt.encode(to_encode, settings.SECRET_KEY, algorithm=ALGORITHM)
def decode_token(token: str) -> dict:
"""解码并验证 JWT Token,返回 payload"""
return jwt.decode(token, settings.SECRET_KEY, algorithms=[ALGORITHM])
def verify_password(plain_password: str, hashed_password: str) -> bool: def verify_password(plain_password: str, hashed_password: str) -> bool:
"""验证明文密码 vs 哈希密码""" """验证明文密码 vs 哈希密码"""
return pwd_context.verify(plain_password, hashed_password) return pwd_context.verify(plain_password, hashed_password)

View File

@ -4,6 +4,7 @@ from app.models.production_order import ProductionOrder
from app.models.product import Product from app.models.product import Product
from app.models.task import Task, TaskRecord from app.models.task import Task, TaskRecord
from app.models.task_log import TaskLog from app.models.task_log import TaskLog
from app.models.notification import Notification
__all__ = [ __all__ = [
"Base", "Base",
"ProductionOrder", "ProductionOrder",
@ -11,4 +12,5 @@ __all__ = [
"Task", "Task",
"TaskRecord", "TaskRecord",
"TaskLog", "TaskLog",
"Notification",
] ]

View File

@ -0,0 +1,52 @@
"""通知模型 — 任务转交/驳回等事件的消息提醒"""
import uuid
from datetime import datetime
from sqlalchemy import String, DateTime, Boolean, ForeignKey, Text
from sqlalchemy.dialects.postgresql import UUID
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base
from app.core.time_utils import get_beijing_time
# 通知类型常量
NOTIFY_TRANSFER = "TRANSFER" # 新任务派发/转交
NOTIFY_REJECT = "REJECT" # 品质驳回
class Notification(Base):
__tablename__ = "notifications"
id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True), primary_key=True, default=uuid.uuid4,
)
user_id: Mapped[str] = mapped_column(
String(64), nullable=False, index=True, comment="接收人ID(逻辑外键→老系统)",
)
title: Mapped[str] = mapped_column(
String(200), nullable=False, comment="通知标题",
)
content: Mapped[str] = mapped_column(
Text, nullable=False, comment="通知内容详情",
)
type: Mapped[str] = mapped_column(
String(20), nullable=False, comment="通知类型: TRANSFER(转交派发) | REJECT(驳回)",
)
task_id: Mapped[uuid.UUID | None] = mapped_column(
UUID(as_uuid=True), ForeignKey("tasks.id", ondelete="SET NULL"), nullable=True, comment="关联任务ID",
)
is_read: Mapped[bool] = mapped_column(
Boolean, default=False, comment="是否已读",
)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), default=get_beijing_time, comment="创建时间",
)
def __repr__(self) -> str:
return f"<Notification {self.type} → {self.user_id}>"

View File

@ -84,6 +84,7 @@ class Task(Base):
) )
child_tasks: Mapped[list["Task"]] = relationship( child_tasks: Mapped[list["Task"]] = relationship(
"Task", back_populates="parent_task", lazy="selectin", "Task", back_populates="parent_task", lazy="selectin",
order_by="Task.created_at",
) )
records: Mapped[list["TaskRecord"]] = relationship( records: Mapped[list["TaskRecord"]] = relationship(
"TaskRecord", back_populates="task", lazy="selectin", cascade="all, delete-orphan", "TaskRecord", back_populates="task", lazy="selectin", cascade="all, delete-orphan",

View File

@ -0,0 +1,26 @@
"""通知 Pydantic Schema"""
from __future__ import annotations
import uuid
from datetime import datetime
from pydantic import BaseModel, Field
class NotificationResponse(BaseModel):
"""通知列表响应"""
id: uuid.UUID
user_id: str
title: str
content: str
type: str
task_id: uuid.UUID | None = None
is_read: bool
created_at: datetime
model_config = {"from_attributes": True}
class NotificationListResponse(BaseModel):
"""通知分页列表"""
notifications: list[NotificationResponse]
total: int
unread_count: int

View File

@ -110,6 +110,7 @@ class TaskSummaryResponse(BaseModel):
status: str status: str
notify_parent_on_complete: bool notify_parent_on_complete: bool
is_rework: bool = False is_rework: bool = False
task_type: str | None = None
remark: str | None = None remark: str | None = None
reject_reason: str | None = None reject_reason: str | None = None
received_at: datetime | None = None received_at: datetime | None = None

View File

@ -1,4 +1,4 @@
"""用户 Schemas — 对接 MOM sys_user 表""" """用户 Schemas — 对接 MOM sys_user 表 + 双 Token"""
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
@ -20,5 +20,15 @@ class UserResponse(BaseModel):
class LoginResponse(BaseModel): class LoginResponse(BaseModel):
access_token: str access_token: str
refresh_token: str
token_type: str = "bearer" token_type: str = "bearer"
user: UserResponse user: UserResponse
class RefreshRequest(BaseModel):
refresh_token: str = Field(..., description="Refresh Token")
class RefreshResponse(BaseModel):
access_token: str
token_type: str = "bearer"

View File

@ -1,11 +1,18 @@
"""认证服务 — 对接 MOM 系统 sys_user 表 + Track 自有 JWT""" """认证服务 — 对接 MOM 系统 sys_user 表 + Track 自有 JWT(双 Token 架构)"""
from fastapi import HTTPException, status, Depends from fastapi import HTTPException, status, Depends
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
from jose import JWTError, jwt from jose import JWTError, jwt
from werkzeug.security import check_password_hash from werkzeug.security import check_password_hash
from app.core.config import settings from app.core.config import settings
from app.core.security import create_access_token, ALGORITHM from app.core.security import (
create_access_token,
create_refresh_token,
decode_token,
ALGORITHM,
TOKEN_TYPE_ACCESS,
TOKEN_TYPE_REFRESH,
)
from app.core.mom_database import MomSessionLocal from app.core.mom_database import MomSessionLocal
from app.schemas.user import LoginResponse, UserResponse from app.schemas.user import LoginResponse, UserResponse
@ -13,15 +20,15 @@ security = HTTPBearer()
def login(username: str, password: str) -> LoginResponse: def login(username: str, password: str) -> LoginResponse:
"""登录 — 查询 MOM 数据库 sys_user 表验证""" """登录 — 签发双 Token(Access + Refresh)"""
db = MomSessionLocal() db = MomSessionLocal()
try: try:
# 1. 超级管理员硬编码(和 MOM 系统一致) # 1. 超级管理员硬编码(和 MOM 系统一致)
if username == "IRIS" and password == "123321": if username == "IRIS" and password == "123321":
token_data = {"sub": "0", "role": "SUPER_ADMIN", "username": "IRIS", "display_name": "超级管理员"}
return LoginResponse( return LoginResponse(
access_token=create_access_token( access_token=create_access_token(data=token_data),
data={"sub": "0", "role": "SUPER_ADMIN"} refresh_token=create_refresh_token(data=token_data),
),
user=UserResponse( user=UserResponse(
id="0", id="0",
username="IRIS", username="IRIS",
@ -60,17 +67,16 @@ def login(username: str, password: str) -> LoginResponse:
# 4. 解析 display_name("张三/zhangsan01" → "张三") # 4. 解析 display_name("张三/zhangsan01" → "张三")
display_name = full_username.split("/")[0] if "/" in full_username else full_username display_name = full_username.split("/")[0] if "/" in full_username else full_username
token = create_access_token( token_data = {
data={ "sub": str(user_id),
"sub": str(user_id), "role": role or "operator",
"role": role or "operator", "username": username,
"username": username, "display_name": display_name,
"display_name": display_name, }
}
)
return LoginResponse( return LoginResponse(
access_token=token, access_token=create_access_token(data=token_data),
refresh_token=create_refresh_token(data=token_data),
user=UserResponse( user=UserResponse(
id=str(user_id), id=str(user_id),
username=username, username=username,
@ -83,16 +89,60 @@ def login(username: str, password: str) -> LoginResponse:
db.close() db.close()
def refresh_access_token(refresh_token: str) -> dict:
"""
使用 Refresh Token 换取新的 Access Token。
校验:
1. Token 签名是否有效
2. Token type 是否为 "refresh"
3. Token 是否未过期
"""
try:
payload = decode_token(refresh_token)
except JWTError:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Refresh Token 无效或已过期,请重新登录",
)
# 校验 token 类型
if payload.get("type") != TOKEN_TYPE_REFRESH:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="无效的 Token 类型,仅接受 Refresh Token",
)
# 提取用户信息,签发新的 Access Token
access_token = create_access_token(
data={
"sub": payload.get("sub"),
"role": payload.get("role", "operator"),
"username": payload.get("username", ""),
"display_name": payload.get("display_name", ""),
}
)
return {"access_token": access_token, "token_type": "bearer"}
async def get_current_user( async def get_current_user(
credentials: HTTPAuthorizationCredentials = Depends(security), credentials: HTTPAuthorizationCredentials = Depends(security),
) -> dict: ) -> dict:
"""从 Bearer Token 解析当前用户(不查数据库,直接解 JWT)""" """从 Bearer Token 解析当前用户(仅接受 Access Token)"""
token = credentials.credentials token = credentials.credentials
try: try:
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[ALGORITHM]) payload = decode_token(token)
user_id = payload.get("sub") user_id = payload.get("sub")
if not user_id: if not user_id:
raise HTTPException(status_code=401, detail="无效的 Token") raise HTTPException(status_code=401, detail="无效的 Token")
# 校验:仅接受 access token
if payload.get("type") == TOKEN_TYPE_REFRESH:
raise HTTPException(
status_code=401,
detail="请使用 Access Token 访问 API,Refresh Token 仅用于刷新",
)
return payload return payload
except JWTError: except JWTError:
raise HTTPException(status_code=401, detail="无效的 Token") raise HTTPException(status_code=401, detail="无效的 Token")

View File

@ -2,7 +2,7 @@
from __future__ import annotations from __future__ import annotations
import uuid import uuid
from fastapi import HTTPException, status from fastapi import HTTPException, status
from sqlalchemy import select from sqlalchemy import select, or_, cast, String
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload from sqlalchemy.orm import selectinload
@ -34,6 +34,7 @@ def _task_to_response(task: Task) -> TaskResponse:
status=task.status, status=task.status,
notify_parent_on_complete=task.notify_parent_on_complete, notify_parent_on_complete=task.notify_parent_on_complete,
is_rework=task.is_rework, is_rework=task.is_rework,
task_type=task.task_type,
remark=task.remark, remark=task.remark,
reject_reason=task.reject_reason, reject_reason=task.reject_reason,
received_at=task.received_at, received_at=task.received_at,
@ -277,15 +278,68 @@ async def update_overall_status(db: AsyncSession, serial_number: str, status_val
return await get_product_by_serial(db, serial_number) return await get_product_by_serial(db, serial_number)
async def get_all_products(db: AsyncSession, skip: int = 0, limit: int = 50) -> list[ProductResponse]: async def get_all_products(
"""获取产品列表""" db: AsyncSession,
result = await db.execute( skip: int = 0,
select(Product) limit: int = 50,
.options(selectinload(Product.order)) keyword: str | None = None,
.offset(skip) status_filter: str | None = None,
.limit(limit) ) -> list[ProductResponse]:
.order_by(Product.created_at.desc()) """
) 获取产品列表 — 支持多维 keyword 搜索 + 状态筛选
keyword: 同时模糊匹配 serial_number (产品身份证)、material_name/id (规格型号)、order_no (订单号)
status_filter: 按产品状态过滤 (如 PENDING / WIP / COMPLETED / ARCHIVED)
"""
stmt = select(Product).options(selectinload(Product.order))
# keyword 多字段 OR 模糊搜索
if keyword and keyword.strip():
kw = f"%{keyword.strip()}%"
stmt = stmt.outerjoin(ProductionOrder, Product.order_id == ProductionOrder.id).where(
or_(
Product.serial_number.ilike(kw),
Product.material_name.ilike(kw),
cast(Product.material_id, String).ilike(kw),
Product.spec_model.ilike(kw),
ProductionOrder.order_no.ilike(kw),
)
).distinct()
# 状态筛选 — 大小写不敏感,支持组合过滤
if status_filter and status_filter.strip():
from sqlalchemy import func
sf = status_filter.strip().upper()
if sf == "DONE":
# "已完成" 匹配 COMPLETED 或 ARCHIVED
stmt = stmt.where(
or_(
func.upper(Product.status) == "COMPLETED",
func.upper(Product.status) == "ARCHIVED",
)
)
elif sf == "PENDING":
# "待流转" — 产品状态 PENDING 且所有顶层任务均未分配人
stmt = (
stmt.outerjoin(Task, Task.product_id == Product.id)
.where(func.upper(Product.status) == "PENDING")
.where(Task.assignee_id.is_(None))
.distinct()
)
elif sf == "PENDING_ASSIGNED":
# "待接收" — 产品状态 PENDING 但已有任务被分配(等待工人扫码)
stmt = (
stmt.outerjoin(Task, Task.product_id == Product.id)
.where(func.upper(Product.status) == "PENDING")
.where(Task.assignee_id.isnot(None))
.distinct()
)
else:
stmt = stmt.where(func.upper(Product.status) == sf)
stmt = stmt.offset(skip).limit(limit).order_by(Product.created_at.desc())
result = await db.execute(stmt)
products = result.scalars().all() products = result.scalars().all()
return [ return [
ProductResponse( ProductResponse(

View File

@ -7,6 +7,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload from sqlalchemy.orm import selectinload
from app.models.task import Task, TaskRecord, TASK_STATUS_PENDING, TASK_STATUS_WIP, TASK_STATUS_COMPLETED, TASK_STATUS_REJECTED, TASK_STATUS_CANCELED, TASK_STATUS_ARCHIVED from app.models.task import Task, TaskRecord, TASK_STATUS_PENDING, TASK_STATUS_WIP, TASK_STATUS_COMPLETED, TASK_STATUS_REJECTED, TASK_STATUS_CANCELED, TASK_STATUS_ARCHIVED
from app.models.notification import Notification, NOTIFY_TRANSFER, NOTIFY_REJECT
from app.core.time_utils import get_beijing_time from app.core.time_utils import get_beijing_time
from app.models.product import Product from app.models.product import Product
from app.models.task_log import TaskLog from app.models.task_log import TaskLog
@ -288,6 +289,12 @@ async def end_task(
""" """
task = await _get_task_or_404(db, task_id) task = await _get_task_or_404(db, task_id)
# 校验:仅 SPAWN 协助分支可以结束,主分支(TRANSFER/RECOVERY)不能通过此接口终止
if not task.parent_task_id:
raise HTTPException(status_code=409, detail="根任务无法结束,请使用完工转交")
if task.task_type != "SPAWN":
raise HTTPException(status_code=409, detail="仅协助分支可以结束,主分支请使用完工转交")
# 校验:必须等待所有协助分支完成 # 校验:必须等待所有协助分支完成
await _check_children_done(db, task_id) await _check_children_done(db, task_id)
@ -562,7 +569,7 @@ async def reject_task(
task_name=task.task_name, task_name=task.task_name,
assignee_id=rework_assignee_id, assignee_id=rework_assignee_id,
status=TASK_STATUS_PENDING, status=TASK_STATUS_PENDING,
task_type="TRANSFER", task_type=task.task_type, # 🚀 继承被驳回任务的基因:主线→主线,协助→协助
notify_parent_on_complete=task.notify_parent_on_complete, notify_parent_on_complete=task.notify_parent_on_complete,
is_rework=True, is_rework=True,
) )
@ -576,6 +583,24 @@ async def reject_task(
remark=f"返工任务(驳回自「{task.task_name}」,原因: {request.reason}),分配给 {rework_assignee_id}", remark=f"返工任务(驳回自「{task.task_name}」,原因: {request.reason}),分配给 {rework_assignee_id}",
) )
# 🔔 通知:品质驳回
product_sn = ""
try:
product_result = await db.execute(select(Product).where(Product.id == task.product_id))
p = product_result.scalar_one_or_none()
if p:
product_sn = p.serial_number or ""
except Exception:
pass
if rework_assignee_id:
db.add(Notification(
user_id=rework_assignee_id,
title="🔴 品质驳回提醒",
content=f"产品 [{product_sn}] 的「{task.task_name}」被驳回,原因: {request.reason}",
type=NOTIFY_REJECT,
task_id=rework_task.id,
))
await db.commit() await db.commit()
await db.refresh(task) await db.refresh(task)
@ -697,6 +722,21 @@ async def transfer_task(
operator_id=operator_id, operator_id=operator_id,
remark=request.note or f"由任务「{task.task_name}」裂变转交创建,分配给 {nt.assignee_id}", remark=request.note or f"由任务「{task.task_name}」裂变转交创建,分配给 {nt.assignee_id}",
) )
# 🔔 通知:新任务派发
if nt.assignee_id:
product_sn = ""
try:
if product:
product_sn = product.serial_number or ""
except Exception:
pass
db.add(Notification(
user_id=nt.assignee_id,
title=f"🟢 新任务派发",
content=f"产品 [{product_sn}] 的「{nt.task_name}」任务已分配给你",
type=NOTIFY_TRANSFER,
task_id=nt.id,
))
# --- 更新 Product 的 current_location_id --- # --- 更新 Product 的 current_location_id ---
product_result = await db.execute( product_result = await db.execute(