主要结果: PPT: 已完成步骤可自由返回,刷新后保留当前步骤。 结果页改为固定三栏视口,缩略图和质量面板独立滚动。 修复画布宽度计算和宽高自适应缩放。 返回配置页时恢复上次模板。 Worker 禁止静默替换模板,不兼容时明确失败。 结果页显示实际应用的模板。 海报解析与合规: 修复 insured/policy 嵌套字段映射。 增加 partial/failed 解析质量判断和缺失字段提示。 传入险种、产品、保司和别名上下文。 修复性别、币种归一化和错误字段计数。 新增服务端字符级合规接口与生成前复检。 不合规文字可高亮、点击定位并选中对应文字。 海报编辑: AI 图片现在作为背景资产,不再替换整个海报。 单图、长图生成后仍可编辑标题、正文、CTA、数据和卖点。 背景生成失败不会清空当前编辑内容。 下载时合成背景、文字、数据、图表和免责声明。 增加可编辑文档快照、最终 PNG 保存和历史继续编辑。 新增数据库迁移:[migrate_028.py](/D:/work/code/python/coding/baodanagent/api/insurance/db/migrate_028.py) 验证结果: 专项及相关后端测试:26 passed, 1 skipped Python 全模块编译检查:通过 前端生产构建:通过 git diff --check:通过 全量测试收集受本机缺少 python-pptx 依赖影响,报错为 ModuleNotFoundError: pptx,不是本次修改产生的测试失败。
451 lines
20 KiB
Python
451 lines
20 KiB
Python
"""海报业务逻辑服务。"""
|
||
import asyncio
|
||
import copy
|
||
import json
|
||
import os
|
||
import re
|
||
import uuid
|
||
import logging
|
||
from flask import current_app
|
||
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}
|
||
|
||
# 保存文件(使用持久化存储)
|
||
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,
|
||
parse_status="pending",
|
||
)
|
||
db.session.add(record)
|
||
db.session.commit()
|
||
|
||
# 启动后台解析任务
|
||
from insurance.poster.tasks import start_case_parse_task
|
||
if start_case_parse_task(current_app._get_current_object(), record.id):
|
||
record.parse_status = "queued"
|
||
db.session.commit()
|
||
else:
|
||
record.parse_status = "failed"
|
||
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 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", {})
|
||
size = data.get("size", "1024x1792")
|
||
reference_image = data.get("referenceImage")
|
||
output_mode = data.get("outputMode", "single")
|
||
force_regenerate = data.get("regenerate", False)
|
||
|
||
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}
|
||
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}
|
||
|
||
# 幂等检查(TASK-P1-02):相同参数的未完成任务直接返回
|
||
if not force_regenerate and case_upload_id:
|
||
existing = PosterRecord.query.filter(
|
||
PosterRecord.user_id == user_id,
|
||
PosterRecord.case_upload_id == case_upload_id,
|
||
PosterRecord.template_id == template_id,
|
||
PosterRecord.task_status.in_(["pending", "queued", "generating"]),
|
||
).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,
|
||
"compliance": compliance_result,
|
||
}, ensure_ascii=False),
|
||
document_json=json.dumps({
|
||
"revision": 1,
|
||
"templateId": template_id,
|
||
"outputMode": output_mode,
|
||
"copy": copy_content,
|
||
"facts": (
|
||
json.loads(case.confirmed_data)
|
||
if case_upload_id and case and case.confirmed_data else {}
|
||
),
|
||
"style": {},
|
||
}, 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["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}
|
||
image_bytes = file.read()
|
||
if len(image_bytes) > 30 * 1024 * 1024:
|
||
return {"code": 1002, "message": "海报文件不能超过 30MB", "data": None}
|
||
if not image_bytes.startswith(b"\x89PNG\r\n\x1a\n"):
|
||
return {"code": 1002, "message": "仅支持 PNG 海报文件", "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")
|
||
with open(filepath, "wb") as output:
|
||
output.write(image_bytes)
|
||
|
||
document = dict(document or {})
|
||
document["revision"] = revision
|
||
record.document_json = json.dumps(document, ensure_ascii=False)
|
||
record.document_revision = revision
|
||
record.export_url = filepath
|
||
record.export_format = "png"
|
||
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()}
|