baodan/api/insurance/poster/tasks.py
2026-07-28 16:45:14 +08:00

318 lines
11 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")
use_masked_data = bool(data.get("useMaskedData"))
# 获取模板和产品信息
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
# 脱敏处理
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
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_dict,
company=company_dict,
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)