baodan/api/insurance/generation/task_service.py
wsb1224 5f78598b3a input_revision 正确记录
task_service.py — create_task() 新增 input_revision 参数,存入 GenerationTask.input_revision
generation_task.py — 模型的静态 create_task 方法同步新增参数
ppt/routes.py — 三处调用(parse/generate/regenerate)均传入 session.draft_revision
poster/service.py — 海报生成传入 record.draft_revision
现在每个任务快照都记录了创建时的草稿版本号,不再一直是默认值 1
2. 海报任务进度实时同步 
celery_tasks.py — 在海报生成的每个阶段(preparing_data→20%、building_prompt→40%、requesting_image→60%、saving→85%)调用 sync_poster_progress() 同步到 PosterRecord.task_status/task_progress
任务领取时立即将 PosterRecord.task_status 设为 running
解决了"页面轮询海报记录时一直显示排队中"的问题
3. 后端测试覆盖 
test_task_state_sync.py — 新增 13 个测试用例:
sync_workspace_status 同步 cancelled/failed 到 PPT session(4 个)
sync_workspace_status 同步 cancelled/failed 到 PosterRecord(2 个)
旧任务跳过同步的安全检查(1 个)
cancel_task 状态更新和非 queued 拒绝(2 个)
mark_task_viewed 正确标记和权限校验(2 个)
input_revision 正确记录和默认值(2 个)
list_active_tasks 排除终态任务(1 个)
所有 13 个新测试 + 原有 10 个测试全部通过
4. useAutoSave TypeScript 类型修复 
useAutoSave.ts — getEndpoint 变量显式标注 () => string 类型,消除 TS2345 编译错误
2026-07-30 13:54:44 +08:00

361 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""统一任务服务:创建、查询、取消任务。"""
import json
import logging
import uuid
from datetime import datetime
logger = logging.getLogger(__name__)
def create_task(user_id: str, artifact_type: str, operation: str,
workspace_id: str, title: str = "",
input_snapshot: dict = None, idempotency_key: str = None,
input_revision: int = None) -> dict:
"""创建生成任务并提交到 Celery。
返回 {"code": 0, "data": task_dict} 或错误。
"""
from insurance.db.compat import db
from insurance.models.generation_task import GenerationTask
# 幂等检查
if idempotency_key:
existing = GenerationTask.query.filter_by(
idempotency_key=idempotency_key,
).filter(GenerationTask.status.in_(["queued", "running", "done"])).first()
if existing:
return {"code": 0, "data": existing.to_dict(), "message": "任务已存在"}
# 检查同一工作区是否有运行中的任务
active = GenerationTask.query.filter_by(
workspace_id=workspace_id,
).filter(GenerationTask.status.in_(["queued", "running"])).first()
if active:
return {"code": 1001, "message": "该工作区有正在执行的任务,请等待完成", "data": None}
task = GenerationTask(
id=uuid.uuid4().hex,
user_id=user_id,
artifact_type=artifact_type,
operation=operation,
workspace_id=workspace_id,
title_snapshot=title,
input_snapshot_json=json.dumps(input_snapshot, ensure_ascii=False) if input_snapshot else None,
input_revision=input_revision or 1,
idempotency_key=idempotency_key,
)
db.session.add(task)
db.session.commit()
# 提交到 Celery。投递失败时必须结束数据库任务不能留下永久 queued。
try:
_dispatch_to_celery(task)
except Exception as exc:
message = f"任务提交失败: {exc}"
logger.error("任务 %s 提交到 Celery 失败: %s", task.id, exc, exc_info=True)
task.status = "failed"
task.error_code = "dispatch_failed"
task.error_message = message[:1000]
task.finished_at = datetime.now()
db.session.commit()
_sync_failed_ppt_session(task, message)
return {"code": 9999, "message": message, "data": task.to_dict()}
return {"code": 0, "data": task.to_dict()}
def _dispatch_to_celery(task):
"""根据任务类型分发到对应的 Celery 任务。"""
from insurance.generation.celery_tasks import (
parse_ppt_task, generate_ppt_task, regenerate_ppt_task,
generate_poster_task,
)
task_map = {
("ppt", "parse"): parse_ppt_task,
("ppt", "generate"): generate_ppt_task,
("ppt", "regenerate"): regenerate_ppt_task,
("poster", "generate"): generate_poster_task,
}
celery_task_fn = task_map.get((task.artifact_type, task.operation))
if not celery_task_fn:
logger.error(f"未知任务类型: {task.artifact_type}/{task.operation}")
task.status = "failed"
task.error_code = "unknown_type"
task.error_message = f"未知任务类型: {task.artifact_type}/{task.operation}"
from insurance.db.compat import db
db.session.commit()
return
result = celery_task_fn.apply_async(args=[task.id], queue="insurance")
task.celery_task_id = result.id
from insurance.db.compat import db
db.session.commit()
logger.info(f"任务 {task.id} 已提交到 Celery: {result.id}")
def _sync_failed_ppt_session(task, message: str):
"""将 PPT 解析任务失败同步到工作区,避免页面永久停在 parsing。"""
if task.artifact_type != "ppt" or task.operation != "parse":
return
from insurance.db.compat import db
from insurance.models.ppt_session import PptSession
session = PptSession.query.get(task.workspace_id)
if not session or session.status != "parsing":
return
if session.latest_task_id and session.latest_task_id != task.id:
return
session.status = "error"
session.parse_progress = 100
session.parse_message = "处理失败"
session.parse_error = message[:1000]
session.parse_finished_at = datetime.now()
db.session.commit()
# ─── 工作区状态同步 ────────────────────────────────────────
def sync_workspace_status(task):
"""将任务终态同步回工作区,保证页面恢复时状态一致。
调用时机:任务进入 done/failed/cancelled 状态时。
安全检查:只有当 task 是工作区的最新任务时才同步,防止旧任务覆盖新状态。
"""
from insurance.db.compat import db
if task.artifact_type == "ppt":
_sync_ppt_workspace(task, db)
elif task.artifact_type == "poster":
_sync_poster_workspace(task, db)
def _sync_ppt_workspace(task, db):
"""同步 PPT 工作区状态。"""
from insurance.models.ppt_session import PptSession
session = PptSession.query.get(task.workspace_id)
if not session:
return
# 安全检查:只同步最新任务
if session.latest_task_id and session.latest_task_id != task.id:
logger.debug("跳过旧任务同步: task=%s, latest_task=%s", task.id, session.latest_task_id)
return
if task.operation == "parse":
if task.status == "cancelled":
session.status = "error"
session.parse_progress = 100
session.parse_message = "解析已取消"
session.parse_error = "用户取消"
session.parse_finished_at = datetime.now()
elif task.status == "failed":
session.status = "error"
session.parse_progress = 100
session.parse_message = "处理失败"
session.parse_error = (task.error_message or "处理失败")[:1000]
session.parse_finished_at = datetime.now()
# done 状态由 celery_tasks 直接写入,这里不重复
elif task.operation in ("generate", "regenerate"):
if task.status == "cancelled":
session.status = "error"
session.parse_message = "生成已取消"
session.parse_error = "用户取消"
elif task.status == "failed":
session.status = "error"
session.parse_message = "生成失败"
session.parse_error = (task.error_message or "生成失败")[:1000]
# done 状态由 celery_tasks 直接写入
db.session.commit()
def _sync_poster_workspace(task, db):
"""同步海报工作区状态。"""
from insurance.models.poster_record import PosterRecord
record = PosterRecord.query.get(task.workspace_id)
if not record:
return
# 安全检查
if record.latest_task_id and str(record.latest_task_id) != task.id:
logger.debug("跳过旧任务同步: task=%s, latest_task=%s", task.id, record.latest_task_id)
return
if task.operation == "generate":
if task.status == "cancelled":
record.task_status = "failed"
record.task_progress = 0
record.task_error = "用户取消"
record.finished_at = datetime.now()
elif task.status == "failed":
record.task_status = "failed"
record.task_progress = 0
record.task_error = (task.error_message or "生成失败")[:1000]
record.finished_at = datetime.now()
# done 状态由 celery_tasks 直接写入
db.session.commit()
def list_active_tasks(user_id: str, artifact_type: str = None) -> list:
"""查询用户活跃任务(任务坞用)。
只返回 queued/running 的任务。终态任务done/failed/cancelled不返回
避免任务坞永久显示已完成任务。
"""
from insurance.models.generation_task import GenerationTask
query = GenerationTask.query.filter_by(user_id=user_id)
query = query.filter(GenerationTask.dock_hidden_at.is_(None))
if artifact_type:
query = query.filter_by(artifact_type=artifact_type)
# 只返回活跃任务
query = query.filter(
GenerationTask.status.in_(["queued", "running"])
)
query = query.order_by(GenerationTask.created_at.desc())
tasks = query.limit(20).all()
return [t.to_dict() for t in tasks]
def list_tasks(user_id: str, artifact_type: str = None,
status: str = None, page: int = 1, page_size: int = 20) -> dict:
"""查询任务列表(任务中心用)。按最后更新时间排序。"""
from insurance.models.generation_task import GenerationTask
query = GenerationTask.query.filter_by(user_id=user_id)
if artifact_type:
query = query.filter_by(artifact_type=artifact_type)
if status:
query = query.filter_by(status=status)
query = query.order_by(GenerationTask.updated_at.desc())
total = query.count()
items = query.offset((page - 1) * page_size).limit(page_size).all()
return {
"total": total,
"items": [t.to_dict() for t in items],
}
def get_task(task_id: str, user_id: str) -> dict:
"""获取任务详情。"""
from insurance.models.generation_task import GenerationTask
task = GenerationTask.query.get(task_id)
if not task or task.user_id != user_id:
return {"code": 1002, "message": "任务不存在", "data": None}
return {"code": 0, "data": task.to_dict()}
def cancel_task(task_id: str, user_id: str) -> dict:
"""取消排队中的任务,并同步工作区状态。"""
from insurance.db.compat import db
from insurance.models.generation_task import GenerationTask
task = GenerationTask.query.get(task_id)
if not task or task.user_id != user_id:
return {"code": 1002, "message": "任务不存在", "data": None}
if task.status != "queued":
return {"code": 1001, "message": "只能取消排队中的任务", "data": None}
task.status = "cancelled"
task.finished_at = datetime.now()
db.session.commit()
# 同步工作区状态
try:
sync_workspace_status(task)
except Exception as exc:
logger.warning("取消任务后同步工作区失败: %s", exc)
return {"code": 0, "data": task.to_dict()}
def hide_task_from_dock(task_id: str, user_id: str) -> dict:
"""从任务坞隐藏任务。"""
from insurance.db.compat import db
from insurance.models.generation_task import GenerationTask
task = GenerationTask.query.get(task_id)
if not task or task.user_id != user_id:
return {"code": 1002, "message": "任务不存在", "data": None}
task.dock_hidden_at = datetime.now()
db.session.commit()
return {"code": 0, "data": task.to_dict()}
def mark_task_viewed(task_id: str, user_id: str) -> dict:
"""标记任务为已查看。"""
from insurance.db.compat import db
from insurance.models.generation_task import GenerationTask
task = GenerationTask.query.get(task_id)
if not task or task.user_id != user_id:
return {"code": 1002, "message": "任务不存在", "data": None}
task.viewed_at = datetime.now()
db.session.commit()
return {"code": 0, "data": task.to_dict()}
# ─── 过期任务恢复 ──────────────────────────────────────────
STALE_TASK_TIMEOUT = 600 # 10 分钟无心跳视为过期
def recover_stale_tasks():
"""启动时恢复过期任务:将长时间无心跳的 running 任务标记为 failed。
应在应用启动时调用。
"""
from insurance.db.compat import db
from insurance.models.generation_task import GenerationTask
cutoff = datetime.now().timestamp() - STALE_TASK_TIMEOUT
# 恢复无心跳的 running 任务
stale_running = GenerationTask.query.filter(
GenerationTask.status == "running",
GenerationTask.heartbeat_at < datetime.fromtimestamp(cutoff),
).all()
for task in stale_running:
task.status = "failed"
task.error_code = "stale"
task.error_message = "任务因服务重启而中断,请重新提交"
task.finished_at = datetime.now()
try:
sync_workspace_status(task)
except Exception as exc:
logger.warning("恢复过期任务同步工作区失败: %s", exc)
logger.warning(f"恢复过期任务: {task.id} (workspace={task.workspace_id})")
# 恢复长时间 queued 的任务(可能 Celery 未消费)
stale_queued = GenerationTask.query.filter(
GenerationTask.status == "queued",
GenerationTask.created_at < datetime.fromtimestamp(cutoff),
).all()
for task in stale_queued:
task.status = "failed"
task.error_code = "queue_timeout"
task.error_message = "任务排队超时,请重新提交"
task.finished_at = datetime.now()
try:
sync_workspace_status(task)
except Exception as exc:
logger.warning("恢复排队超时任务同步工作区失败: %s", exc)
logger.warning(f"恢复排队超时任务: {task.id}")
if stale_running or stale_queued:
db.session.commit()
logger.info(f"已恢复 {len(stale_running)} 个运行过期任务, {len(stale_queued)} 个排队超时任务")