"""海报后台任务管理。 使用与 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)