2026-07-27 13:52:09 +08:00
|
|
|
|
"""海报后台任务管理。
|
|
|
|
|
|
|
|
|
|
|
|
使用与 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")
|
2026-07-28 16:45:14 +08:00
|
|
|
|
use_masked_data = bool(data.get("useMaskedData"))
|
2026-07-27 13:52:09 +08:00
|
|
|
|
|
|
|
|
|
|
# 获取模板和产品信息
|
|
|
|
|
|
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
|
|
|
|
|
|
|
2026-07-28 16:45:14 +08:00
|
|
|
|
# 脱敏处理
|
|
|
|
|
|
if use_masked_data:
|
|
|
|
|
|
from insurance.ppt.masking import apply_product_mask, apply_company_mask, mask_text
|
|
|
|
|
|
product_dict = product.to_dict() if product else None
|
|
|
|
|
|
company_dict = company.to_dict() if company else None
|
|
|
|
|
|
if product_dict:
|
|
|
|
|
|
apply_product_mask(product_dict, True)
|
|
|
|
|
|
if company_dict:
|
|
|
|
|
|
apply_company_mask(company_dict, True)
|
|
|
|
|
|
# 替换 copy_content 中的产品名和保司名
|
|
|
|
|
|
if product and company:
|
|
|
|
|
|
real_name = product.display_name
|
|
|
|
|
|
masked_name = (product_dict or {}).get("displayName", "")
|
|
|
|
|
|
real_company = company.display_name
|
|
|
|
|
|
masked_company = (company_dict or {}).get("displayName", "")
|
|
|
|
|
|
replacements = {}
|
|
|
|
|
|
if real_name and masked_name and real_name != masked_name:
|
|
|
|
|
|
replacements[real_name] = masked_name
|
|
|
|
|
|
if real_company and masked_company and real_company != masked_company:
|
|
|
|
|
|
replacements[real_company] = masked_company
|
|
|
|
|
|
if replacements:
|
|
|
|
|
|
for key in copy_content:
|
|
|
|
|
|
if isinstance(copy_content[key], str):
|
|
|
|
|
|
copy_content[key] = mask_text(copy_content[key], replacements)
|
|
|
|
|
|
else:
|
|
|
|
|
|
product_dict = product.to_dict() if product else None
|
|
|
|
|
|
company_dict = {"displayName": company.display_name} if company else None
|
|
|
|
|
|
|
2026-07-27 13:52:09 +08:00
|
|
|
|
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,
|
2026-07-28 16:45:14 +08:00
|
|
|
|
product=product_dict,
|
|
|
|
|
|
company=company_dict,
|
2026-07-27 13:52:09 +08:00
|
|
|
|
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)
|