685 lines
30 KiB
Python
685 lines
30 KiB
Python
"""PPT/海报管理后台服务。"""
|
||
import json
|
||
import logging
|
||
import os
|
||
import uuid
|
||
from flask import current_app, request
|
||
from insurance.db.compat import db
|
||
from insurance.models.ppt_config import PptCompany, PptProduct, PptTemplate
|
||
from insurance.models.ppt_history import PptHistory
|
||
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.models.system_setting import SystemSetting
|
||
from sqlalchemy import or_
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
def _safe_page_params(params: dict) -> tuple[int, int]:
|
||
"""安全获取分页参数,防止负数和超大值。"""
|
||
page = max(1, params.get("page", 1))
|
||
page_size = min(100, max(1, params.get("page_size", 20)))
|
||
return page, page_size
|
||
|
||
|
||
class PptAdminService:
|
||
"""PPT/海报管理后台业务逻辑。"""
|
||
|
||
# ---- 保司管理 ----
|
||
|
||
def list_companies(self, params: dict) -> dict:
|
||
query = db.session.query(PptCompany)
|
||
if params.get("status") is not None:
|
||
query = query.filter(PptCompany.status == params["status"])
|
||
if params.get("keyword"):
|
||
kw = f"%{params['keyword']}%"
|
||
query = query.filter(or_(
|
||
PptCompany.display_name.ilike(kw),
|
||
PptCompany.name_zh.ilike(kw),
|
||
PptCompany.name_en.ilike(kw),
|
||
))
|
||
query = query.order_by(PptCompany.sort_order.asc(), PptCompany.id.asc())
|
||
page, page_size = _safe_page_params(params)
|
||
total = query.count()
|
||
items = query.offset((page - 1) * page_size).limit(page_size).all()
|
||
return {
|
||
"code": 0,
|
||
"data": {
|
||
"total": total,
|
||
"items": [c.to_dict() for c in items],
|
||
},
|
||
}
|
||
|
||
def create_company(self, data: dict) -> dict:
|
||
if not data.get("id") or not data.get("displayName"):
|
||
return {"code": 1001, "message": "id 和 displayName 必填", "data": None}
|
||
if PptCompany.query.get(data["id"]):
|
||
return {"code": 1001, "message": "公司 ID 已存在", "data": None}
|
||
company = PptCompany(
|
||
id=data["id"],
|
||
display_name=data["displayName"],
|
||
aliases_json=json.dumps(data.get("aliases", []), ensure_ascii=False),
|
||
name_zh=data.get("nameZh"),
|
||
name_en=data.get("nameEn"),
|
||
short_en=data.get("shortEn"),
|
||
logo_url=data.get("logoUrl"),
|
||
company_intro=data.get("companyIntro"),
|
||
status=data.get("status", 1),
|
||
sort_order=data.get("sortOrder", 0),
|
||
)
|
||
db.session.add(company)
|
||
db.session.commit()
|
||
return {"code": 0, "data": company.to_dict()}
|
||
|
||
def update_company(self, company_id: str, data: dict) -> dict:
|
||
company = PptCompany.query.get(company_id)
|
||
if not company:
|
||
return {"code": 1002, "message": "公司不存在", "data": None}
|
||
if "displayName" in data:
|
||
company.display_name = data["displayName"]
|
||
if "aliases" in data:
|
||
company.aliases_json = json.dumps(data["aliases"], ensure_ascii=False)
|
||
if "nameZh" in data:
|
||
company.name_zh = data["nameZh"]
|
||
if "nameEn" in data:
|
||
company.name_en = data["nameEn"]
|
||
if "shortEn" in data:
|
||
company.short_en = data["shortEn"]
|
||
if "logoUrl" in data:
|
||
company.logo_url = data["logoUrl"]
|
||
if "companyIntro" in data:
|
||
company.company_intro = data["companyIntro"]
|
||
if "status" in data:
|
||
company.status = data["status"]
|
||
if "sortOrder" in data:
|
||
company.sort_order = data["sortOrder"]
|
||
db.session.commit()
|
||
return {"code": 0, "data": company.to_dict()}
|
||
|
||
def update_company_status(self, company_id: str, data: dict) -> dict:
|
||
company = PptCompany.query.get(company_id)
|
||
if not company:
|
||
return {"code": 1002, "message": "公司不存在", "data": None}
|
||
company.status = data.get("status", 1)
|
||
db.session.commit()
|
||
return {"code": 0, "data": company.to_dict()}
|
||
|
||
# ---- 产品管理 ----
|
||
|
||
def list_products(self, params: dict) -> dict:
|
||
query = db.session.query(PptProduct)
|
||
if params.get("company_id"):
|
||
query = query.filter(PptProduct.company_id == params["company_id"])
|
||
if params.get("plan_type"):
|
||
query = query.filter(PptProduct.plan_type == params["plan_type"])
|
||
if params.get("status") is not None:
|
||
query = query.filter(PptProduct.status == params["status"])
|
||
query = query.order_by(PptProduct.sort_order.asc(), PptProduct.id.asc())
|
||
page, page_size = _safe_page_params(params)
|
||
total = query.count()
|
||
items = query.offset((page - 1) * page_size).limit(page_size).all()
|
||
return {
|
||
"code": 0,
|
||
"data": {
|
||
"total": total,
|
||
"items": [p.to_dict() for p in items],
|
||
},
|
||
}
|
||
|
||
def create_product(self, data: dict) -> dict:
|
||
if not data.get("id") or not data.get("displayName"):
|
||
return {"code": 1001, "message": "id 和 displayName 必填", "data": None}
|
||
if PptProduct.query.get(data["id"]):
|
||
return {"code": 1001, "message": "产品 ID 已存在", "data": None}
|
||
product = PptProduct(
|
||
id=data["id"],
|
||
company_id=data.get("companyId", ""),
|
||
plan_type=data.get("planType", "savings"),
|
||
display_name=data["displayName"],
|
||
aliases_json=json.dumps(data.get("aliases", []), ensure_ascii=False),
|
||
required_modules_json=json.dumps(data.get("requiredModules", []), ensure_ascii=False),
|
||
product_code=data.get("productCode"),
|
||
product_type=data.get("productType"),
|
||
coverage_period=data.get("coveragePeriod"),
|
||
payment_period=data.get("paymentPeriod"),
|
||
insured_age_range=data.get("insuredAgeRange"),
|
||
waiting_period=data.get("waitingPeriod"),
|
||
highlights=json.dumps(data.get("highlights", []), ensure_ascii=False) if data.get("highlights") else None,
|
||
extra_fields=json.dumps(data.get("extraFields", {}), ensure_ascii=False) if data.get("extraFields") else None,
|
||
status=data.get("status", 1),
|
||
sort_order=data.get("sortOrder", 0),
|
||
)
|
||
db.session.add(product)
|
||
db.session.commit()
|
||
return {"code": 0, "data": product.to_dict()}
|
||
|
||
def update_product(self, product_id: str, data: dict) -> dict:
|
||
product = PptProduct.query.get(product_id)
|
||
if not product:
|
||
return {"code": 1002, "message": "产品不存在", "data": None}
|
||
simple_fields = {
|
||
"displayName": "display_name", "companyId": "company_id",
|
||
"planType": "plan_type", "productCode": "product_code",
|
||
"productType": "product_type", "coveragePeriod": "coverage_period",
|
||
"paymentPeriod": "payment_period", "insuredAgeRange": "insured_age_range",
|
||
"waitingPeriod": "waiting_period", "status": "status", "sortOrder": "sort_order",
|
||
}
|
||
for key, attr in simple_fields.items():
|
||
if key in data:
|
||
setattr(product, attr, data[key])
|
||
json_fields = {
|
||
"aliases": "aliases_json", "requiredModules": "required_modules_json",
|
||
"highlights": "highlights", "extraFields": "extra_fields",
|
||
}
|
||
for key, attr in json_fields.items():
|
||
if key in data:
|
||
setattr(product, attr, json.dumps(data[key], ensure_ascii=False) if data[key] else None)
|
||
db.session.commit()
|
||
return {"code": 0, "data": product.to_dict()}
|
||
|
||
def update_product_status(self, product_id: str, data: dict) -> dict:
|
||
product = PptProduct.query.get(product_id)
|
||
if not product:
|
||
return {"code": 1002, "message": "产品不存在", "data": None}
|
||
product.status = data.get("status", 1)
|
||
db.session.commit()
|
||
return {"code": 0, "data": product.to_dict()}
|
||
|
||
# ---- 产品小册子 ----
|
||
|
||
def upload_manual(self, product_id: str, file) -> dict:
|
||
product = PptProduct.query.get(product_id)
|
||
if not product:
|
||
return {"code": 1002, "message": "产品不存在", "data": None}
|
||
if not file.filename or not file.filename.lower().endswith(".pdf"):
|
||
return {"code": 1003, "message": "仅支持 PDF 文件", "data": None}
|
||
upload_dir = os.path.join(current_app.config.get("UPLOAD_FOLDER", "uploads"), "manuals")
|
||
os.makedirs(upload_dir, exist_ok=True)
|
||
filename = f"{product_id}_{uuid.uuid4().hex[:8]}.pdf"
|
||
filepath = os.path.join(upload_dir, filename)
|
||
file.save(filepath)
|
||
product.manual_file_url = filepath
|
||
product.manual_parse_status = "pending"
|
||
db.session.commit()
|
||
return {"code": 0, "data": {"fileUrl": filepath}}
|
||
|
||
def parse_manual(self, product_id: str) -> dict:
|
||
product = PptProduct.query.get(product_id)
|
||
if not product:
|
||
return {"code": 1002, "message": "产品不存在", "data": None}
|
||
if not product.manual_file_url:
|
||
return {"code": 1003, "message": "请先上传小册子", "data": None}
|
||
if not os.path.exists(product.manual_file_url):
|
||
return {"code": 1004, "message": "小册子文件不存在", "data": None}
|
||
product.manual_parse_status = "parsing"
|
||
db.session.commit()
|
||
# 调用 LLM 解析 PDF
|
||
try:
|
||
import asyncio
|
||
from insurance.poster.manual_parser import parse_manual_pdf
|
||
from insurance.poster.service import _run_async
|
||
result = _run_async(parse_manual_pdf(product.manual_file_url))
|
||
product.manual_parsed_rules = json.dumps(result, ensure_ascii=False)
|
||
product.manual_parse_status = "parsed"
|
||
db.session.commit()
|
||
return {"code": 0, "data": {"status": "parsed", "rules": result}}
|
||
except Exception as e:
|
||
logger.warning(f"小册子解析失败: {e}")
|
||
product.manual_parse_status = "none"
|
||
db.session.commit()
|
||
return {"code": 1005, "message": f"解析失败: {e}", "data": None}
|
||
|
||
def review_manual(self, product_id: str, data: dict) -> dict:
|
||
product = PptProduct.query.get(product_id)
|
||
if not product:
|
||
return {"code": 1002, "message": "产品不存在", "data": None}
|
||
if product.manual_parse_status not in ("parsed", "reviewed"):
|
||
return {"code": 1003, "message": "产品尚未解析,无法核对", "data": None}
|
||
if data.get("rules"):
|
||
product.manual_parsed_rules = json.dumps(data["rules"], ensure_ascii=False)
|
||
product.manual_parse_status = "reviewed"
|
||
product.manual_reviewed_by = data.get("reviewedBy", "")
|
||
from datetime import datetime
|
||
product.manual_reviewed_at = datetime.now()
|
||
db.session.commit()
|
||
return {"code": 0, "data": product.to_dict()}
|
||
|
||
# ---- PPT 模板管理 ----
|
||
|
||
def list_templates(self, params: dict) -> dict:
|
||
query = db.session.query(PptTemplate)
|
||
if params.get("plan_type"):
|
||
query = query.filter(PptTemplate.plan_type == params["plan_type"])
|
||
if params.get("status") is not None:
|
||
query = query.filter(PptTemplate.status == params["status"])
|
||
page, page_size = _safe_page_params(params)
|
||
total = query.count()
|
||
items = query.offset((page - 1) * page_size).limit(page_size).all()
|
||
return {
|
||
"code": 0,
|
||
"data": {
|
||
"total": total,
|
||
"items": [t.to_dict() for t in items],
|
||
},
|
||
}
|
||
|
||
def create_template(self, data: dict) -> dict:
|
||
if not data.get("id"):
|
||
return {"code": 1001, "message": "id 必填", "data": None}
|
||
if PptTemplate.query.get(data["id"]):
|
||
return {"code": 1001, "message": "模板 ID 已存在", "data": None}
|
||
template = PptTemplate(
|
||
id=data["id"],
|
||
plan_type=data.get("planType", "savings"),
|
||
style_preset=data.get("stylePreset", "broker"),
|
||
name=data.get("name"),
|
||
scenario_tag=data.get("scenarioTag"),
|
||
preview_image=data.get("previewImage"),
|
||
applicable_company_ids=json.dumps(data.get("applicableCompanyIds", []), ensure_ascii=False) if data.get("applicableCompanyIds") else None,
|
||
applicable_product_ids=json.dumps(data.get("applicableProductIds", []), ensure_ascii=False) if data.get("applicableProductIds") else None,
|
||
slides_config_json=json.dumps(data.get("slidesConfig", []), ensure_ascii=False) if data.get("slidesConfig") else None,
|
||
status=data.get("status", 1),
|
||
)
|
||
db.session.add(template)
|
||
db.session.commit()
|
||
return {"code": 0, "data": template.to_dict()}
|
||
|
||
def update_template(self, template_id: str, data: dict) -> dict:
|
||
template = PptTemplate.query.get(template_id)
|
||
if not template:
|
||
return {"code": 1002, "message": "模板不存在", "data": None}
|
||
for key, attr in [("name", "name"), ("scenarioTag", "scenario_tag"), ("previewImage", "preview_image")]:
|
||
if key in data:
|
||
setattr(template, attr, data[key])
|
||
if "applicableCompanyIds" in data:
|
||
template.applicable_company_ids = json.dumps(data["applicableCompanyIds"], ensure_ascii=False) if data["applicableCompanyIds"] else None
|
||
if "applicableProductIds" in data:
|
||
template.applicable_product_ids = json.dumps(data["applicableProductIds"], ensure_ascii=False) if data["applicableProductIds"] else None
|
||
if "slidesConfig" in data:
|
||
template.slides_config_json = json.dumps(data["slidesConfig"], ensure_ascii=False) if data["slidesConfig"] else None
|
||
if "status" in data:
|
||
template.status = data["status"]
|
||
db.session.commit()
|
||
return {"code": 0, "data": template.to_dict()}
|
||
|
||
def update_template_status(self, template_id: str, data: dict) -> dict:
|
||
template = PptTemplate.query.get(template_id)
|
||
if not template:
|
||
return {"code": 1002, "message": "模板不存在", "data": None}
|
||
template.status = data.get("status", 1)
|
||
db.session.commit()
|
||
return {"code": 0, "data": template.to_dict()}
|
||
|
||
# ---- 海报模板管理 ----
|
||
|
||
def list_poster_templates(self, params: dict) -> dict:
|
||
query = db.session.query(PosterTemplate)
|
||
if params.get("status") is not None:
|
||
query = query.filter(PosterTemplate.status == params["status"])
|
||
page, page_size = _safe_page_params(params)
|
||
total = query.count()
|
||
items = query.offset((page - 1) * page_size).limit(page_size).all()
|
||
return {
|
||
"code": 0,
|
||
"data": {
|
||
"total": total,
|
||
"items": [t.to_dict() for t in items],
|
||
},
|
||
}
|
||
|
||
def create_poster_template(self, data: dict) -> dict:
|
||
if not data.get("name") or not data.get("styleDescription"):
|
||
return {"code": 1001, "message": "name 和 styleDescription 必填", "data": None}
|
||
template = PosterTemplate(
|
||
name=data["name"],
|
||
scenario_tag=data.get("scenarioTag"),
|
||
style_description=data["styleDescription"],
|
||
color_scheme=json.dumps(data["colorScheme"], ensure_ascii=False) if data.get("colorScheme") else None,
|
||
reference_image=data.get("referenceImage"),
|
||
preview_image=data.get("previewImage"),
|
||
status=data.get("status", 1),
|
||
)
|
||
db.session.add(template)
|
||
db.session.commit()
|
||
return {"code": 0, "data": template.to_dict()}
|
||
|
||
def update_poster_template(self, template_id: int, data: dict) -> dict:
|
||
template = PosterTemplate.query.get(template_id)
|
||
if not template:
|
||
return {"code": 1002, "message": "模板不存在", "data": None}
|
||
for key, attr in [("name", "name"), ("scenarioTag", "scenario_tag"),
|
||
("styleDescription", "style_description"),
|
||
("referenceImage", "reference_image"), ("previewImage", "preview_image")]:
|
||
if key in data:
|
||
setattr(template, attr, data[key])
|
||
if "colorScheme" in data:
|
||
template.color_scheme = json.dumps(data["colorScheme"], ensure_ascii=False) if data["colorScheme"] else None
|
||
if "status" in data:
|
||
template.status = data["status"]
|
||
db.session.commit()
|
||
return {"code": 0, "data": template.to_dict()}
|
||
|
||
def update_poster_template_status(self, template_id: int, data: dict) -> dict:
|
||
template = PosterTemplate.query.get(template_id)
|
||
if not template:
|
||
return {"code": 1002, "message": "模板不存在", "data": None}
|
||
template.status = data.get("status", 1)
|
||
db.session.commit()
|
||
return {"code": 0, "data": template.to_dict()}
|
||
|
||
# ---- 文案模板管理 ----
|
||
|
||
def list_copy_templates(self, params: dict) -> dict:
|
||
query = db.session.query(PosterCopyTemplate)
|
||
if params.get("status") is not None:
|
||
query = query.filter(PosterCopyTemplate.status == params["status"])
|
||
page, page_size = _safe_page_params(params)
|
||
total = query.count()
|
||
items = query.offset((page - 1) * page_size).limit(page_size).all()
|
||
return {
|
||
"code": 0,
|
||
"data": {
|
||
"total": total,
|
||
"items": [t.to_dict() for t in items],
|
||
},
|
||
}
|
||
|
||
def create_copy_template(self, data: dict) -> dict:
|
||
if not data.get("name") or not data.get("content"):
|
||
return {"code": 1001, "message": "name 和 content 必填", "data": None}
|
||
template = PosterCopyTemplate(
|
||
name=data["name"],
|
||
scenario_tag=data.get("scenarioTag"),
|
||
content=data["content"],
|
||
variables=json.dumps(data.get("variables", []), ensure_ascii=False) if data.get("variables") else None,
|
||
status=data.get("status", 1),
|
||
)
|
||
db.session.add(template)
|
||
db.session.commit()
|
||
return {"code": 0, "data": template.to_dict()}
|
||
|
||
def update_copy_template(self, template_id: int, data: dict) -> dict:
|
||
template = PosterCopyTemplate.query.get(template_id)
|
||
if not template:
|
||
return {"code": 1002, "message": "模板不存在", "data": None}
|
||
for key, attr in [("name", "name"), ("scenarioTag", "scenario_tag"), ("content", "content")]:
|
||
if key in data:
|
||
setattr(template, attr, data[key])
|
||
if "variables" in data:
|
||
template.variables = json.dumps(data["variables"], ensure_ascii=False) if data["variables"] else None
|
||
if "status" in data:
|
||
template.status = data["status"]
|
||
db.session.commit()
|
||
return {"code": 0, "data": template.to_dict()}
|
||
|
||
def delete_copy_template(self, template_id: int) -> dict:
|
||
template = PosterCopyTemplate.query.get(template_id)
|
||
if not template:
|
||
return {"code": 1002, "message": "模板不存在", "data": None}
|
||
db.session.delete(template)
|
||
db.session.commit()
|
||
return {"code": 0, "data": None}
|
||
|
||
# ---- 历史记录管理 ----
|
||
|
||
def list_history(self, params: dict) -> dict:
|
||
query = db.session.query(PptHistory)
|
||
if params.get("user_id"):
|
||
query = query.filter(PptHistory.user_id == params["user_id"])
|
||
if params.get("company_id"):
|
||
query = query.filter(PptHistory.company_id == params["company_id"])
|
||
if params.get("action_type"):
|
||
query = query.filter(PptHistory.action_type == params["action_type"])
|
||
query = query.order_by(PptHistory.created_at.desc())
|
||
page, page_size = _safe_page_params(params)
|
||
total = query.count()
|
||
items = query.offset((page - 1) * page_size).limit(page_size).all()
|
||
return {
|
||
"code": 0,
|
||
"data": {
|
||
"total": total,
|
||
"items": [h.to_dict() for h in items],
|
||
},
|
||
}
|
||
|
||
def delete_history(self, history_id: int) -> dict:
|
||
record = PptHistory.query.get(history_id)
|
||
if not record:
|
||
return {"code": 1002, "message": "记录不存在", "data": None}
|
||
db.session.delete(record)
|
||
db.session.commit()
|
||
return {"code": 0, "data": None}
|
||
|
||
def export_history_csv(self, params: dict):
|
||
"""导出历史记录为 CSV。"""
|
||
import csv
|
||
import io
|
||
query = db.session.query(PptHistory)
|
||
if params.get("user_id"):
|
||
query = query.filter(PptHistory.user_id == params["user_id"])
|
||
if params.get("company_id"):
|
||
query = query.filter(PptHistory.company_id == params["company_id"])
|
||
if params.get("action_type"):
|
||
query = query.filter(PptHistory.action_type == params["action_type"])
|
||
query = query.order_by(PptHistory.created_at.desc())
|
||
items = query.limit(10000).all() # 限制最大导出行数
|
||
|
||
output = io.StringIO()
|
||
output.write("") # UTF-8 BOM,确保 Excel 正确显示中文
|
||
writer = csv.writer(output)
|
||
writer.writerow(["ID", "用户ID", "会话ID", "操作类型", "保司ID", "产品ID", "文件地址", "时间"])
|
||
for h in items:
|
||
writer.writerow([
|
||
h.id, h.user_id, h.session_id or "", h.action_type,
|
||
h.company_id or "", h.product_id or "",
|
||
h.file_url or "", str(h.created_at) if h.created_at else "",
|
||
])
|
||
return output.getvalue()
|
||
|
||
# ---- 系统配置 ----
|
||
|
||
# 内置常用模型列表(Dify 不可用时的降级方案)
|
||
_BUILTIN_MODELS = [
|
||
{"provider": "deepseek", "model": "deepseek-chat", "label": "DeepSeek Chat"},
|
||
{"provider": "deepseek", "model": "deepseek-reasoner", "label": "DeepSeek Reasoner"},
|
||
{"provider": "minimax", "model": "MiniMax-2.7-Flash", "label": "MiniMax 2.7 Flash"},
|
||
{"provider": "gemini", "model": "gemini-2.5-flash", "label": "Gemini 2.5 Flash"},
|
||
{"provider": "gemini", "model": "gemini-2.5-pro", "label": "Gemini 2.5 Pro"},
|
||
{"provider": "openai", "model": "gpt-4o", "label": "GPT-4o"},
|
||
{"provider": "openai", "model": "gpt-4o-mini", "label": "GPT-4o Mini"},
|
||
{"provider": "openai", "model": "gpt-image-1", "label": "GPT Image 1"},
|
||
]
|
||
|
||
def get_available_models(self) -> dict:
|
||
"""获取可用 LLM 模型列表,优先从 Dify 获取,失败时返回内置列表。"""
|
||
# 尝试从 Dify 获取
|
||
dify_models = self._fetch_dify_models()
|
||
if dify_models is not None:
|
||
return {"code": 0, "data": {"models": dify_models, "source": "dify"}}
|
||
return {"code": 0, "data": {"models": self._BUILTIN_MODELS, "source": "builtin"}}
|
||
|
||
# 模型名前缀 → 品牌名映射(Dify 用 openai_api_compatible 统一接口,需要从模型名推断品牌)
|
||
_MODEL_BRAND_MAP = {
|
||
"deepseek": "DeepSeek",
|
||
"gpt": "OpenAI",
|
||
"o1": "OpenAI",
|
||
"o3": "OpenAI",
|
||
"o4": "OpenAI",
|
||
"claude": "Anthropic",
|
||
"gemini": "Google",
|
||
"qwen": "通义千问",
|
||
"glm": "智谱",
|
||
"cogview": "智谱",
|
||
"MiniMax": "MiniMax",
|
||
"moonshot": "月之暗面",
|
||
"doubao": "豆包",
|
||
}
|
||
|
||
def _infer_brand(self, model_name: str) -> str:
|
||
"""从模型名推断品牌名。"""
|
||
for prefix, brand in self._MODEL_BRAND_MAP.items():
|
||
if model_name.lower().startswith(prefix.lower()):
|
||
return brand
|
||
return "其他"
|
||
|
||
def _fetch_dify_models(self) -> list | None:
|
||
"""从 Dify 数据库直接查询已配置的 LLM 模型列表。"""
|
||
try:
|
||
rows = db.session.execute(db.text(
|
||
"SELECT DISTINCT pm.model_name, pm.provider_name "
|
||
"FROM provider_models pm "
|
||
"WHERE pm.model_type = 'llm' AND pm.is_valid = true "
|
||
"ORDER BY pm.provider_name, pm.model_name"
|
||
)).fetchall()
|
||
if not rows:
|
||
return None
|
||
models = []
|
||
for row in rows:
|
||
model_name = row[0]
|
||
provider_name = row[1]
|
||
brand = self._infer_brand(model_name)
|
||
models.append({
|
||
"provider": brand,
|
||
"model": model_name,
|
||
"label": f"{brand} / {model_name}",
|
||
})
|
||
return models
|
||
except Exception as e:
|
||
logger.warning(f"从 Dify 数据库获取模型列表失败: {e}")
|
||
return None
|
||
|
||
# 需要掩码处理的敏感 key(包含 api_key 或 secret 等)
|
||
_SENSITIVE_KEYS = {"api_key", "secret", "password", "token"}
|
||
|
||
@staticmethod
|
||
def _mask_value(key: str, value: str) -> str:
|
||
"""对敏感配置项进行掩码,不返回完整密钥。"""
|
||
if not value:
|
||
return value
|
||
key_lower = key.lower()
|
||
if any(s in key_lower for s in PptAdminService._SENSITIVE_KEYS):
|
||
if len(value) <= 4:
|
||
return "****"
|
||
return f"{value[:2]}****{value[-2:]}"
|
||
return value
|
||
|
||
def get_settings(self) -> dict:
|
||
settings = SystemSetting.query.all()
|
||
data = {}
|
||
for s in settings:
|
||
data[s.key] = s.value
|
||
return {
|
||
"code": 0,
|
||
"data": data,
|
||
}
|
||
|
||
# 允许通过 API 写入的设置键白名单
|
||
_ALLOWED_SETTING_KEYS = {
|
||
"ppt_llm_provider", "ppt_llm_model", "ppt_llm_api_key", "ppt_llm_base_url", "ppt_llm_timeout_ms",
|
||
"poster_llm_provider", "poster_llm_model", "poster_llm_api_key", "poster_llm_base_url",
|
||
"poster_image_provider", "poster_image_model", "poster_image_api_key", "poster_image_base_url",
|
||
"poster_history_retention_days", "ppt_history_retention_days",
|
||
"dify_workspace_api_key",
|
||
}
|
||
|
||
@staticmethod
|
||
def _is_masked_value(value: str) -> bool:
|
||
"""判断值是否为掩码格式(如 sk-****abcd)。"""
|
||
return isinstance(value, str) and "****" in value
|
||
|
||
@staticmethod
|
||
def _validate_settings(data: dict) -> str | None:
|
||
"""校验设置值的安全性,返回错误信息或 None。"""
|
||
from insurance.utils.security import is_safe_base_url
|
||
allowed = PptAdminService._ALLOWED_SETTING_KEYS
|
||
for key, value in data.items():
|
||
if key not in allowed:
|
||
return f"不允许的配置键: {key}"
|
||
if not isinstance(value, str):
|
||
continue
|
||
# 校验 Base URL 字段
|
||
if "base_url" in key.lower() and value.strip():
|
||
is_safe, err_msg = is_safe_base_url(value.strip())
|
||
if not is_safe:
|
||
return f"{key}: {err_msg}"
|
||
return None
|
||
|
||
def update_settings(self, data: dict, updated_by: str = "") -> dict:
|
||
# 安全校验(SEC-P1-02)
|
||
validation_error = self._validate_settings(data)
|
||
if validation_error:
|
||
return {"code": 1001, "message": validation_error, "data": None}
|
||
|
||
for key, value in data.items():
|
||
value_str = str(value) if value is not None else ""
|
||
setting = SystemSetting.query.filter_by(key=key).first()
|
||
if setting:
|
||
# 如果是掩码值,跳过更新(保留真实密钥)
|
||
if self._is_masked_value(value_str):
|
||
continue
|
||
setting.value = value_str
|
||
setting.updated_by = updated_by
|
||
else:
|
||
# 新建设置项,不允许掩码值
|
||
if self._is_masked_value(value_str):
|
||
continue
|
||
db.session.add(SystemSetting(key=key, value=value_str, updated_by=updated_by))
|
||
db.session.commit()
|
||
return {"code": 0, "data": None}
|
||
|
||
def sync_models(self, provider: str, api_key: str, base_url: str = "") -> dict:
|
||
"""从供应商 API 拉取可用模型列表。"""
|
||
if not api_key:
|
||
return {"code": 1001, "message": "请填写 API Key", "data": None}
|
||
|
||
try:
|
||
models = self._fetch_provider_models(provider, api_key, base_url)
|
||
except Exception as e:
|
||
logger.warning(f"同步模型失败: {e}")
|
||
return {"code": 5001, "message": f"同步失败: {e}", "data": None}
|
||
|
||
if not models:
|
||
return {"code": 0, "data": {"models": [], "message": "未获取到模型"}}
|
||
|
||
return {"code": 0, "data": {"models": models}}
|
||
|
||
def _fetch_provider_models(self, provider: str, api_key: str, base_url: str) -> list:
|
||
"""根据供应商类型拉取模型列表。"""
|
||
import httpx
|
||
|
||
if provider == "deepseek":
|
||
url = "https://api.deepseek.com/v1/models"
|
||
headers = {"Authorization": f"Bearer {api_key}"}
|
||
resp = httpx.get(url, headers=headers, timeout=10)
|
||
resp.raise_for_status()
|
||
return [m["id"] for m in resp.json().get("data", [])]
|
||
|
||
elif provider == "gemini":
|
||
url = f"https://generativelanguage.googleapis.com/v1/models?key={api_key}"
|
||
resp = httpx.get(url, timeout=10)
|
||
resp.raise_for_status()
|
||
return [m["name"].split("/")[-1] for m in resp.json().get("models", [])]
|
||
|
||
elif provider == "minimax":
|
||
url = "https://api.minimax.chat/v1/models"
|
||
headers = {"Authorization": f"Bearer {api_key}"}
|
||
resp = httpx.get(url, headers=headers, timeout=10)
|
||
resp.raise_for_status()
|
||
data = resp.json()
|
||
models = data.get("models") or data.get("data", [])
|
||
return [m.get("id") or m.get("model", "") for m in models if m]
|
||
|
||
elif provider == "dify":
|
||
return [m["model"] for m in (self._fetch_dify_models() or [])]
|
||
|
||
else:
|
||
# 自定义供应商:OpenAI 兼容 /models 端点
|
||
if not base_url:
|
||
return []
|
||
url = f"{base_url.rstrip('/')}/models"
|
||
headers = {"Authorization": f"Bearer {api_key}"}
|
||
resp = httpx.get(url, headers=headers, timeout=10)
|
||
resp.raise_for_status()
|
||
return [m["id"] for m in resp.json().get("data", [])]
|