255 lines
10 KiB
Python
255 lines
10 KiB
Python
|
|
"""海报业务逻辑服务。"""
|
|||
|
|
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 + 触发解析。"""
|
|||
|
|
product = PptProduct.query.get(product_id)
|
|||
|
|
if not product:
|
|||
|
|
return {"code": 1002, "message": "产品不存在", "data": None}
|
|||
|
|
|
|||
|
|
# 保存文件
|
|||
|
|
safe_uid = _safe_user_id(user_id)
|
|||
|
|
upload_dir = os.path.join(current_app.config.get("UPLOAD_FOLDER", "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)
|
|||
|
|
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()
|
|||
|
|
|
|||
|
|
# 异步解析(同步执行,后续可改为 Celery)
|
|||
|
|
try:
|
|||
|
|
import asyncio
|
|||
|
|
from insurance.ppt.extraction import ExtractionOrchestrator
|
|||
|
|
orchestrator = ExtractionOrchestrator(use_cache=False)
|
|||
|
|
parsed = _run_async(orchestrator.extract_for_poster(filepath))
|
|||
|
|
record.parsed_data = json.dumps(parsed, ensure_ascii=False)
|
|||
|
|
record.parse_status = "parsed"
|
|||
|
|
db.session.commit()
|
|||
|
|
except Exception as e:
|
|||
|
|
logger.warning(f"计划书解析失败: {e}")
|
|||
|
|
record.parse_status = "failed"
|
|||
|
|
db.session.commit()
|
|||
|
|
|
|||
|
|
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)
|
|||
|
|
|
|||
|
|
# 获取客户数据
|
|||
|
|
customer_data = {}
|
|||
|
|
if case_upload_id:
|
|||
|
|
case = PosterCaseUpload.query.get(case_upload_id)
|
|||
|
|
if case and case.confirmed_data:
|
|||
|
|
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:
|
|||
|
|
"""生成海报图片。"""
|
|||
|
|
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")
|
|||
|
|
|
|||
|
|
# 获取模板
|
|||
|
|
poster_template = PosterTemplate.query.get(template_id) if template_id else None
|
|||
|
|
product = PptProduct.query.get(product_id) if product_id else None
|
|||
|
|
company = PptCompany.query.get(product.company_id) if product else None
|
|||
|
|
|
|||
|
|
# 组装 prompt
|
|||
|
|
from insurance.poster.image_generator import PosterImageGenerator
|
|||
|
|
generator = PosterImageGenerator()
|
|||
|
|
prompt = generator.build_prompt(
|
|||
|
|
template=poster_template.to_dict() if poster_template else None,
|
|||
|
|
product=product.to_dict() if product else None,
|
|||
|
|
company={"displayName": company.display_name} if company else None,
|
|||
|
|
copy=copy_content,
|
|||
|
|
size=size,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
# 生成图片
|
|||
|
|
try:
|
|||
|
|
image_bytes = generator.generate(prompt, size=size, reference_image=reference_image)
|
|||
|
|
except Exception as e:
|
|||
|
|
logger.warning(f"GPT image API 失败,使用降级方案: {e}")
|
|||
|
|
from insurance.poster.image_generator import generate_fallback
|
|||
|
|
image_bytes = generate_fallback(copy_content, size=size)
|
|||
|
|
|
|||
|
|
# 保存文件
|
|||
|
|
output_dir = os.path.join(current_app.config.get("UPLOAD_FOLDER", "uploads"), "posters")
|
|||
|
|
os.makedirs(output_dir, exist_ok=True)
|
|||
|
|
safe_uid = _safe_user_id(user_id)
|
|||
|
|
filename = f"poster_{safe_uid}_{uuid.uuid4().hex[:8]}.png"
|
|||
|
|
filepath = os.path.join(output_dir, filename)
|
|||
|
|
with open(filepath, "wb") as f:
|
|||
|
|
f.write(image_bytes)
|
|||
|
|
|
|||
|
|
# 创建记录(含合规留痕字段)
|
|||
|
|
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_url=filepath,
|
|||
|
|
export_format="png",
|
|||
|
|
export_size=size,
|
|||
|
|
reference_image_used=reference_image,
|
|||
|
|
prompt_used=prompt[:2000] if prompt else None,
|
|||
|
|
)
|
|||
|
|
db.session.add(record)
|
|||
|
|
db.session.commit()
|
|||
|
|
|
|||
|
|
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()}
|