baodan/api/insurance/poster/service.py
wsb1224 b8c4e8b672 主要完成内容:
修复 PPT 异步任务无法生成的问题,包括任务变量引用错误、失败状态回写、心跳缺失任务恢复。
脱敏改为保司/产品后台统一配置,生成端不再让用户选择;任务创建时保存策略快照。
保司支持独立控制 PPT、海报 Logo 显示。
PPT 核验新增吸烟状态、币种及三个条件字段。
利益演示、退保提取调整为警告,不再阻止生成。
PPT 生成完成后可以直接返回数据核验页修改。
建立不同险种、单图/长图共六套海报字段画像。
PPT“生成场景”支持后台新增、启停和删除。
保司、产品、PPT 模板、文案模板均支持安全删除。
内置模板禁止删除,只允许停用;存在关联数据时拒绝危险删除。
补充策略变更及删除审计日志。
更新 API 文档、部署文档及修复计划实施记录。
关键交付文件:
[数据库迁移 migrate_027.py](D:/work/code/python/coding/baodanagent/api/insurance/db/migrate_027.py)
[海报字段画像 field_profiles.py](D:/work/code/python/coding/baodanagent/api/insurance/poster/field_profiles.py)
[动态场景服务 scenarios.py](D:/work/code/python/coding/baodanagent/api/insurance/ppt/scenarios.py)
[新增回归测试](D:/work/code/python/coding/baodanagent/tests/ppt_poster_optimization_test.py)
[优化修复计划书](D:/work/code/python/coding/baodanagent/docs/保险智能客服系统_PPT与海报优化修复计划书_20260731.md)
验证结果:
核心链路测试:37 passed,1 skipped
扩展回归测试:140 passed
PPT 渲染器测试:6 passed
前端生产构建:通过
Python 编译检查:通过
完整测试集:190 passed,1 failed
唯一失败为 tests/test_chat_save.py::test_chat_logs_query 未建立 Flask application context,与本次 PPT/海报链路无关。
2026-07-31 14:10:24 +08:00

373 lines
16 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.

"""海报业务逻辑服务。"""
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 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)
# 提取参考图元数据(不含 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,
}, ensure_ascii=False),
)
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()}