baodan/api/insurance/generation/routes.py
wsb1224 2b54d078f2 针对“A 账号生成、B 账号也能看到”的问题,目前不会再发生,前提是 A、B 都通过保险前端各自重新登录。
当前已验证:
用户 1 和用户 3 的 PPT、海报、任务列表完全隔离。
未登录或访客身份会直接返回 401。
任务详情、下载、工作区操作都校验所属用户。
跨用户幂等任务复用漏洞已封堵。
35 项相关测试、前端构建和部署健康检查均通过。
2026-08-01 19:38:56 +08:00

490 lines
18 KiB
Python

"""工作区与任务 API 路由。
按 docs/0728修复文件.md 十、API 设计实施。
"""
import json
import logging
from flask import Blueprint, request
from insurance.middleware.auth_middleware import account_required
from insurance.utils.response import success, error, ErrorCode
logger = logging.getLogger(__name__)
workspace_bp = Blueprint("workspace", __name__)
# ─── PPT 工作区 ────────────────────────────────────────────
@workspace_bp.route("/ppt/workspaces", methods=["GET"])
@account_required
def list_ppt_workspaces():
"""查询用户 PPT 工作区列表。"""
from insurance.db.compat import db
from insurance.models.ppt_session import PptSession
user_id = str(getattr(request, "user_id", "guest"))
status = request.args.get("status", "active")
query = PptSession.query.filter_by(user_id=user_id)
if status == "active":
query = query.filter(PptSession.archived_at.is_(None))
elif status == "archived":
query = query.filter(PptSession.archived_at.isnot(None))
query = query.order_by(PptSession.updated_at.desc())
sessions = query.limit(50).all()
return success({
"items": [_ppt_workspace_summary(s) for s in sessions],
})
@workspace_bp.route("/ppt/workspaces/<session_id>", methods=["GET"])
@account_required
def get_ppt_workspace(session_id):
"""获取 PPT 工作区详情。"""
from insurance.models.ppt_session import PptSession
user_id = str(getattr(request, "user_id", "guest"))
session = PptSession.query.filter_by(id=session_id, user_id=user_id).first()
if not session:
return error(ErrorCode.NOT_FOUND, "工作区不存在")
return success(session.to_dict())
@workspace_bp.route("/ppt/workspaces/<session_id>/rename", methods=["PUT"])
@account_required
def rename_ppt_workspace(session_id):
"""重命名 PPT 工作区。"""
from insurance.db.compat import db
from insurance.models.ppt_session import PptSession
user_id = str(getattr(request, "user_id", "guest"))
session = PptSession.query.filter_by(id=session_id, user_id=user_id).first()
if not session:
return error(ErrorCode.NOT_FOUND, "工作区不存在")
data = request.get_json(silent=True) or {}
title = data.get("title", "").strip()
if not title:
return error(ErrorCode.PARAM_ERROR, "名称不能为空")
session.title = title[:200]
db.session.commit()
return success({"id": session.id, "title": session.title})
@workspace_bp.route("/ppt/workspaces/<session_id>/draft", methods=["POST", "PATCH"])
@account_required
def save_ppt_draft(session_id):
"""自动保存 PPT 草稿(乐观锁)。
请求体:
draft_options: 草稿选项(模板、公司、脱敏等)
expected_revision: 期望的草稿版本(乐观锁,可选)
"""
from insurance.db.compat import db
from insurance.models.ppt_session import PptSession
user_id = str(getattr(request, "user_id", "guest"))
session = PptSession.query.filter_by(id=session_id, user_id=user_id).first()
if not session:
return error(ErrorCode.NOT_FOUND, "工作区不存在")
data = request.get_json(silent=True) or {}
# 乐观锁检查
expected = data.get("expected_revision")
if expected is not None and expected != session.draft_revision:
return error(ErrorCode.PARAM_ERROR, f"草稿版本冲突:当前 v{session.draft_revision},请求 v{expected}")
if "draft_options" in data:
session.draft_options_json = json.dumps(data["draft_options"], ensure_ascii=False)
if "workflow_step" in data:
session.workflow_step = data["workflow_step"]
session.draft_revision = (session.draft_revision or 1) + 1
db.session.commit()
return success({
"id": session.id,
"draft_revision": session.draft_revision,
"generated_revision": session.generated_revision or 0,
"has_unsaved_changes": session.draft_revision > (session.generated_revision or 0),
})
@workspace_bp.route("/ppt/workspaces/<session_id>/archive", methods=["PUT"])
@account_required
def archive_ppt_workspace(session_id):
"""归档 PPT 工作区。"""
from datetime import datetime
from insurance.db.compat import db
from insurance.models.ppt_session import PptSession
user_id = str(getattr(request, "user_id", "guest"))
session = PptSession.query.filter_by(id=session_id, user_id=user_id).first()
if not session:
return error(ErrorCode.NOT_FOUND, "工作区不存在")
session.archived_at = datetime.now()
db.session.commit()
return success({"id": session.id, "archived_at": str(session.archived_at)})
@workspace_bp.route("/ppt/workspaces/<session_id>/unarchive", methods=["PUT"])
@account_required
def unarchive_ppt_workspace(session_id):
"""取消归档 PPT 工作区。"""
from insurance.db.compat import db
from insurance.models.ppt_session import PptSession
user_id = str(getattr(request, "user_id", "guest"))
session = PptSession.query.filter_by(id=session_id, user_id=user_id).first()
if not session:
return error(ErrorCode.NOT_FOUND, "工作区不存在")
session.archived_at = None
db.session.commit()
return success({"id": session.id})
@workspace_bp.route("/ppt/workspaces/<session_id>/copy", methods=["POST"])
@account_required
def copy_ppt_workspace(session_id):
"""复制 PPT 工作区为新任务。"""
from insurance.db.compat import db
from insurance.models.ppt_session import PptSession
import uuid
user_id = str(getattr(request, "user_id", "guest"))
source = PptSession.query.filter_by(id=session_id, user_id=user_id).first()
if not source:
return error(ErrorCode.NOT_FOUND, "工作区不存在")
new_id = uuid.uuid4().hex
new_session = PptSession(
id=new_id,
user_id=user_id,
status="created",
files_json=source.files_json,
extractions_json=source.extractions_json,
title=f"{source.title or 'PPT'} (副本)",
workflow_step=source.workflow_step,
draft_options_json=source.draft_options_json,
)
db.session.add(new_session)
db.session.commit()
return success({"id": new_id, "title": new_session.title})
# ─── 海报工作区 ────────────────────────────────────────────
@workspace_bp.route("/poster/workspaces", methods=["GET"])
@account_required
def list_poster_workspaces():
"""查询用户海报工作区列表。"""
from insurance.models.poster_record import PosterRecord
user_id = str(getattr(request, "user_id", "guest"))
status = request.args.get("status", "active")
query = PosterRecord.query.filter_by(user_id=user_id)
if status == "active":
query = query.filter(PosterRecord.archived_at.is_(None))
elif status == "archived":
query = query.filter(PosterRecord.archived_at.isnot(None))
query = query.order_by(PosterRecord.created_at.desc())
records = query.limit(50).all()
return success({
"items": [_poster_workspace_summary(r) for r in records],
})
@workspace_bp.route("/poster/workspaces/<int:record_id>", methods=["GET"])
@account_required
def get_poster_workspace(record_id):
"""获取海报工作区详情。"""
from insurance.models.poster_record import PosterRecord
user_id = str(getattr(request, "user_id", "guest"))
record = PosterRecord.query.filter_by(id=record_id, user_id=user_id).first()
if not record:
return error(ErrorCode.NOT_FOUND, "工作区不存在")
return success(record.to_dict())
@workspace_bp.route("/poster/workspaces/<int:record_id>/rename", methods=["PUT"])
@account_required
def rename_poster_workspace(record_id):
"""重命名海报工作区。"""
from insurance.db.compat import db
from insurance.models.poster_record import PosterRecord
user_id = str(getattr(request, "user_id", "guest"))
record = PosterRecord.query.filter_by(id=record_id, user_id=user_id).first()
if not record:
return error(ErrorCode.NOT_FOUND, "工作区不存在")
data = request.get_json(silent=True) or {}
title = data.get("title", "").strip()
if not title:
return error(ErrorCode.PARAM_ERROR, "名称不能为空")
record.title = title[:200]
db.session.commit()
return success({"id": record.id, "title": record.title})
@workspace_bp.route("/poster/workspaces/<int:record_id>/draft", methods=["POST", "PATCH"])
@account_required
def save_poster_draft(record_id):
"""自动保存海报草稿(乐观锁)。"""
from insurance.db.compat import db
from insurance.models.poster_record import PosterRecord
user_id = str(getattr(request, "user_id", "guest"))
record = PosterRecord.query.filter_by(id=record_id, user_id=user_id).first()
if not record:
return error(ErrorCode.NOT_FOUND, "工作区不存在")
data = request.get_json(silent=True) or {}
expected = data.get("expected_revision")
if expected is not None and expected != record.draft_revision:
return error(ErrorCode.PARAM_ERROR, f"草稿版本冲突:当前 v{record.draft_revision},请求 v{expected}")
if "workflow_step" in data:
record.workflow_step = data["workflow_step"]
if "copy_content" in data:
import json
record.copy_content = json.dumps(data["copy_content"], ensure_ascii=False)
record.draft_revision = (record.draft_revision or 1) + 1
db.session.commit()
return success({
"id": record.id,
"draft_revision": record.draft_revision,
"generated_revision": record.generated_revision or 0,
"has_unsaved_changes": record.draft_revision > (record.generated_revision or 0),
})
@workspace_bp.route("/poster/workspaces/<int:record_id>/archive", methods=["PUT"])
@account_required
def archive_poster_workspace(record_id):
"""归档海报工作区。"""
from datetime import datetime
from insurance.db.compat import db
from insurance.models.poster_record import PosterRecord
user_id = str(getattr(request, "user_id", "guest"))
record = PosterRecord.query.filter_by(id=record_id, user_id=user_id).first()
if not record:
return error(ErrorCode.NOT_FOUND, "工作区不存在")
record.archived_at = datetime.now()
record.draft_status = "archived"
db.session.commit()
return success({"id": record.id, "archived_at": str(record.archived_at)})
@workspace_bp.route("/poster/workspaces/<int:record_id>/copy", methods=["POST"])
@account_required
def copy_poster_workspace(record_id):
"""复制海报工作区为新任务。"""
from insurance.db.compat import db
from insurance.models.poster_record import PosterRecord
user_id = str(getattr(request, "user_id", "guest"))
source = PosterRecord.query.filter_by(id=record_id, user_id=user_id).first()
if not source:
return error(ErrorCode.NOT_FOUND, "工作区不存在")
new_record = PosterRecord(
user_id=user_id,
product_id=source.product_id,
case_upload_id=source.case_upload_id,
template_id=source.template_id,
copy_mode=source.copy_mode,
copy_content=source.copy_content,
export_size=source.export_size,
title=f"{source.title or f'海报 #{source.id}'} (副本)",
workflow_step="template",
extra_data=source.extra_data,
)
db.session.add(new_record)
db.session.commit()
return success({"id": new_record.id, "title": new_record.title})
# ─── 统一任务接口 ──────────────────────────────────────────
@workspace_bp.route("/tasks", methods=["GET"])
@account_required
def list_tasks():
"""查询任务列表(任务中心)。"""
from insurance.generation import task_service
user_id = str(getattr(request, "user_id", "guest"))
artifact_type = request.args.get("artifact_type", "")
status = request.args.get("status", "")
page = max(1, request.args.get("page", 1, type=int))
page_size = min(100, max(1, request.args.get("page_size", 20, type=int)))
result = task_service.list_tasks(
user_id=user_id,
artifact_type=artifact_type or None,
status=status or None,
page=page,
page_size=page_size,
)
return success(result)
@workspace_bp.route("/tasks/active", methods=["GET"])
@account_required
def list_active_tasks():
"""查询活跃任务(任务坞)。"""
from insurance.generation import task_service
user_id = str(getattr(request, "user_id", "guest"))
artifact_type = request.args.get("artifact_type", "")
items = task_service.list_active_tasks(user_id, artifact_type or None)
return success({"items": items})
@workspace_bp.route("/tasks/<task_id>", methods=["GET"])
@account_required
def get_task(task_id):
"""获取任务详情。"""
from insurance.generation import task_service
user_id = str(getattr(request, "user_id", "guest"))
result = task_service.get_task(task_id, user_id)
if result.get("code") != 0:
return error(ErrorCode.NOT_FOUND, result.get("message", "任务不存在"))
return success(result["data"])
@workspace_bp.route("/tasks/<task_id>/cancel", methods=["PUT"])
@account_required
def cancel_task(task_id):
"""取消排队中的任务。"""
from insurance.generation import task_service
user_id = str(getattr(request, "user_id", "guest"))
result = task_service.cancel_task(task_id, user_id)
if result.get("code") != 0:
return error(ErrorCode.PARAM_ERROR, result.get("message", "操作失败"))
return success(result["data"])
@workspace_bp.route("/tasks/<task_id>/hide", methods=["PUT"])
@account_required
def hide_task(task_id):
"""从任务坞隐藏任务。"""
from insurance.generation import task_service
user_id = str(getattr(request, "user_id", "guest"))
result = task_service.hide_task_from_dock(task_id, user_id)
if result.get("code") != 0:
return error(ErrorCode.NOT_FOUND, result.get("message", "操作失败"))
return success(result["data"])
@workspace_bp.route("/tasks/<task_id>/viewed", methods=["PUT"])
@account_required
def mark_task_viewed(task_id):
"""标记任务为已查看。"""
from insurance.generation import task_service
user_id = str(getattr(request, "user_id", "guest"))
result = task_service.mark_task_viewed(task_id, user_id)
if result.get("code") != 0:
return error(ErrorCode.NOT_FOUND, result.get("message", "操作失败"))
return success(result["data"])
@workspace_bp.route("/tasks/<task_id>/download", methods=["GET"])
@account_required
def download_task_output(task_id):
"""按任务下载成品(确保下载的是该任务的版本,而非工作区最新版本)。"""
import json as _json
import os
from flask import send_file
from insurance.models.generation_task import GenerationTask
user_id = str(getattr(request, "user_id", "guest"))
task = GenerationTask.query.get(task_id)
if not task or task.user_id != user_id:
return error(ErrorCode.NOT_FOUND, "任务不存在")
if task.status != "done":
return error(ErrorCode.PARAM_ERROR, "任务未完成,无法下载")
output = _json.loads(task.output_json) if task.output_json else {}
file_path = output.get("filePath", "")
if not file_path or not os.path.exists(file_path):
return error(ErrorCode.NOT_FOUND, "文件不存在或已过期")
# 根据文件扩展名确定 MIME 类型
if file_path.endswith(".pptx"):
mimetype = "application/vnd.openxmlformats-officedocument.presentationml.presentation"
download_name = f"ppt_{task.workspace_id}_{task_id[:8]}.pptx"
elif file_path.endswith(".png"):
mimetype = "image/png"
download_name = f"poster_{task.workspace_id}_{task_id[:8]}.png"
else:
mimetype = "application/octet-stream"
download_name = os.path.basename(file_path)
return send_file(
file_path,
mimetype=mimetype,
as_attachment=True,
download_name=download_name,
)
# ─── 辅助函数 ──────────────────────────────────────────────
def _ppt_workspace_summary(session):
"""PPT 工作区摘要(列表用)。"""
import json
extractions = json.loads(session.extractions_json) if session.extractions_json else []
product_names = [e.get("productName", "") for e in extractions if e.get("productName")]
return {
"id": session.id,
"title": session.title or session.id[:8],
"status": session.status,
"workflowStep": session.workflow_step or "upload",
"draftRevision": session.draft_revision or 1,
"generatedRevision": session.generated_revision or 0,
"fileCount": len(json.loads(session.files_json) if session.files_json else []),
"productNames": product_names,
"createdAt": str(session.created_at) if session.created_at else None,
"updatedAt": str(session.updated_at) if session.updated_at else None,
}
def _poster_workspace_summary(record):
"""海报工作区摘要(列表用)。"""
return {
"id": record.id,
"title": record.title or f"海报 #{record.id}",
"taskStatus": record.task_status,
"workflowStep": record.workflow_step or "product",
"draftRevision": record.draft_revision or 1,
"generatedRevision": record.generated_revision or 0,
"productId": record.product_id,
"createdAt": record.created_at.isoformat() if record.created_at else None,
}