baodan/api/insurance/poster/service.py
wsb1224 5f78598b3a input_revision 正确记录
task_service.py — create_task() 新增 input_revision 参数,存入 GenerationTask.input_revision
generation_task.py — 模型的静态 create_task 方法同步新增参数
ppt/routes.py — 三处调用(parse/generate/regenerate)均传入 session.draft_revision
poster/service.py — 海报生成传入 record.draft_revision
现在每个任务快照都记录了创建时的草稿版本号,不再一直是默认值 1
2. 海报任务进度实时同步 
celery_tasks.py — 在海报生成的每个阶段(preparing_data→20%、building_prompt→40%、requesting_image→60%、saving→85%)调用 sync_poster_progress() 同步到 PosterRecord.task_status/task_progress
任务领取时立即将 PosterRecord.task_status 设为 running
解决了"页面轮询海报记录时一直显示排队中"的问题
3. 后端测试覆盖 
test_task_state_sync.py — 新增 13 个测试用例:
sync_workspace_status 同步 cancelled/failed 到 PPT session(4 个)
sync_workspace_status 同步 cancelled/failed 到 PosterRecord(2 个)
旧任务跳过同步的安全检查(1 个)
cancel_task 状态更新和非 queued 拒绝(2 个)
mark_task_viewed 正确标记和权限校验(2 个)
input_revision 正确记录和默认值(2 个)
list_active_tasks 排除终态任务(1 个)
所有 13 个新测试 + 原有 10 个测试全部通过
4. useAutoSave TypeScript 类型修复 
useAutoSave.ts — getEndpoint 变量显式标注 () => string 类型,消除 TS2345 编译错误
2026-07-30 13:54:44 +08:00

363 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 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
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)
class PosterService:
"""海报业务逻辑。"""
def get_reviewed_products(self) -> dict:
"""获取已 reviewed 的产品列表(供选择)。"""
products = PptProduct.query.filter(
PptProduct.manual_parse_status == "reviewed",
PptProduct.status == 1,
).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,
).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 = "") -> dict:
"""上传计划书 PDF + 排队解析(异步)。"""
product = PptProduct.query.filter_by(
id=product_id, status=1, manual_parse_status="reviewed"
).first()
if not product:
return {"code": 1002, "message": "产品不存在", "data": None}
# 文件安全校验SEC-P1-01
from insurance.utils.security import prepare_pdf_upload
if not PptCompany.query.filter_by(id=product.company_id, status=1).first():
return {"code": 1002, "message": "产品所属保司不存在或已停用", "data": None}
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=product_id,
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).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")
product_id = data.get("productId")
use_masked_data = bool(data.get("useMaskedData"))
# 获取产品规则
product_rules = {}
product = None
if product_id:
product = PptProduct.query.filter_by(
id=product_id, status=1, manual_parse_status="reviewed"
).first()
if not product:
return {"code": 1002, "message": "产品不存在、未启用或尚未审核", "data": None}
if product and product.manual_parsed_rules:
try:
product_rules = json.loads(product.manual_parsed_rules)
except (json.JSONDecodeError, TypeError):
logger.warning("产品规则 JSON 解析失败: product_id=%s", product_id)
product_rules = {}
product_data = product.to_dict()
product_rules.setdefault("product_name", product.display_name)
product_rules.setdefault("coverage_period", product.coverage_period or "")
product_rules.setdefault("payment_period", product.payment_period or "")
product_rules.setdefault("insured_age_range", product.insured_age_range or "")
if not product_rules.get("features") and product_data.get("highlights"):
product_rules["features"] = [
{"title": item, "summary": ""}
if isinstance(item, str) else item
for item in product_data["highlights"]
]
company = PptCompany.query.filter_by(
id=product.company_id, status=1
).first()
if not company:
return {"code": 1002, "message": "产品所属保司不存在或已停用", "data": None}
product_rules.setdefault("company_name", company.display_name)
# 脱敏处理:替换产品规则中的产品名和保司名
if use_masked_data and product_rules:
from insurance.ppt.masking import (
apply_product_mask, apply_company_mask,
fallback_mask_name, mask_text, build_name_replacements,
)
if product:
product_dict = product.to_dict()
apply_product_mask(product_dict, True)
product_rules["product_name"] = product_dict.get("displayName", product_rules.get("product_name", ""))
# 替换规则中的产品名引用
masked_name = product_rules["product_name"]
real_name = product.display_name
if real_name and masked_name and real_name != masked_name:
for key in product_rules:
if isinstance(product_rules[key], str):
product_rules[key] = product_rules[key].replace(real_name, masked_name)
# 也替换保司名
company = PptCompany.query.get(product.company_id)
if company:
company_dict = company.to_dict()
apply_company_mask(company_dict, True)
masked_company = company_dict.get("displayName", "")
real_company = company.display_name
if real_company and masked_company and real_company != masked_company:
for key in product_rules:
if isinstance(product_rules[key], str):
product_rules[key] = product_rules[key].replace(real_company, masked_company)
# 获取客户数据(校验 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 product_id and case.product_id != product_id:
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
).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")
product_id = data.get("productId")
reference_image = data.get("referenceImage")
force_regenerate = data.get("regenerate", False)
use_masked_data = bool(data.get("useMaskedData"))
template = PosterTemplate.query.filter_by(id=template_id, status=1).first()
if not template:
return {"code": 1002, "message": "海报模板不存在或已停用", "data": None}
product = PptProduct.query.filter_by(
id=product_id, status=1, manual_parse_status="reviewed"
).first()
if not product:
return {"code": 1002, "message": "产品不存在、未启用或尚未审核", "data": None}
company = PptCompany.query.filter_by(id=product.company_id, status=1).first()
if not company:
return {"code": 1002, "message": "产品所属保司不存在或已停用", "data": None}
# 校验 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 case.product_id != product_id:
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=product_id,
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({"useMaskedData": use_masked_data}, ensure_ascii=False),
)
db.session.add(record)
db.session.commit()
# 启动后台生成任务Celery
from insurance.generation import task_service
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=data,
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()}