baodan/api/insurance/generation/task_service.py
wsb1224 c7df34d0cd 07-27 阶段 状态 进度
阶段 1:数据库与工作区	 完成	100%
阶段 2:Celery 后台任务	 完成	100%
阶段 3:前端刷新恢复与多工作区	 完成	100%
阶段 4:任务坞与任务中心	 完成	100%
阶段 5:版本化编辑	 完成	100%
阶段 6:测试与灰度	 未开始	0%
2026-07-28 17:53:14 +08:00

207 lines
7.3 KiB
Python

"""统一任务服务:创建、查询、取消任务。"""
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) -> 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,
idempotency_key=idempotency_key,
)
db.session.add(task)
db.session.commit()
# 提交到 Celery
_dispatch_to_celery(task)
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,
parse_poster_task, generate_poster_task,
)
task_map = {
("ppt", "parse"): parse_ppt_task,
("ppt", "generate"): generate_ppt_task,
("poster", "parse"): parse_poster_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.delay(task.id)
task.celery_task_id = result.id
from insurance.db.compat import db
db.session.commit()
logger.info(f"任务 {task.id} 已提交到 Celery: {result.id}")
def list_active_tasks(user_id: str, artifact_type: str = None) -> list:
"""查询用户活跃任务(任务坞用)。"""
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", "done", "failed"])
)
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.created_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()
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()}
# ─── 过期任务恢复 ──────────────────────────────────────────
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()
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()
logger.warning(f"恢复排队超时任务: {task.id}")
if stale_running or stale_queued:
db.session.commit()
logger.info(f"已恢复 {len(stale_running)} 个运行过期任务, {len(stale_queued)} 个排队超时任务")