baodan/api/insurance/poster/service.py

257 lines
11 KiB
Python
Raw Normal View History

2026-07-23 15:04:16 +08:00
"""海报业务逻辑服务。"""
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()
# 按公司分组
company_map = {}
for p in PptProduct.query.filter(PptProduct.status == 1).all():
company_map[p.id] = p.company_id
result = []
for p in products:
result.append({
**p.to_dict(),
"companyName": self._get_company_name(p.company_id),
})
return {"code": 0, "data": result}
def _get_company_name(self, company_id: str) -> str:
company = PptCompany.query.get(company_id)
return company.display_name if company else company_id
def upload_case(self, user_id: str, product_id: str, file) -> dict:
"""上传计划书 PDF + 排队解析(异步)。"""
2026-07-23 15:04:16 +08:00
product = PptProduct.query.get(product_id)
if not product:
return {"code": 1002, "message": "产品不存在", "data": None}
# 文件安全校验SEC-P1-01
from insurance.utils.security import validate_pdf_upload
is_valid, err_msg = validate_pdf_upload(file)
if not is_valid:
return {"code": 4002, "message": err_msg, "data": None}
# 保存文件(使用持久化存储)
2026-07-23 15:04:16 +08:00
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")
2026-07-23 15:04:16 +08:00
os.makedirs(upload_dir, exist_ok=True)
filename = f"{safe_uid}_{uuid.uuid4().hex[:8]}.pdf"
filepath = os.path.join(upload_dir, filename)
file.save(filepath)
# 创建记录
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"
2026-07-23 15:04:16 +08:00
db.session.commit()
else:
2026-07-23 15:04:16 +08:00
record.parse_status = "failed"
db.session.commit()
return {"code": 5001, "message": "任务排队失败,请重试", "data": None}
2026-07-23 15:04:16 +08:00
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).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).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")
# 获取产品规则
product_rules = {}
if product_id:
product = PptProduct.query.get(product_id)
if product and product.manual_parsed_rules:
product_rules = json.loads(product.manual_parsed_rules)
# 获取客户数据(校验 case 所有权 — 防止越权读取他人客户数据)
2026-07-23 15:04:16 +08:00
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 case.confirmed_data:
2026-07-23 15:04:16 +08:00
customer_data = json.loads(case.confirmed_data)
if mode == "template":
template_id = data.get("templateId")
template = PosterCopyTemplate.query.get(template_id)
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:
"""生成海报图片(异步 — 排队后立即返回)。"""
2026-07-23 15:04:16 +08:00
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)
2026-07-23 15:04:16 +08:00
# 校验 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.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
2026-07-23 15:04:16 +08:00
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,
2026-07-23 15:04:16 +08:00
)
db.session.add(record)
db.session.commit()
# 启动后台生成任务
from insurance.poster.tasks import start_poster_generate_task
if start_poster_generate_task(current_app._get_current_object(), record.id, user_id, data):
record.task_status = "queued"
db.session.commit()
else:
record.task_status = "failed"
record.task_error = "任务排队失败"
db.session.commit()
return {"code": 5001, "message": "任务排队失败,请重试", "data": None}
2026-07-23 15:04:16 +08:00
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:
"""记录详情(含任务状态,用于轮询)。"""
2026-07-23 15:04:16 +08:00
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()}