"""海报业务逻辑服务。""" import asyncio import copy import hashlib import json import os import re import uuid import logging from datetime import datetime, timedelta from insurance.db.compat import db from insurance.models.ppt_config import PptProduct, PptCompany from insurance.models.poster_case_upload import PosterCaseUpload from insurance.models.poster_template_model import PosterTemplate from insurance.models.poster_copy_template import PosterCopyTemplate from insurance.models.poster_record import PosterRecord from insurance.poster.product_source_resolver import ( ProductSourceError, case_matches_product, product_snapshot, resolve_product_source, ) logger = logging.getLogger(__name__) def _safe_user_id(user_id: str) -> str: """清理 user_id 中的路径分隔符,防止路径穿越。""" return re.sub(r'[/\\.]', '_', user_id) def _run_async(coro): """安全执行异步函数,处理事件循环已存在的情况。""" try: loop = asyncio.get_running_loop() except RuntimeError: loop = None if loop and loop.is_running(): # 已有运行中的事件循环,用线程池执行 import concurrent.futures with concurrent.futures.ThreadPoolExecutor() as pool: return pool.submit(asyncio.run, coro).result() return asyncio.run(coro) def _replace_names(value, replacements: dict[str, str]): """递归替换规则、卖点和文案中的产品/保司名称。""" if isinstance(value, str): for real_name, masked_name in replacements.items(): value = value.replace(real_name, masked_name) return value if isinstance(value, list): return [_replace_names(item, replacements) for item in value] if isinstance(value, dict): return {key: _replace_names(item, replacements) for key, item in value.items()} return value def _configured_context(context: dict) -> tuple[dict, dict, dict, dict]: """按后台品牌策略返回规则、产品和保司副本。""" from insurance.ppt.masking import apply_brand_policy, build_brand_policy product_data = copy.deepcopy(context.get("productData") or {}) company_data = copy.deepcopy(context.get("companyData") or {}) rules = copy.deepcopy(context.get("rules") or {}) real_product_name = context.get("productName") or "" real_company_name = context.get("companyName") or "" brand_policy = build_brand_policy(company_data, [product_data]) company_data, product_data = apply_brand_policy( company_data, product_data, brand_policy ) replacements = {} masked_product_name = product_data.get("displayName", "") masked_company_name = company_data.get("displayName", "") if real_product_name and masked_product_name and real_product_name != masked_product_name: replacements[real_product_name] = masked_product_name if real_company_name and masked_company_name and real_company_name != masked_company_name: replacements[real_company_name] = masked_company_name return _replace_names(rules, replacements), product_data, company_data, brand_policy class PosterService: """海报业务逻辑。""" def get_reviewed_products(self) -> dict: """获取已 reviewed 的产品列表(供选择)。""" products = PptProduct.query.filter( PptProduct.manual_parse_status == "reviewed", PptProduct.status == 1, PptProduct.deleted_at.is_(None), ).all() if not products: return {"code": 0, "data": []} # 批量查询公司(避免 N+1) company_ids = list({p.company_id for p in products}) companies = { c.id: c for c in PptCompany.query.filter( PptCompany.id.in_(company_ids), PptCompany.status == 1, PptCompany.deleted_at.is_(None), ).all() } result = [] for p in products: company = companies.get(p.company_id) if not company: continue result.append({ **p.to_dict(), "companyName": company.display_name, }) return {"code": 0, "data": result} def upload_case(self, user_id: str, product_id: str, file, password: str = "", product_source: dict | None = None) -> dict: """上传计划书 PDF + 排队解析(异步)。""" source_data = { "productId": product_id, "productSource": product_source or {}, } try: context = resolve_product_source(user_id, source_data) except ProductSourceError as exc: return {"code": exc.code, "message": exc.message, "data": None} # 文件安全校验(SEC-P1-01) from insurance.utils.security import prepare_pdf_upload is_valid, err_msg, pdf_bytes = prepare_pdf_upload(file, password) if not is_valid: code = 4003 if "密码" in err_msg else 4002 return {"code": code, "message": err_msg, "data": None} digest = hashlib.sha256(pdf_bytes).hexdigest() existing = PosterCaseUpload.query.filter_by( user_id=user_id, file_hash=digest, product_source_type=context["sourceType"], product_source_id=context["sourceId"], ).filter(PosterCaseUpload.parse_status.in_([ "queued", "parsing", "parsed", "partial", ])).order_by(PosterCaseUpload.id.desc()).first() if existing: data = existing.to_dict() data["deduplicated"] = True message = ( "相同计划书正在解析" if existing.parse_status in ("queued", "parsing") else "已使用相同计划书的解析结果" ) return {"code": 0, "message": message, "data": data} # 保存文件(使用持久化存储) safe_uid = _safe_user_id(user_id) from insurance.config import get_storage_root upload_dir = os.path.join(get_storage_root(), "uploads", "poster-cases") os.makedirs(upload_dir, exist_ok=True) filename = f"{safe_uid}_{uuid.uuid4().hex[:8]}.pdf" filepath = os.path.join(upload_dir, filename) with open(filepath, "wb") as output: output.write(pdf_bytes) # 创建记录 record = PosterCaseUpload( user_id=user_id, product_id=context.get("productId") or "", product_source_type=context["sourceType"], product_source_id=context["sourceId"], product_snapshot_json=json.dumps(product_snapshot(context), ensure_ascii=False), source_file_url=filepath, file_hash=digest, parse_status="queued", parse_progress=5, parse_message="任务已提交,等待解析...", ) db.session.add(record) db.session.commit() # 启动后台解析任务 from insurance.poster.tasks import start_case_parse_task if not start_case_parse_task(record.id): record.parse_status = "failed" record.parse_progress = 100 record.parse_message = "任务提交失败" record.parse_error = "解析任务未能提交到后台队列" record.parse_finished_at = datetime.now() db.session.commit() return {"code": 5001, "message": "任务排队失败,请重试", "data": None} return {"code": 0, "data": record.to_dict()} def get_case_upload(self, record_id: int, user_id: str) -> dict: """获取解析结果。""" record = PosterCaseUpload.query.get(record_id) if not record or record.user_id != user_id: return {"code": 1002, "message": "记录不存在", "data": None} return {"code": 0, "data": record.to_dict()} def confirm_case_upload(self, record_id: int, user_id: str, data: dict) -> dict: """人工核对/修正解析数据。""" record = PosterCaseUpload.query.get(record_id) if not record or record.user_id != user_id: return {"code": 1002, "message": "记录不存在", "data": None} record.confirmed_data = json.dumps(data.get("confirmedData", {}), ensure_ascii=False) record.confirmed_by = user_id from datetime import datetime record.confirmed_at = datetime.now() db.session.commit() return {"code": 0, "data": record.to_dict()} def reparse_case_upload(self, record_id: int, user_id: str) -> dict: """保留源文件,清除旧结果并使用当前解析器重新解析。""" record = PosterCaseUpload.query.get(record_id) if not record or record.user_id != user_id: return {"code": 1002, "message": "记录不存在", "data": None} if not record.source_file_url or not os.path.exists(record.source_file_url): return {"code": 1002, "message": "源计划书不存在,请重新上传", "data": None} if record.parse_status in ("queued", "parsing"): return {"code": 0, "message": "计划书正在解析中", "data": record.to_dict()} from insurance.poster.tasks import start_case_parse_task record.parse_status = "queued" record.parse_progress = 5 record.parse_message = "任务已提交,等待解析..." record.parse_error = None record.parse_task_id = None record.parse_started_at = None record.parse_heartbeat_at = None record.parse_finished_at = None record.parsed_data = None record.confirmed_data = None record.confirmed_by = None record.confirmed_at = None db.session.commit() if not start_case_parse_task(record.id): record.parse_status = "failed" record.parse_progress = 100 record.parse_message = "任务提交失败" record.parse_error = "解析任务未能提交到后台队列" record.parse_finished_at = datetime.now() db.session.commit() return {"code": 5001, "message": "重新解析任务启动失败,请重试", "data": None} return {"code": 0, "data": record.to_dict()} def get_templates(self) -> dict: """获取可用海报模板列表。""" templates = PosterTemplate.query.filter_by(status=1).order_by( PosterTemplate.id.asc() ).all() return {"code": 0, "data": [t.to_dict() for t in templates]} def get_copy_templates(self) -> dict: """获取可用文案模板列表。""" templates = PosterCopyTemplate.query.filter_by( status=1, deleted_at=None ).order_by( PosterCopyTemplate.id.asc() ).all() return {"code": 0, "data": [t.to_dict() for t in templates]} def generate_copy(self, user_id: str, data: dict) -> dict: """生成文案(template/ai 模式)。""" mode = data.get("mode", "template") case_upload_id = data.get("caseUploadId") try: context = resolve_product_source(user_id, data) except ProductSourceError as exc: return {"code": exc.code, "message": exc.message, "data": None} product_rules, _product_data, _company_data, _brand_policy = _configured_context(context) # 获取客户数据(校验 case 所有权 — 防止越权读取他人客户数据) customer_data = {} if case_upload_id: case = PosterCaseUpload.query.get(case_upload_id) if not case or case.user_id != user_id: return {"code": 404, "message": "记录不存在", "data": None} if not case_matches_product(case, context): return {"code": 1002, "message": "计划书与所选产品不一致", "data": None} if case.confirmed_data: try: customer_data = json.loads(case.confirmed_data) except (json.JSONDecodeError, TypeError): logger.warning("confirmed_data JSON 解析失败: case_id=%s", case_upload_id) customer_data = {} if mode == "template": template_id = data.get("templateId") template = PosterCopyTemplate.query.filter_by( id=template_id, status=1, deleted_at=None ).first() if not template: return {"code": 1002, "message": "文案模板不存在", "data": None} from insurance.poster.copy_generator import CopyGenerator generator = CopyGenerator() result = generator.generate_template_copy(template.content, product_rules, customer_data) return {"code": 0, "data": {"copy": result, "mode": "template"}} else: style = data.get("style", "专业") import asyncio from insurance.poster.copy_generator import CopyGenerator generator = CopyGenerator() result = _run_async(generator.generate_ai_copy(product_rules, customer_data, style)) return {"code": 0, "data": {"copy": result, "mode": "ai"}} def check_compliance(self, data: dict) -> dict: """检查文案并返回可定位的问题。""" from insurance.poster.compliance import check_copy_compliance copy_content = data.get("copyContent") if "copyContent" in data else data return {"code": 0, "data": check_copy_compliance(copy_content or {})} def generate_poster(self, user_id: str, data: dict) -> dict: """生成海报图片(异步 — 排队后立即返回)。""" case_upload_id = data.get("caseUploadId") template_id = data.get("templateId") copy_content = data.get("copyContent", {}) reference_image = data.get("referenceImage") output_mode = data.get("outputMode", "single") force_regenerate = data.get("regenerate", False) if not template_id: return {"code": 1001, "message": "请选择海报模板", "data": None} from insurance.poster.format_registry import ( PosterFormatError, resolve_custom_long_height, resolve_poster_format, ) try: format_spec = resolve_poster_format( format_id=data.get("formatId"), legacy_size=data.get("size"), output_mode=output_mode, ) custom_height = resolve_custom_long_height(data.get("customHeight"), output_mode) except PosterFormatError as exc: return {"code": 1002, "message": str(exc), "data": None} format_id = format_spec["id"] size = ( f"{format_spec['output']['width']}x{custom_height}" if custom_height is not None else format_spec["exportSize"] ) from insurance.poster.compliance import check_copy_compliance compliance_result = check_copy_compliance(copy_content) if compliance_result["status"] == "block": return { "code": 1002, "message": "文案包含不合规表达,请修改后再生成", "data": compliance_result, } # 提取参考图元数据(不含 base64 数据,仅保存文件名供展示) ref_images_meta = [] if reference_image: ref_images_meta = [{"name": "参考图", "hasData": True}] template = PosterTemplate.query.filter_by(id=template_id, status=1).first() if not template: return {"code": 1002, "message": "海报模板不存在或已停用", "data": None} if not template.supports(output_mode, format_id): return {"code": 1002, "message": "所选模板不支持当前海报格式", "data": None} try: context = resolve_product_source(user_id, data) except ProductSourceError as exc: return {"code": exc.code, "message": exc.message, "data": None} rules, product_data, company_data, brand_policy = _configured_context(context) source_snapshot = product_snapshot(context) # 校验 case 所有权(SEC-P0-03) if case_upload_id: case = PosterCaseUpload.query.get(case_upload_id) if not case or case.user_id != user_id: return {"code": 404, "message": "记录不存在", "data": None} if not case_matches_product(case, context): return {"code": 1002, "message": "计划书与所选产品不一致", "data": None} if case.confirmed_data is None: return {"code": 1002, "message": "请先确认解析数据", "data": None} case_facts = {} if case_upload_id and case and case.confirmed_data: try: case_facts = json.loads(case.confirmed_data) except (json.JSONDecodeError, TypeError): return {"code": 1002, "message": "已确认的计划书数据无效", "data": None} from insurance.poster.render_document_builder import build_render_document render_document = build_render_document( format_spec=format_spec, template=template.to_dict(), copy_content=copy_content, case_facts=case_facts, product_rules=rules, plan_type=context.get("planType") or "other", compliance_revision=compliance_result["revision"], custom_height=custom_height, sections=data.get("sections"), brand={ "productName": case_facts.get("product_name") or product_data.get("displayName", ""), "companyName": case_facts.get("company_name") or company_data.get("displayName", ""), "logoUrl": company_data.get("logoUrl", ""), }, ) # 幂等检查(TASK-P1-02):拦截同一页面的连续重复提交。 if not force_regenerate: active_query = PosterRecord.query.filter( PosterRecord.user_id == user_id, PosterRecord.template_id == template_id, PosterRecord.export_size == size, PosterRecord.task_status.in_(["pending", "queued", "generating"]), PosterRecord.created_at >= datetime.now() - timedelta(seconds=15), ) if case_upload_id: active_query = active_query.filter( PosterRecord.case_upload_id == case_upload_id, ) else: active_query = active_query.filter( PosterRecord.product_source_type == context["sourceType"], PosterRecord.product_source_id == context["sourceId"], ) existing = active_query.order_by(PosterRecord.created_at.desc()).first() if existing: return {"code": 0, "data": existing.to_dict()} # 创建记录(任务状态为 pending) ai_raw_content = data.get("aiRawContent") record = PosterRecord( user_id=user_id, product_id=context.get("productId"), product_source_type=context["sourceType"], product_source_id=context["sourceId"], product_snapshot_json=json.dumps(source_snapshot, ensure_ascii=False), case_upload_id=case_upload_id, template_id=template_id, copy_mode=data.get("copyMode", "ai"), copy_content=json.dumps(copy_content, ensure_ascii=False) if copy_content else None, ai_raw_content=json.dumps(ai_raw_content, ensure_ascii=False) if ai_raw_content else None, export_size=size, reference_image_used=reference_image, task_status="pending", task_progress=0, extra_data=json.dumps({ "brandPolicy": brand_policy, "referenceImages": ref_images_meta, "outputMode": output_mode, "formatId": format_id, "customHeight": custom_height, "compliance": compliance_result, }, ensure_ascii=False), document_json=json.dumps(render_document, ensure_ascii=False), document_revision=1, compliance_json=json.dumps(compliance_result, ensure_ascii=False), compliance_revision=compliance_result["revision"], ) db.session.add(record) db.session.commit() # 启动后台生成任务(Celery) from insurance.generation import task_service task_input = dict(data) task_input.pop("useMaskedData", None) task_input["formatId"] = format_id task_input["size"] = size task_input["renderDocument"] = render_document task_input["brandPolicy"] = brand_policy task_input["productSource"] = { "type": context["sourceType"], "id": context["sourceId"], } task_result = task_service.create_task( user_id=user_id, artifact_type="poster", operation="generate", workspace_id=str(record.id), title=f"海报 #{record.id}", input_snapshot=task_input, input_revision=record.draft_revision or 1, idempotency_key=f"poster_gen_{record.id}", ) if task_result.get("code") == 0: record.task_status = "queued" record.latest_task_id = task_result["data"]["id"] db.session.commit() else: record.task_status = "failed" record.task_error = task_result.get("message", "任务排队失败") db.session.commit() return {"code": 5001, "message": "任务排队失败,请重试", "data": None} return {"code": 0, "data": record.to_dict()} def list_records(self, user_id: str, params: dict) -> dict: """海报生成记录列表。""" query = db.session.query(PosterRecord).filter(PosterRecord.user_id == user_id) query = query.order_by(PosterRecord.created_at.desc()) page = max(1, params.get("page", 1)) page_size = min(100, max(1, params.get("page_size", 20))) total = query.count() items = query.offset((page - 1) * page_size).limit(page_size).all() return { "code": 0, "data": { "total": total, "items": [r.to_dict() for r in items], }, } def get_record(self, record_id: int, user_id: str) -> dict: """记录详情(含任务状态,用于轮询)。""" record = PosterRecord.query.get(record_id) if not record or record.user_id != user_id: return {"code": 1002, "message": "记录不存在", "data": None} return {"code": 0, "data": record.to_dict()} def save_document(self, record_id: int, user_id: str, document: dict) -> dict: """保存可编辑海报文档,不覆盖历史导出文件。""" record = PosterRecord.query.get(record_id) if not record or record.user_id != user_id: return {"code": 1002, "message": "记录不存在", "data": None} record.document_revision = (record.document_revision or 1) + 1 document = dict(document or {}) document["revision"] = record.document_revision record.document_json = json.dumps(document, ensure_ascii=False) record.draft_revision = (record.draft_revision or 1) + 1 db.session.commit() return {"code": 0, "data": record.to_dict()} def save_rendered(self, record_id: int, user_id: str, file, document: dict) -> dict: """保存浏览器合成后的最终 PNG,并绑定同一文档版本。""" record = PosterRecord.query.get(record_id) if not record or record.user_id != user_id: return {"code": 1002, "message": "记录不存在", "data": None} if not file: return {"code": 1002, "message": "请上传最终海报文件", "data": None} if not isinstance(document, dict): return {"code": 1002, "message": "海报文档必须是 JSON 对象", "data": None} image_bytes = file.read() if len(image_bytes) > 30 * 1024 * 1024: return {"code": 1002, "message": "海报文件不能超过 30MB", "data": None} try: extra_data = json.loads(record.extra_data or "{}") except (json.JSONDecodeError, TypeError): extra_data = {} if not isinstance(extra_data, dict): extra_data = {} format_id = extra_data.get("formatId") or (document or {}).get("formatId") if not format_id: from insurance.poster.format_registry import PosterFormatError, resolve_poster_format try: format_id = resolve_poster_format( legacy_size=record.export_size, output_mode=extra_data.get("outputMode", "single"), )["id"] except PosterFormatError as exc: return {"code": 1002, "message": str(exc), "data": None} from insurance.poster.render_validation import ( PosterRenderValidationError, validate_rendered_png, ) requested_output = document.get("requestedOutput") or {} expected_height = ( requested_output.get("height") if format_id == "long_1242_auto" else None ) try: actual_output = validate_rendered_png( image_bytes, format_id, expected_height=expected_height, ) except PosterRenderValidationError as exc: return {"code": 1002, "message": str(exc), "data": None} from insurance.config import get_storage_root output_dir = os.path.join( get_storage_root(), "outputs", "posters", str(record.latest_task_id or record.id) ) os.makedirs(output_dir, exist_ok=True) revision = (record.document_revision or 1) + 1 filepath = os.path.join(output_dir, f"final_v{revision}.png") temp_path = f"{filepath}.{uuid.uuid4().hex}.tmp" try: with open(temp_path, "wb") as output: output.write(image_bytes) os.replace(temp_path, filepath) finally: if os.path.exists(temp_path): os.remove(temp_path) document = dict(document or {}) document["revision"] = revision document["formatId"] = format_id document["actualOutput"] = actual_output document["byteSize"] = len(image_bytes) document["validationResult"] = "passed" record.document_json = json.dumps(document, ensure_ascii=False) record.document_revision = revision if format_id == "long_1242_auto": record.export_size = ( f"1242x{expected_height}" if expected_height is not None else "1242xauto" ) record.export_url = filepath record.export_format = "png" record.task_status = "done" record.task_progress = 100 extra_data.update({ "formatId": format_id, "actualOutput": actual_output, "byteSize": len(image_bytes), "validationResult": "passed", "customHeight": expected_height, }) record.extra_data = json.dumps(extra_data, ensure_ascii=False) record.draft_revision = (record.draft_revision or 1) + 1 record.generated_revision = record.draft_revision db.session.commit() return {"code": 0, "data": record.to_dict()}