749 lines
29 KiB
Python
749 lines
29 KiB
Python
"""管理后台服务:用户、角色、LLM、Prompt 管理。"""
|
||
import bcrypt
|
||
import requests
|
||
import uuid
|
||
from flask import current_app
|
||
from insurance.db.compat import db
|
||
from insurance.models.wecom_group import WeComGroupConfig
|
||
from insurance.models.wecom_user import WeComUserMapping
|
||
from sqlalchemy import or_
|
||
|
||
|
||
class AdminService:
|
||
"""管理后台业务逻辑。"""
|
||
|
||
# ---- 用户管理 ----
|
||
|
||
def list_users(self, params: dict) -> dict:
|
||
query = db.session.query(WeComUserMapping)
|
||
if params.get("role"):
|
||
query = query.filter(WeComUserMapping.role == params["role"])
|
||
if params.get("department"):
|
||
query = query.filter(WeComUserMapping.department == params["department"])
|
||
|
||
total = query.count()
|
||
items = query.offset((params["page"] - 1) * params["page_size"]).limit(params["page_size"]).all()
|
||
|
||
# 检查每个用户是否已同步到 Dify
|
||
dify_users = set()
|
||
try:
|
||
result = db.session.execute(db.text("SELECT name FROM accounts"))
|
||
dify_users = {row[0] for row in result.fetchall()}
|
||
except Exception:
|
||
pass
|
||
|
||
return {
|
||
"code": 0,
|
||
"data": {
|
||
"total": total,
|
||
"items": [
|
||
{
|
||
"id": str(u.id),
|
||
"username": u.username,
|
||
"email": u.email,
|
||
"real_name": u.real_name,
|
||
"email_verified": u.email_verified,
|
||
"wecom_userid": u.wecom_userid,
|
||
"role": u.role,
|
||
"department": u.department,
|
||
"status": u.status,
|
||
"synced_to_dify": u.username in dify_users,
|
||
"created_at": str(u.created_at) if u.created_at else None,
|
||
"last_active_at": str(u.last_active_at) if u.last_active_at else None,
|
||
}
|
||
for u in items
|
||
],
|
||
},
|
||
}
|
||
|
||
def create_user(self, data: dict) -> dict:
|
||
username = (data.get("username") or "").strip()
|
||
email = (data.get("email") or "").strip().lower()
|
||
if username and db.session.query(WeComUserMapping).filter_by(username=username).first():
|
||
return {"code": 1001, "message": "用户名已存在", "data": None}
|
||
if email and db.session.query(WeComUserMapping).filter_by(email=email).first():
|
||
return {"code": 1001, "message": "邮箱已被使用", "data": None}
|
||
|
||
wecom_userid = (data.get("wecom_userid") or "").strip() or f"local_{uuid.uuid4().hex[:8]}"
|
||
internal_user_id = (data.get("internal_user_id") or "").strip() or f"manual_{wecom_userid}"
|
||
user = WeComUserMapping(
|
||
wecom_userid=wecom_userid,
|
||
internal_user_id=internal_user_id,
|
||
username=username,
|
||
email=email,
|
||
real_name=(data.get("real_name") or "").strip(),
|
||
email_verified="true" if data.get("email") else "false",
|
||
role=data.get("role", "sales"),
|
||
department=data.get("department", ""),
|
||
)
|
||
if data.get("password"):
|
||
user.password_hash = bcrypt.hashpw(data["password"].encode(), bcrypt.gensalt()).decode()
|
||
db.session.add(user)
|
||
try:
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
raise
|
||
|
||
# 同步到 Dify accounts 表
|
||
if data.get("password"):
|
||
self._sync_user_to_dify(user.username, data["password"], user.email)
|
||
|
||
from insurance.utils.audit import log_operation
|
||
log_operation("system", "create", "user", str(user.id), {"username": user.username})
|
||
|
||
return {"code": 0, "message": "success", "data": {"id": str(user.id)}}
|
||
|
||
def _sync_user_to_dify(self, username: str, password: str, email: str = "") -> bool:
|
||
"""同步用户到 Dify accounts 表。"""
|
||
try:
|
||
import uuid as uuid_mod
|
||
from datetime import datetime, timezone
|
||
|
||
# 检查是否已存在
|
||
result = db.session.execute(
|
||
db.text("SELECT id FROM accounts WHERE name = :name"),
|
||
{"name": username}
|
||
)
|
||
if result.fetchone():
|
||
return True
|
||
|
||
# 创建 Dify 账号
|
||
account_id = str(uuid_mod.uuid4())
|
||
account_email = email or f"{username}@baodan.local"
|
||
now = datetime.now(timezone.utc)
|
||
|
||
# Dify 期望密码已经是 bcrypt 哈希后的值
|
||
password_hash = bcrypt.hashpw(password.encode(), bcrypt.gensalt()).decode()
|
||
db.session.execute(
|
||
db.text("""
|
||
INSERT INTO accounts (id, name, email, password, avatar, interface_language,
|
||
interface_theme, timezone, last_login_at, last_active_at,
|
||
status, created_at, updated_at)
|
||
VALUES (:id, :name, :email, :password, '', 'zh-Hans', 'light', 'Asia/Shanghai',
|
||
:now, :now, 'active', :now, :now)
|
||
"""),
|
||
{"id": account_id, "name": username, "email": account_email, "password": password_hash, "now": now}
|
||
)
|
||
|
||
# 关联到默认租户
|
||
tenant_result = db.session.execute(db.text("SELECT id FROM tenants LIMIT 1"))
|
||
tenant = tenant_result.fetchone()
|
||
|
||
if tenant:
|
||
db.session.execute(
|
||
db.text("""
|
||
INSERT INTO tenant_account_joins (id, tenant_id, account_id, role,
|
||
invited_by, created_at, updated_at)
|
||
VALUES (:id, :tenant_id, :account_id, 'normal', NULL, :now, :now)
|
||
"""),
|
||
{"id": str(uuid_mod.uuid4()), "tenant_id": str(tenant[0]), "account_id": account_id, "now": now}
|
||
)
|
||
|
||
db.session.commit()
|
||
return True
|
||
except Exception as e:
|
||
print(f"同步用户到 Dify 失败: {e}")
|
||
db.session.rollback()
|
||
return False
|
||
|
||
def update_user(self, user_id: str, data: dict) -> dict:
|
||
user = db.session.query(WeComUserMapping).filter_by(id=int(user_id)).first()
|
||
if not user:
|
||
return {"code": 1005, "message": "用户不存在", "data": None}
|
||
if "username" in data and data["username"] != user.username:
|
||
existing = db.session.query(WeComUserMapping).filter_by(username=data["username"]).first()
|
||
if existing:
|
||
return {"code": 1001, "message": "用户名已存在", "data": None}
|
||
if "email" in data:
|
||
new_email = (data["email"] or "").strip().lower()
|
||
if new_email and new_email != user.email:
|
||
existing = db.session.query(WeComUserMapping).filter_by(email=new_email).first()
|
||
if existing:
|
||
return {"code": 1001, "message": "邮箱已被使用", "data": None}
|
||
for key in ("username", "email", "real_name", "role", "department", "status"):
|
||
if key in data:
|
||
value = data[key]
|
||
if key == "email":
|
||
value = (value or "").strip().lower()
|
||
user.email_verified = "true" if value else "false"
|
||
elif key == "real_name":
|
||
value = (value or "").strip()
|
||
setattr(user, key, value)
|
||
try:
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
raise
|
||
|
||
# 同步状态到 Dify
|
||
if "status" in data:
|
||
self._sync_user_status_to_dify(user.username, data["status"])
|
||
|
||
return {"code": 0, "message": "success", "data": None}
|
||
|
||
# ---- 企微群配置 ----
|
||
|
||
def list_wecom_groups(self, params: dict) -> dict:
|
||
query = db.session.query(WeComGroupConfig)
|
||
if params.get("status"):
|
||
query = query.filter(WeComGroupConfig.status == params["status"])
|
||
if params.get("keyword"):
|
||
keyword = f"%{params['keyword']}%"
|
||
query = query.filter(
|
||
or_(
|
||
WeComGroupConfig.chatid.ilike(keyword),
|
||
WeComGroupConfig.name.ilike(keyword),
|
||
)
|
||
)
|
||
|
||
total = query.count()
|
||
items = (
|
||
query.order_by(WeComGroupConfig.updated_at.desc())
|
||
.offset((params["page"] - 1) * params["page_size"])
|
||
.limit(params["page_size"])
|
||
.all()
|
||
)
|
||
return {"code": 0, "data": {"total": total, "items": [g.to_dict() for g in items]}}
|
||
|
||
def update_wecom_group(self, group_id: str, data: dict) -> dict:
|
||
group = db.session.query(WeComGroupConfig).filter_by(id=int(group_id)).first()
|
||
if not group:
|
||
return {"code": 1005, "message": "企微群配置不存在", "data": None}
|
||
|
||
if "name" in data:
|
||
group.name = (data.get("name") or "").strip()
|
||
if "webhook_url" in data:
|
||
group.webhook_url = (data.get("webhook_url") or "").strip()
|
||
if "status" in data:
|
||
status = data.get("status") or "active"
|
||
if status not in ("active", "disabled"):
|
||
return {"code": 1001, "message": "群配置状态无效", "data": None}
|
||
group.status = status
|
||
|
||
try:
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
raise
|
||
|
||
return {"code": 0, "message": "success", "data": group.to_dict()}
|
||
|
||
def delete_user(self, user_id: str) -> dict:
|
||
user = db.session.query(WeComUserMapping).filter_by(id=int(user_id)).first()
|
||
if not user:
|
||
return {"code": 1005, "message": "用户不存在", "data": None}
|
||
user.status = "disabled"
|
||
try:
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
raise
|
||
|
||
# 同步停用状态到 Dify
|
||
self._sync_user_status_to_dify(user.username, "disabled")
|
||
|
||
from insurance.utils.audit import log_operation
|
||
log_operation("system", "disable", "user", user_id, {"username": user.username})
|
||
|
||
return {"code": 0, "message": "success", "data": None}
|
||
|
||
def _sync_user_status_to_dify(self, username: str, status: str) -> bool:
|
||
"""同步用户状态到 Dify。"""
|
||
try:
|
||
dify_status = "active" if status == "active" else "banned"
|
||
db.session.execute(
|
||
db.text("UPDATE accounts SET status = :status, updated_at = NOW() WHERE name = :name"),
|
||
{"status": dify_status, "name": username}
|
||
)
|
||
db.session.commit()
|
||
return True
|
||
except Exception as e:
|
||
print(f"同步用户状态到 Dify 失败: {e}")
|
||
db.session.rollback()
|
||
return False
|
||
|
||
def batch_import_users(self, file) -> dict:
|
||
"""批量导入用户(从Excel文件)。"""
|
||
import openpyxl
|
||
import uuid
|
||
|
||
try:
|
||
# 读取Excel文件
|
||
wb = openpyxl.load_workbook(file)
|
||
ws = wb.active
|
||
|
||
# 获取表头
|
||
headers = [cell.value for cell in ws[1]]
|
||
username_idx = headers.index("用户名") if "用户名" in headers else 0
|
||
password_idx = headers.index("密码") if "密码" in headers else 1
|
||
role_idx = headers.index("角色") if "角色" in headers else 2
|
||
department_idx = headers.index("部门") if "部门" in headers else 3
|
||
|
||
created = 0
|
||
skipped = 0
|
||
errors = []
|
||
|
||
for row_idx, row in enumerate(ws.iter_rows(min_row=2, values_only=True), start=2):
|
||
username = row[username_idx] if username_idx < len(row) else ""
|
||
password = row[password_idx] if password_idx < len(row) else ""
|
||
role = row[role_idx] if role_idx < len(row) else "sales"
|
||
department = row[department_idx] if department_idx < len(row) else ""
|
||
|
||
if not username or not password:
|
||
errors.append(f"第{row_idx}行:用户名或密码为空")
|
||
continue
|
||
|
||
# 检查用户名是否已存在
|
||
existing = db.session.query(WeComUserMapping).filter_by(username=str(username)).first()
|
||
if existing:
|
||
skipped += 1
|
||
continue
|
||
|
||
# 创建用户
|
||
password_hash = bcrypt.hashpw(str(password).encode(), bcrypt.gensalt()).decode()
|
||
mapping = WeComUserMapping(
|
||
wecom_userid=f"local_{uuid.uuid4().hex[:8]}",
|
||
internal_user_id=f"user_{uuid.uuid4().hex[:8]}",
|
||
username=str(username),
|
||
password_hash=password_hash,
|
||
role=str(role) if role else "sales",
|
||
department=str(department) if department else "",
|
||
status="active",
|
||
)
|
||
db.session.add(mapping)
|
||
created += 1
|
||
|
||
try:
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
raise
|
||
|
||
from insurance.utils.audit import log_operation
|
||
log_operation("system", "batch_import", "user", "", {
|
||
"created": created,
|
||
"skipped": skipped,
|
||
"errors": errors[:10], # 最多记录10个错误
|
||
})
|
||
|
||
return {
|
||
"code": 0,
|
||
"message": f"导入完成:成功 {created} 个,跳过 {skipped} 个",
|
||
"data": {
|
||
"created": created,
|
||
"skipped": skipped,
|
||
"errors": errors,
|
||
},
|
||
}
|
||
|
||
except Exception as e:
|
||
db.session.rollback()
|
||
return {"code": 5001, "message": f"导入失败: {str(e)}", "data": None}
|
||
|
||
def sync_all_users_to_dify(self) -> dict:
|
||
"""同步所有本地用户到 Dify。"""
|
||
try:
|
||
# 获取所有本地用户
|
||
users = db.session.query(WeComUserMapping).filter(
|
||
WeComUserMapping.wecom_userid.like('local_%')
|
||
).all()
|
||
|
||
# 获取已同步的用户
|
||
result = db.session.execute(db.text("SELECT name FROM accounts"))
|
||
existing_users = {row[0] for row in result.fetchall()}
|
||
|
||
synced = 0
|
||
skipped = 0
|
||
errors = 0
|
||
|
||
for user in users:
|
||
if user.username in existing_users:
|
||
skipped += 1
|
||
continue
|
||
|
||
# 生成默认密码
|
||
default_password = f"Baodan@{user.id}2024"
|
||
if self._sync_user_to_dify(user.username, default_password, user.email):
|
||
synced += 1
|
||
else:
|
||
errors += 1
|
||
|
||
return {
|
||
"code": 0,
|
||
"message": f"同步完成: 成功 {synced}, 跳过 {skipped}, 失败 {errors}",
|
||
"data": {"synced": synced, "skipped": skipped, "errors": errors}
|
||
}
|
||
except Exception as e:
|
||
return {"code": 500, "message": f"同步失败: {str(e)}", "data": None}
|
||
|
||
# ---- 角色管理 ----
|
||
|
||
_builtin_roles_initialized = False
|
||
|
||
def _init_builtin_roles(self):
|
||
"""初始化内置角色(如果不存在)。只在首次调用时执行。"""
|
||
if AdminService._builtin_roles_initialized:
|
||
return
|
||
|
||
import json
|
||
from insurance.models.role import Role
|
||
|
||
builtin_roles = [
|
||
{"name": "super_admin", "permissions": ["*"], "description": "超级管理员"},
|
||
{"name": "admin", "permissions": ["*"], "description": "管理员"},
|
||
{"name": "manager", "permissions": ["chat", "proposal", "kb_view", "team_stats"], "description": "销售主管"},
|
||
{"name": "sales", "permissions": ["chat", "proposal"], "description": "销售人员"},
|
||
{"name": "client", "permissions": ["chat"], "description": "客户"},
|
||
]
|
||
|
||
for role_data in builtin_roles:
|
||
existing = db.session.query(Role).filter_by(name=role_data["name"]).first()
|
||
if existing:
|
||
# 同步内置角色的权限(确保代码中的定义与数据库一致)
|
||
new_perms = json.dumps(role_data["permissions"])
|
||
if existing.permissions != new_perms:
|
||
existing.permissions = new_perms
|
||
else:
|
||
role = Role(
|
||
name=role_data["name"],
|
||
description=role_data["description"],
|
||
permissions=json.dumps(role_data["permissions"]),
|
||
builtin=True,
|
||
)
|
||
db.session.add(role)
|
||
try:
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
raise
|
||
AdminService._builtin_roles_initialized = True
|
||
|
||
def list_roles(self) -> dict:
|
||
import json
|
||
from insurance.models.role import Role
|
||
|
||
# 确保内置角色存在(首次调用时初始化)
|
||
self._init_builtin_roles()
|
||
|
||
roles = db.session.query(Role).all()
|
||
return {"code": 0, "data": [r.to_dict() for r in roles]}
|
||
|
||
def create_role(self, data: dict) -> dict:
|
||
import json
|
||
from insurance.models.role import Role
|
||
|
||
# 检查角色名是否已存在
|
||
existing = db.session.query(Role).filter_by(name=data.get("name")).first()
|
||
if existing:
|
||
return {"code": 1006, "message": "角色名已存在", "data": None}
|
||
|
||
role = Role(
|
||
name=data.get("name", ""),
|
||
description=data.get("description", ""),
|
||
permissions=json.dumps(data.get("permissions", [])),
|
||
builtin=False,
|
||
)
|
||
db.session.add(role)
|
||
try:
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
raise
|
||
|
||
from insurance.utils.audit import log_operation
|
||
log_operation("system", "create", "role", str(role.id), {"name": role.name})
|
||
|
||
return {"code": 0, "message": "success", "data": {"id": f"role-{role.id}"}}
|
||
|
||
def update_role(self, role_id: str, data: dict) -> dict:
|
||
import json
|
||
from insurance.models.role import Role
|
||
|
||
role_id_int = int(role_id.replace("role-", ""))
|
||
role = db.session.query(Role).filter_by(id=role_id_int).first()
|
||
if not role:
|
||
return {"code": 1005, "message": "角色不存在", "data": None}
|
||
|
||
if role.builtin and "name" in data:
|
||
return {"code": 1007, "message": "内置角色不可修改名称", "data": None}
|
||
|
||
if "description" in data:
|
||
role.description = data["description"]
|
||
if "permissions" in data:
|
||
role.permissions = json.dumps(data["permissions"])
|
||
|
||
try:
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
raise
|
||
|
||
from insurance.utils.audit import log_operation
|
||
log_operation("system", "update", "role", role_id, {"name": role.name})
|
||
|
||
return {"code": 0, "message": "success", "data": None}
|
||
|
||
def delete_role(self, role_id: str) -> dict:
|
||
from insurance.models.role import Role
|
||
|
||
role_id_int = int(role_id.replace("role-", ""))
|
||
role = db.session.query(Role).filter_by(id=role_id_int).first()
|
||
if not role:
|
||
return {"code": 1005, "message": "角色不存在", "data": None}
|
||
|
||
if role.builtin:
|
||
return {"code": 1008, "message": "内置角色不可删除", "data": None}
|
||
|
||
db.session.delete(role)
|
||
try:
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
raise
|
||
|
||
from insurance.utils.audit import log_operation
|
||
log_operation("system", "delete", "role", role_id, {"name": role.name})
|
||
|
||
return {"code": 0, "message": "success", "data": None}
|
||
|
||
# ---- LLM 配置 ----
|
||
|
||
def list_llm_configs(self) -> dict:
|
||
# 调用 BaoDan API 获取模型列表
|
||
base_url = current_app.config.get("BAODAN_API_URL", "http://localhost:5001")
|
||
try:
|
||
resp = requests.get(f"{base_url}/v1/models", timeout=10)
|
||
return {"code": 0, "data": resp.json()}
|
||
except Exception:
|
||
return {"code": 0, "data": []}
|
||
|
||
def create_llm_config(self, data: dict) -> dict:
|
||
# TODO: 存储到数据库并同步到 BaoDan
|
||
return {"code": 0, "message": "success", "data": {"id": "llm-new"}}
|
||
|
||
def update_llm_config(self, config_id: str, data: dict) -> dict:
|
||
return {"code": 0, "message": "success", "data": None}
|
||
|
||
def ping_llm(self, config_id: str) -> dict:
|
||
# TODO: 调用 BaoDan 模型测试接口
|
||
return {"code": 0, "data": {"reachable": True, "latency_ms": 0, "model": ""}}
|
||
|
||
# ---- Prompt 管理 ----
|
||
|
||
def list_prompts(self) -> dict:
|
||
import json
|
||
from insurance.models.prompt import PromptTemplate
|
||
|
||
prompts = db.session.query(PromptTemplate).all()
|
||
return {"code": 0, "data": [p.to_dict() for p in prompts]}
|
||
|
||
def get_prompt(self, prompt_id: str) -> dict:
|
||
from insurance.models.prompt import PromptTemplate
|
||
|
||
prompt_id_int = int(prompt_id.replace("prompt-", ""))
|
||
prompt = db.session.query(PromptTemplate).filter_by(id=prompt_id_int).first()
|
||
if not prompt:
|
||
return {"code": 1005, "message": "模板不存在", "data": None}
|
||
return {"code": 0, "data": prompt.to_dict()}
|
||
|
||
def save_prompt(self, data: dict) -> dict:
|
||
import json
|
||
from insurance.models.prompt import PromptTemplate, PromptVersion
|
||
|
||
prompt_id = data.get("id")
|
||
|
||
if prompt_id:
|
||
# 更新现有模板
|
||
prompt_id_int = int(prompt_id.replace("prompt-", ""))
|
||
prompt = db.session.query(PromptTemplate).filter_by(id=prompt_id_int).first()
|
||
if not prompt:
|
||
return {"code": 1005, "message": "模板不存在", "data": None}
|
||
|
||
# 记录变更前的值
|
||
old_value = {
|
||
"name": prompt.name,
|
||
"description": prompt.description,
|
||
"content": prompt.content,
|
||
"variables": json.loads(prompt.variables) if prompt.variables else [],
|
||
"category": prompt.category,
|
||
}
|
||
|
||
# 保存旧版本
|
||
latest_version = db.session.query(PromptVersion).filter_by(
|
||
prompt_id=prompt_id_int
|
||
).order_by(PromptVersion.version.desc()).first()
|
||
|
||
new_version = (latest_version.version + 1) if latest_version else 1
|
||
|
||
version_record = PromptVersion(
|
||
prompt_id=prompt_id_int,
|
||
version=new_version,
|
||
content=prompt.content,
|
||
variables=prompt.variables,
|
||
change_note=data.get("change_note", ""),
|
||
)
|
||
db.session.add(version_record)
|
||
|
||
# 更新模板
|
||
prompt.name = data.get("name", prompt.name)
|
||
prompt.description = data.get("description", prompt.description)
|
||
prompt.content = data.get("content", prompt.content)
|
||
prompt.variables = json.dumps(data.get("variables", []))
|
||
prompt.category = data.get("category", prompt.category)
|
||
|
||
try:
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
raise
|
||
|
||
# 记录配置变更日志(含变更前后值)
|
||
new_value = {
|
||
"name": prompt.name,
|
||
"description": prompt.description,
|
||
"content": prompt.content,
|
||
"variables": json.loads(prompt.variables) if prompt.variables else [],
|
||
"category": prompt.category,
|
||
}
|
||
from insurance.utils.audit import log_config_change
|
||
log_config_change("system", "prompt", prompt_id, old_value, new_value)
|
||
|
||
return {"code": 0, "data": {"id": f"prompt-{prompt.id}", "version": new_version}}
|
||
else:
|
||
# 创建新模板
|
||
prompt = PromptTemplate(
|
||
name=data.get("name", ""),
|
||
description=data.get("description", ""),
|
||
content=data.get("content", ""),
|
||
variables=json.dumps(data.get("variables", [])),
|
||
category=data.get("category", "general"),
|
||
)
|
||
db.session.add(prompt)
|
||
db.session.flush()
|
||
|
||
# 创建初始版本
|
||
version_record = PromptVersion(
|
||
prompt_id=prompt.id,
|
||
version=1,
|
||
content=prompt.content,
|
||
variables=prompt.variables,
|
||
change_note="初始版本",
|
||
)
|
||
db.session.add(version_record)
|
||
try:
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
raise
|
||
|
||
from insurance.utils.audit import log_operation
|
||
log_operation("system", "create", "prompt", str(prompt.id), {"name": prompt.name})
|
||
|
||
return {"code": 0, "data": {"id": f"prompt-{prompt.id}", "version": 1}}
|
||
|
||
def delete_prompt(self, prompt_id: str) -> dict:
|
||
from insurance.models.prompt import PromptTemplate, PromptVersion
|
||
|
||
prompt_id_int = int(prompt_id.replace("prompt-", ""))
|
||
prompt = db.session.query(PromptTemplate).filter_by(id=prompt_id_int).first()
|
||
if not prompt:
|
||
return {"code": 1005, "message": "模板不存在", "data": None}
|
||
|
||
# 删除版本历史
|
||
db.session.query(PromptVersion).filter_by(prompt_id=prompt_id_int).delete()
|
||
# 删除模板
|
||
db.session.delete(prompt)
|
||
db.session.commit()
|
||
|
||
from insurance.utils.audit import log_operation
|
||
log_operation("system", "delete", "prompt", prompt_id, {"name": prompt.name})
|
||
|
||
return {"code": 0, "message": "success", "data": None}
|
||
|
||
def get_prompt_versions(self, prompt_id: str) -> dict:
|
||
from insurance.models.prompt import PromptVersion
|
||
|
||
prompt_id_int = int(prompt_id.replace("prompt-", ""))
|
||
versions = db.session.query(PromptVersion).filter_by(
|
||
prompt_id=prompt_id_int
|
||
).order_by(PromptVersion.version.desc()).all()
|
||
|
||
return {"code": 0, "data": [v.to_dict() for v in versions]}
|
||
|
||
def rollback_prompt(self, prompt_id: str, version: int) -> dict:
|
||
import json
|
||
from insurance.models.prompt import PromptTemplate, PromptVersion
|
||
|
||
prompt_id_int = int(prompt_id.replace("prompt-", ""))
|
||
prompt = db.session.query(PromptTemplate).filter_by(id=prompt_id_int).first()
|
||
if not prompt:
|
||
return {"code": 1005, "message": "模板不存在", "data": None}
|
||
|
||
version_record = db.session.query(PromptVersion).filter_by(
|
||
prompt_id=prompt_id_int, version=version
|
||
).first()
|
||
if not version_record:
|
||
return {"code": 1005, "message": "版本不存在", "data": None}
|
||
|
||
# 保存当前版本为新版本
|
||
latest_version = db.session.query(PromptVersion).filter_by(
|
||
prompt_id=prompt_id_int
|
||
).order_by(PromptVersion.version.desc()).first()
|
||
|
||
new_version = (latest_version.version + 1) if latest_version else 1
|
||
|
||
rollback_record = PromptVersion(
|
||
prompt_id=prompt_id_int,
|
||
version=new_version,
|
||
content=prompt.content,
|
||
variables=prompt.variables,
|
||
change_note=f"回滚到版本 {version}",
|
||
)
|
||
db.session.add(rollback_record)
|
||
|
||
# 回滚内容
|
||
prompt.content = version_record.content
|
||
prompt.variables = version_record.variables
|
||
db.session.commit()
|
||
|
||
from insurance.utils.audit import log_operation
|
||
log_operation("system", "rollback", "prompt", prompt_id, {"version": version})
|
||
|
||
return {"code": 0, "data": {"id": f"prompt-{prompt.id}", "version": new_version}}
|
||
|
||
def test_prompt(self, data: dict) -> dict:
|
||
api_key = current_app.config.get("BAODAN_CHAT_API_KEY", "")
|
||
base_url = current_app.config.get("BAODAN_API_URL", "http://localhost:5001")
|
||
|
||
# 构建输入
|
||
inputs = data.get("variables", {})
|
||
test_query = data.get("test_query", "")
|
||
|
||
# 如果有模板内容,替换变量
|
||
template_content = data.get("content", "")
|
||
if template_content and inputs:
|
||
for key, value in inputs.items():
|
||
template_content = template_content.replace(f"{{{{{key}}}}}", str(value))
|
||
|
||
resp = requests.post(
|
||
f"{base_url}/v1/chat-messages",
|
||
json={
|
||
"inputs": inputs,
|
||
"query": test_query,
|
||
"response_mode": "blocking",
|
||
"user": "admin",
|
||
},
|
||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||
timeout=30,
|
||
)
|
||
result = resp.json()
|
||
return {
|
||
"code": 0,
|
||
"data": {
|
||
"response": result.get("answer", ""),
|
||
"latency_ms": 0,
|
||
"token_usage": result.get("metadata", {}).get("usage", {}),
|
||
},
|
||
}
|