baodan/api/insurance/poster/tasks.py
wsb1224 e3479f0546 上线阻断问题全部修复
#	问题	修复	文件
1	前端构建失败(引号错误)	size="small type=" → size="small" type="	PosterHistoryPage.vue
2	migrate_014 ORM vs 缺失列	全部改为原始 SQL,不再引用 ORM 模型	migrate_014.py
3	cleanup 字段名错误	output_path → ppt_path	cleanup.py
4	文案生成 case 越权	添加 case.user_id != user_id 校验	poster/service.py
5	存储路径未接通持久化卷	全部改用 get_storage_root()(默认 /app/api/storage/insurance)	config.py, ppt/routes.py, poster/service.py, poster/tasks.py
高风险问题修复
#	问题	修复	文件
6	migrate_019 rollback 撤销成功字段	每个 ALTER 后立即 commit,失败只回滚当前语句	migrate_019.py
7	迁移锁 Windows 不兼容 + 句柄未持久化	全局变量保存锁句柄,支持 Windows msvcrt	api/insurance/db/__init__.py
8	PDF 校验异常时放行	异常返回 False(文件损坏)	security.py
9	健康检查始终返回成功	缺少关键资源时返回 503 + missing 列表	poster/routes.py
10	短密钥掩码泄露原值	≤4 字符返回 ****	ppt_admin_service.py
11	设置无键名白名单	添加 _ALLOWED_SETTING_KEYS 白名单	ppt_admin_service.py
12	容器重启任务永久 stuck	添加 recover_stale_tasks() 启动恢复函数	poster/tasks.py, ppt/parse_worker.py
2026-07-27 13:52:09 +08:00

289 lines
9.4 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.

"""海报后台任务管理。
使用与 PPT parse_worker 相同的后台线程 + Redis 锁模式。
任务状态通过数据库字段追踪,前端轮询获取进度。
"""
import json
import logging
import re
import threading
import uuid
import os
from datetime import datetime
from insurance.db.compat import db
logger = logging.getLogger(__name__)
_local_locks: set = set()
_local_locks_guard = threading.Lock()
_redis_locks: set = set()
# 超过此时间仍为 queued/generating 的任务视为过期(秒)
STALE_TASK_TIMEOUT = 600 # 10 分钟
def recover_stale_tasks():
"""启动时恢复过期任务:将长时间 queued/generating 的任务标记为 failed。
应在应用启动时调用。
"""
from insurance.models.poster_record import PosterRecord
from insurance.models.poster_case_upload import PosterCaseUpload
cutoff = datetime.now().timestamp() - STALE_TASK_TIMEOUT
# 恢复过期的海报生成任务
stale_records = PosterRecord.query.filter(
PosterRecord.task_status.in_(["queued", "generating"]),
PosterRecord.created_at < datetime.fromtimestamp(cutoff),
).all()
for record in stale_records:
record.task_status = "failed"
record.task_error = "任务因服务重启而中断,请重新生成"
record.finished_at = datetime.now()
logger.warning(f"恢复过期海报任务: record_id={record.id}")
# 恢复过期的计划书解析任务
stale_cases = PosterCaseUpload.query.filter(
PosterCaseUpload.parse_status.in_(["queued", "parsing"]),
PosterCaseUpload.created_at < datetime.fromtimestamp(cutoff),
).all()
for case in stale_cases:
case.parse_status = "failed"
logger.warning(f"恢复过期解析任务: case_id={case.id}")
if stale_records or stale_cases:
db.session.commit()
logger.info(f"已恢复 {len(stale_records)} 个海报任务, {len(stale_cases)} 个解析任务")
# ─── 计划书解析任务 ───────────────────────────────────────
def start_case_parse_task(app, case_upload_id: int) -> bool:
"""启动后台计划书解析任务,返回是否新启动。"""
lock_key = f"poster_case_parse:{case_upload_id}"
if not _acquire_lock(lock_key):
return False
thread = threading.Thread(
target=_run_case_parse_task,
args=(app, case_upload_id, lock_key),
daemon=True,
)
thread.start()
return True
def _run_case_parse_task(app, case_upload_id: int, lock_key: str):
with app.app_context():
try:
_execute_case_parse(case_upload_id)
except Exception as exc:
logger.error(f"计划书解析任务失败 [{case_upload_id}]: {exc}", exc_info=True)
_mark_case_failed(case_upload_id, str(exc))
finally:
_release_lock(lock_key)
db.session.remove()
def _execute_case_parse(case_upload_id: int):
from insurance.models.poster_case_upload import PosterCaseUpload
record = PosterCaseUpload.query.get(case_upload_id)
if not record:
return
filepath = record.source_file_url
if not filepath or not os.path.exists(filepath):
record.parse_status = "failed"
db.session.commit()
return
record.parse_status = "parsing"
db.session.commit()
from insurance.ppt.extraction import ExtractionOrchestrator
orchestrator = ExtractionOrchestrator(use_cache=False)
import asyncio
parsed = asyncio.run(orchestrator.extract_for_poster(filepath))
record = PosterCaseUpload.query.get(case_upload_id)
if not record:
return
record.parsed_data = json.dumps(parsed, ensure_ascii=False)
record.parse_status = "parsed"
db.session.commit()
def _mark_case_failed(case_upload_id: int, error: str):
from insurance.models.poster_case_upload import PosterCaseUpload
record = PosterCaseUpload.query.get(case_upload_id)
if not record:
return
record.parse_status = "failed"
db.session.commit()
# ─── 海报图片生成任务 ──────────────────────────────────────
def start_poster_generate_task(app, record_id: int, user_id: str, data: dict) -> bool:
"""启动后台海报图片生成任务,返回是否新启动。"""
lock_key = f"poster_generate:{record_id}"
if not _acquire_lock(lock_key):
return False
thread = threading.Thread(
target=_run_poster_generate_task,
args=(app, record_id, user_id, data, lock_key),
daemon=True,
)
thread.start()
return True
def _run_poster_generate_task(app, record_id: int, user_id: str, data: dict, lock_key: str):
with app.app_context():
try:
_execute_poster_generate(record_id, user_id, data)
except Exception as exc:
logger.error(f"海报生成任务失败 [{record_id}]: {exc}", exc_info=True)
_mark_poster_failed(record_id, str(exc))
finally:
_release_lock(lock_key)
db.session.remove()
def _execute_poster_generate(record_id: int, user_id: str, data: dict):
from insurance.models.poster_record import PosterRecord
from insurance.models.poster_template_model import PosterTemplate
from insurance.models.ppt_config import PptProduct, PptCompany
record = PosterRecord.query.get(record_id)
if not record:
return
record.task_status = "generating"
record.task_progress = 10
record.started_at = datetime.now()
db.session.commit()
template_id = data.get("templateId")
product_id = data.get("productId")
copy_content = data.get("copyContent", {})
size = data.get("size", "1024x1792")
reference_image = data.get("referenceImage")
# 获取模板和产品信息
poster_template = PosterTemplate.query.get(template_id) if template_id else None
product = PptProduct.query.get(product_id) if product_id else None
company = PptCompany.query.get(product.company_id) if product else None
record.task_progress = 30
db.session.commit()
# 组装 prompt
from insurance.poster.image_generator import PosterImageGenerator
generator = PosterImageGenerator()
prompt = generator.build_prompt(
template=poster_template.to_dict() if poster_template else None,
product=product.to_dict() if product else None,
company={"displayName": company.display_name} if company else None,
copy=copy_content,
size=size,
)
record.task_progress = 50
db.session.commit()
# 生成图片
generation_mode = "ai"
provider_info = {}
try:
image_bytes, provider_info = generator.generate(prompt, size=size, reference_image=reference_image)
except Exception as e:
logger.warning(f"图片 API 失败,使用降级方案: {e}")
from insurance.poster.image_generator import generate_fallback
image_bytes = generate_fallback(copy_content, size=size)
generation_mode = "fallback"
record.task_progress = 80
db.session.commit()
# 保存文件(使用持久化存储)
from insurance.config import get_storage_root
output_dir = os.path.join(get_storage_root(), "outputs", "posters")
os.makedirs(output_dir, exist_ok=True)
safe_uid = re.sub(r'[/\\.]', '_', user_id)
filename = f"poster_{safe_uid}_{uuid.uuid4().hex[:8]}.png"
filepath = os.path.join(output_dir, filename)
with open(filepath, "wb") as f:
f.write(image_bytes)
# 更新记录
record = PosterRecord.query.get(record_id)
if not record:
return
record.export_url = filepath
record.export_format = "png"
record.generation_mode = generation_mode
record.image_provider = provider_info.get("provider", "")
record.image_model = provider_info.get("model", "")
record.prompt_used = prompt[:2000] if prompt else None
record.task_status = "done"
record.task_progress = 100
record.finished_at = datetime.now()
db.session.commit()
def _mark_poster_failed(record_id: int, error: str):
from insurance.models.poster_record import PosterRecord
record = PosterRecord.query.get(record_id)
if not record:
return
record.task_status = "failed"
record.task_error = error[:1000]
record.task_progress = 100
record.finished_at = datetime.now()
db.session.commit()
# ─── 锁管理 ────────────────────────────────────────────────
def _acquire_lock(key: str) -> bool:
"""获取任务锁(优先 Redis降级本地内存"""
redis_key = f"poster_task_lock:{key}"
try:
from insurance.db.compat import redis_client
if redis_client and redis_client.set(redis_key, "1", nx=True, ex=1800):
_redis_locks.add(key)
return True
if redis_client:
return False
except Exception:
pass
with _local_locks_guard:
if key in _local_locks:
return False
_local_locks.add(key)
return True
def _release_lock(key: str):
"""释放任务锁。"""
if key in _redis_locks:
try:
from insurance.db.compat import redis_client
if redis_client:
redis_client.delete(f"poster_task_lock:{key}")
except Exception:
pass
_redis_locks.discard(key)
with _local_locks_guard:
_local_locks.discard(key)