1961 lines
88 KiB
Python
1961 lines
88 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""配置后台的数据层:库结构、账号与角色、配置校验、模型连通性测试。
|
||
|
||
**这个模块不再自己起 HTTP 服务。** 原来那套服务端渲染的网页后台(8765)已经
|
||
整体退役,界面由 Vue 管理端承担,接口由 `admin_api.py`(8766)提供。
|
||
|
||
留下来的是被三方共用的地基:
|
||
|
||
- `Database` —— 建表与迁移、用户/角色/令牌、配置版本、模型清单、调用留痕
|
||
- 配置层 —— `CONFIG_KEYS` / `CONFIG_DEFAULTS` / `effective_config` /
|
||
`validate_config_form` / `validate_review_rules`
|
||
- `desktop_config_payload` —— 下发给桌面端的那一整份,两处调用方共用一份实现
|
||
- 模型连通性测试、网关地址推算、`write_runtime_info` 本机发现
|
||
|
||
`admin_api.py` 和 `model_gateway.py` 都 import 它,所以文件本身不能删——
|
||
退役的只是那层 HTTP 界面。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import base64
|
||
import hashlib
|
||
import hmac
|
||
import html
|
||
import http.client
|
||
import ipaddress
|
||
import json
|
||
import os
|
||
import re
|
||
import secrets
|
||
import sqlite3
|
||
import sys
|
||
|
||
import model_protocol
|
||
import threading
|
||
import time
|
||
import urllib.parse
|
||
from collections import defaultdict, deque
|
||
from datetime import datetime, timedelta
|
||
from http.cookies import SimpleCookie
|
||
from pathlib import Path
|
||
from typing import Any
|
||
|
||
|
||
SCRIPT_DIR = Path(__file__).resolve().parent
|
||
DEFAULT_DB = SCRIPT_DIR / "backend.db"
|
||
DEFAULT_RUNTIME_FILE = SCRIPT_DIR / "backend_runtime.json"
|
||
DEFAULT_HOST = "127.0.0.1"
|
||
DEFAULT_PORT = 8765
|
||
DEFAULT_PORT_ATTEMPTS = 100
|
||
DEFAULT_ADMIN_PASSWORD = "Admin@123456"
|
||
DEFAULT_APP_VERSION = "1.0.0"
|
||
DEFAULT_DESKTOP_SYNC_KEY = "wcrpa-v1-H3q9mT7xK2pN8cR5vL4sF6dB1yG0uJ"
|
||
DESKTOP_SYNC_KEY = os.environ.get(
|
||
"WECOM_DESKTOP_SYNC_KEY", DEFAULT_DESKTOP_SYNC_KEY
|
||
).strip()
|
||
ROLES = ("admin", "operator", "viewer")
|
||
ROLE_LABELS = {"admin": "管理员", "operator": "配置员", "viewer": "只读用户"}
|
||
APP_VERSION_PATTERN = re.compile(
|
||
r"^v?(0|[1-9]\d*)\.(0|[1-9]\d*)\.(0|[1-9]\d*)"
|
||
r"(?:[-+][0-9A-Za-z.-]+)?$"
|
||
)
|
||
PBKDF2_ITERATIONS = 310_000
|
||
# 下发给桌面端的配置项。
|
||
#
|
||
# 这里**没有**模型连接参数(服务类型 / API 地址 / API Key / 模型名称 / 温度 /
|
||
# max_tokens / 超时)。它们全部搬到了「模型清单 + 角色编排」:
|
||
#
|
||
# * 同一件事有两个地方能配,迟早会不一致,而且出问题时没人说得清以哪边为准。
|
||
# * 密钥留在这份配置里就必须明文下发到每一台客户端。模型清单里是加密存储、
|
||
# 接口只返遮罩值,真正的调用走模型网关——密钥根本不出后端。
|
||
# * 一个客户只能配一个模型。best-of-N、裁判、降级链在这种形状下无从谈起。
|
||
#
|
||
# 留下的都是**客户端行为**,不是模型参数:开不开自动回复、记几轮上下文、
|
||
# 人格叫什么、MCP 工具连哪里。
|
||
CONFIG_KEYS = (
|
||
"AI_ENABLED",
|
||
"AI_DEVELOPMENT_MODE",
|
||
"AI_USE_VISION",
|
||
"AI_UI_GUARD_ENABLED",
|
||
"AI_CONTEXT_ENABLED",
|
||
"AI_CONTEXT_MAX_ROUNDS",
|
||
"AI_COUNTER_INSULT_ENABLED",
|
||
"AI_AGENT_NAME",
|
||
"AI_HOSPITAL_NAME",
|
||
"AI_MCP_ENABLED",
|
||
"AI_MCP_MAX_ROUNDS",
|
||
"AI_MCP_SERVERS",
|
||
"AI_GATEWAY_URL",
|
||
# 选择性审核规则:命中的才进人工审核队列,其余照常自动发送。
|
||
# 见 CONFIG_DEFAULTS 里 AI_REVIEW_RULES 的形状说明。
|
||
"AI_REVIEW_RULES",
|
||
)
|
||
|
||
# 已经被「角色编排」取代的旧字段。库里的老记录还带着它们,读的时候原样保留
|
||
# (不删用户数据),但不再校验、不再下发、界面上也不再出现。
|
||
RETIRED_CONFIG_KEYS = (
|
||
"AI_PROVIDER_TYPE",
|
||
"AI_API_BASE",
|
||
"AI_API_KEY",
|
||
"AI_MODEL",
|
||
"AI_MAX_TOKENS",
|
||
"AI_TEMPERATURE",
|
||
"AI_TIMEOUT",
|
||
)
|
||
BOOL_KEYS = {
|
||
"AI_ENABLED",
|
||
"AI_DEVELOPMENT_MODE",
|
||
"AI_USE_VISION",
|
||
"AI_UI_GUARD_ENABLED",
|
||
"AI_CONTEXT_ENABLED",
|
||
"AI_COUNTER_INSULT_ENABLED",
|
||
"AI_MCP_ENABLED",
|
||
}
|
||
PROVIDER_TYPES = {
|
||
"openai": "OpenAI 兼容(GPT / DeepSeek / vLLM / SGLang 等)",
|
||
"dify": "Dify 应用(chat-messages 接口)",
|
||
"comfyui": "ComfyUI 文生图",
|
||
}
|
||
|
||
|
||
def pbkdf2_sha256(
|
||
password: bytes, salt: bytes, iterations: int = PBKDF2_ITERATIONS
|
||
) -> bytes:
|
||
"""PBKDF2-HMAC-SHA256,兼容未编译 OpenSSL PBKDF2 的 Python。"""
|
||
native_pbkdf2 = getattr(hashlib, "pbkdf2_hmac", None)
|
||
if callable(native_pbkdf2):
|
||
return native_pbkdf2("sha256", password, salt, iterations)
|
||
|
||
# SHA-256 的输出正好是本项目需要的 32 字节,因此只需计算一个块。
|
||
current = hmac.new(password, salt + b"\x00\x00\x00\x01", hashlib.sha256).digest()
|
||
derived = int.from_bytes(current, "big")
|
||
for _ in range(1, iterations):
|
||
current = hmac.new(password, current, hashlib.sha256).digest()
|
||
derived ^= int.from_bytes(current, "big")
|
||
return derived.to_bytes(32, "big")
|
||
|
||
|
||
def now_text() -> str:
|
||
return datetime.now().astimezone().isoformat(timespec="seconds")
|
||
|
||
|
||
def _since_text(days: int) -> str:
|
||
""""最近 N 天"的起点,格式必须和 `now_text()` 写进库里的完全一致。
|
||
|
||
`created_at` 是 TEXT,窗口过滤是字符串比较——比较双方格式不一样的话,日期
|
||
相同的那一天会比错(ISO 的 'T' 大于空格),该留的行被悄悄滤掉。
|
||
"""
|
||
return (
|
||
datetime.now().astimezone() - timedelta(days=days)
|
||
).isoformat(timespec="seconds")
|
||
|
||
|
||
def password_hash(password: str, salt: bytes | None = None) -> tuple[str, str]:
|
||
salt = salt or secrets.token_bytes(16)
|
||
digest = pbkdf2_sha256(password.encode("utf-8"), salt)
|
||
return base64.b64encode(salt).decode("ascii"), base64.b64encode(digest).decode("ascii")
|
||
|
||
|
||
def verify_password(password: str, salt_text: str, digest_text: str) -> bool:
|
||
try:
|
||
salt = base64.b64decode(salt_text)
|
||
expected = base64.b64decode(digest_text)
|
||
except Exception:
|
||
return False
|
||
actual = pbkdf2_sha256(password.encode("utf-8"), salt)
|
||
return hmac.compare_digest(actual, expected)
|
||
|
||
|
||
def token_hash(token: str) -> str:
|
||
return hashlib.sha256(token.encode("utf-8")).hexdigest()
|
||
|
||
|
||
def provider_type(config: dict[str, Any]) -> str:
|
||
value = str(config.get("AI_PROVIDER_TYPE") or "").strip().lower()
|
||
if value in PROVIDER_TYPES:
|
||
return value
|
||
base = str(config.get("AI_API_BASE") or "").lower()
|
||
if "chat-messages" in base or "completion-messages" in base:
|
||
return "dify"
|
||
if "system_stats" in base or "comfyui" in base:
|
||
return "comfyui"
|
||
try:
|
||
if urllib.parse.urlparse(base).port == 8188:
|
||
return "comfyui"
|
||
except ValueError:
|
||
pass
|
||
return "openai"
|
||
|
||
|
||
DEFAULT_GATEWAY_PORT = 8770
|
||
GATEWAY_ANSWER_PATH = "/v1/answer"
|
||
# 反代时网关挂在这个路径前缀下。桌面端只认后台一个地址,网关在哪由后台告诉它。
|
||
GATEWAY_PROXY_PREFIX = "/gateway"
|
||
# 后台自己直连时用的端口。命中这些说明没走反代,网关就在同一台机器的
|
||
# DEFAULT_GATEWAY_PORT 上。
|
||
_DIRECT_BACKEND_PORTS = {"8765", "8766", "8767"}
|
||
|
||
|
||
def derive_gateway_url(configured: str, scheme: str, host: str) -> str:
|
||
"""算出桌面端该把模型请求发到哪儿。
|
||
|
||
这是"只改一个域名"的关键:桌面端配置里只有后台地址,网关地址由后台在同步
|
||
时告诉它。运维改网关位置只需要动后台一处,不用挨个改客户端。
|
||
|
||
三种情况:
|
||
|
||
* 后台配了绝对地址 → 原样使用(网关在别的域名/别的机器上时用这个)
|
||
* 后台配了 `/xxx` → 拼在后台自己的域名后面(同域名不同路径)
|
||
* 留空(默认) → 按后台自己的地址推:
|
||
直连 127.0.0.1:8765 这种 → 同机的 8770 端口
|
||
其它(说明走了反代) → 同域名的 /gateway 前缀
|
||
|
||
留空时的两种推法必须分开:本机开发是三个端口各跑各的,而反代后面只有 80/443
|
||
对外,公网根本连不到 8770。用一套规则套两种部署,总有一边是错的。
|
||
"""
|
||
scheme = (scheme or "http").lower()
|
||
host = str(host or "").strip()
|
||
value = str(configured or "").strip()
|
||
|
||
if value:
|
||
if value.startswith(("http://", "https://")):
|
||
return value.rstrip("/")
|
||
if value.startswith("/"):
|
||
return f"{scheme}://{host}{value.rstrip('/')}" if host else value.rstrip("/")
|
||
# 既不是绝对地址也不是路径,多半是漏了协议头。当成同域名下的路径处理,
|
||
# 比直接拿去发请求(会得到一个看不懂的 URL 解析错误)好。
|
||
return f"{scheme}://{host}/{value.strip('/')}" if host else value
|
||
|
||
if not host:
|
||
return f"http://127.0.0.1:{DEFAULT_GATEWAY_PORT}{GATEWAY_ANSWER_PATH}"
|
||
|
||
hostname, _, port = host.partition(":")
|
||
if port in _DIRECT_BACKEND_PORTS:
|
||
return f"{scheme}://{hostname}:{DEFAULT_GATEWAY_PORT}{GATEWAY_ANSWER_PATH}"
|
||
return f"{scheme}://{host}{GATEWAY_PROXY_PREFIX}{GATEWAY_ANSWER_PATH}"
|
||
|
||
|
||
# CONFIG_KEYS 每一项的出厂默认值。`load_initial_config` 拿它给全新的库打底;
|
||
# `effective_config` 拿它给已存在的库补漏——两处默认值只能有一份,不然新库和
|
||
# 老库升级后看到的默认行为会悄悄不一样。
|
||
CONFIG_DEFAULTS: dict[str, Any] = {
|
||
"AI_ENABLED": True,
|
||
"AI_DEVELOPMENT_MODE": False,
|
||
"AI_USE_VISION": False,
|
||
"AI_UI_GUARD_ENABLED": True,
|
||
"AI_CONTEXT_ENABLED": True,
|
||
"AI_CONTEXT_MAX_ROUNDS": 5,
|
||
"AI_COUNTER_INSULT_ENABLED": False,
|
||
"AI_AGENT_NAME": "客服",
|
||
"AI_HOSPITAL_NAME": "",
|
||
"AI_MCP_ENABLED": False,
|
||
"AI_MCP_MAX_ROUNDS": 5,
|
||
"AI_MCP_SERVERS": [],
|
||
# 留空 = 按后台自己的地址推算,见 derive_gateway_url
|
||
"AI_GATEWAY_URL": "",
|
||
# 选择性审核:命中下面任意一条规则的关键词,这一条回复才会停下来等人工
|
||
# 确认;没命中的照常自动发送。
|
||
#
|
||
# 形状:[{"id": str, "label": str, "keywords": [str, ...], "enabled": bool}, ...]
|
||
# 关键词按子串、不分大小写,同时比对"客户这句话"和"模型准备发的回复"——
|
||
# 客户问的和模型答的任何一边沾上,都算命中,宁可多审一条,不能漏审一条。
|
||
#
|
||
# 这里给的是投用时的起步规则,照着桌面端界面上原来那三个装饰性标签
|
||
# (诊断/用药调整/投诉退款)配的,运营可以随时在「桌面端 → 客户端配置」里
|
||
# 改词、加规则、关掉某一条。
|
||
#
|
||
# 另外有一条不在这份列表里、写死在代码里、关不掉:模型裁判把这条回复判成
|
||
# "高风险"时,不管有没有命中关键词,一样要停下来等人工。裁判只要配了就会
|
||
# 一直在打分(不受 shadow/score_only/arbitrate 影响),这个信号比关键词
|
||
# 匹配更贴近内容本身,没有理由建议关掉它。
|
||
"AI_REVIEW_RULES": [
|
||
{
|
||
"id": "diagnosis",
|
||
"label": "诊断",
|
||
"keywords": ["确诊", "诊断", "是不是得了", "是不是患了", "是什么病"],
|
||
"enabled": True,
|
||
},
|
||
{
|
||
"id": "medication",
|
||
"label": "用药调整",
|
||
"keywords": ["加量", "减量", "停药", "换药", "能不能吃", "剂量"],
|
||
"enabled": True,
|
||
},
|
||
{
|
||
"id": "complaint",
|
||
"label": "投诉退款",
|
||
"keywords": ["投诉", "退款", "退货", "举报", "曝光", "消协"],
|
||
"enabled": True,
|
||
},
|
||
],
|
||
}
|
||
|
||
|
||
def effective_config(stored: dict[str, Any]) -> dict[str, Any]:
|
||
"""把库里存的一份 config_json 投影成 CONFIG_KEYS,缺的/空的都补上默认值。
|
||
|
||
这里补的是两种情况,缺一不可:
|
||
|
||
* 键根本不存在——`AI_GATEWAY_URL` 这种是后来才加进 CONFIG_KEYS 的,
|
||
建库更早的那些库的 config_json 里压根没有这一项。
|
||
* 键存在但值是 `null`——`ai_settings.json` 曾经把未设置的开关存成
|
||
JSON null,那份文件当初就是拿去初始化这张表的,null 就这样被
|
||
原样搬进了库里。
|
||
|
||
不补的后果:这两种情况 `stored.get(key)` 都会返回 `None`,直接发给前端。
|
||
浏览器上的表单把它转存成 JS 的 `null`,保存时再发回来,撞上接口的类型
|
||
校验——`Input should be a valid boolean, input: null`。修好这一处,
|
||
发出去的从一开始就是干净的值,根本不会走到那步。
|
||
"""
|
||
return {
|
||
key: stored[key] if stored.get(key) is not None else CONFIG_DEFAULTS[key]
|
||
for key in CONFIG_KEYS
|
||
}
|
||
|
||
|
||
def load_initial_config() -> dict[str, Any]:
|
||
path = SCRIPT_DIR / "ai_settings.json"
|
||
try:
|
||
saved = json.loads(path.read_text(encoding="utf-8"))
|
||
except (OSError, ValueError, TypeError):
|
||
saved = {}
|
||
defaults = dict(CONFIG_DEFAULTS)
|
||
if isinstance(saved, dict):
|
||
# 只取真正有值的。`null` 是"没设",不是一个值——拿它盖掉默认值,得到的
|
||
# 配置连自家接口的类型校验都过不了(布尔字段变成 None → 422)。
|
||
defaults.update({
|
||
key: saved[key]
|
||
for key in CONFIG_KEYS
|
||
if key in saved and saved[key] is not None
|
||
})
|
||
return defaults
|
||
|
||
|
||
def valid_password(value: str) -> bool:
|
||
"""密码强度下限:10 位以上,且同时含字母和数字。
|
||
|
||
原来挂在老网页后台的请求处理类上。老后台退役了,但这条规则本身和界面无关,
|
||
改密、建用户、重置 admin 三条路径都要用同一把尺子——留在模块级,谁都能用。
|
||
"""
|
||
value = str(value or "")
|
||
return (
|
||
len(value) >= 10
|
||
and any(c.isalpha() for c in value)
|
||
and any(c.isdigit() for c in value)
|
||
)
|
||
|
||
|
||
def desktop_config_payload(db: "Database", *, scheme: str = "http", host: str = "") -> dict[str, Any]:
|
||
"""桌面端同步要的那一整份:配置 + 网关地址 + 模型清单 + 版本信息。
|
||
|
||
同样是从老后台搬出来的。搬而不是各写一份,是因为这份结构决定了每一台客户端
|
||
拿到什么——两份实现迟早会漂移,而漂移的表现是"某台客户端少了个字段",
|
||
极难查。
|
||
"""
|
||
row = db.config()
|
||
release = db.release()
|
||
stored = json.loads(row["config_json"])
|
||
return {
|
||
"version": row["version"],
|
||
"updated_at": row["updated_at"],
|
||
"updated_by": row["updated_by_name"] or "system",
|
||
# 只下发 CONFIG_KEYS。已退休的模型参数(尤其是 AI_API_KEY)留在库里是为了
|
||
# 不删用户数据,但绝不能再发出去——密钥不出后端是这次改版的全部意义。
|
||
"config": effective_config(stored),
|
||
# 桌面端只配后台地址一个,网关在哪由这里告诉它。
|
||
"gateway": {
|
||
"enabled": True,
|
||
"url": derive_gateway_url(
|
||
str(stored.get("AI_GATEWAY_URL") or ""), scheme, host
|
||
),
|
||
},
|
||
# 模型清单和角色编排随配置一起下发。密钥永远不出后端——桌面端拿到的是
|
||
# 遮罩值,真正的调用要走模型网关。
|
||
"models": db.model_providers(),
|
||
"roles": db.model_roles(),
|
||
"release": {
|
||
"latest_version": release["latest_version"],
|
||
"download_url": release["download_url"],
|
||
"release_notes": release["release_notes"],
|
||
"force_upgrade": bool(release["force_upgrade"]),
|
||
"updated_at": release["updated_at"],
|
||
},
|
||
}
|
||
|
||
|
||
class ClosingConnection(sqlite3.Connection):
|
||
"""提交/回滚后立即关闭,避免 Windows 上数据库文件长期被占用。"""
|
||
|
||
def __exit__(self, exc_type, exc_value, traceback):
|
||
try:
|
||
return super().__exit__(exc_type, exc_value, traceback)
|
||
finally:
|
||
self.close()
|
||
|
||
|
||
class Database:
|
||
def __init__(self, path: Path):
|
||
self.path = path
|
||
self.path.parent.mkdir(parents=True, exist_ok=True)
|
||
|
||
def connect(self) -> sqlite3.Connection:
|
||
connection = sqlite3.connect(
|
||
self.path, timeout=10, factory=ClosingConnection
|
||
)
|
||
connection.row_factory = sqlite3.Row
|
||
connection.execute("PRAGMA foreign_keys = ON")
|
||
return connection
|
||
|
||
def migrate(self) -> None:
|
||
"""把库的结构升到当前版本。不建管理员,不写任何业务数据。
|
||
|
||
每一个碰这个库的服务启动时都该调它。以前只有网页后台(console)在
|
||
`initialize()` 里顺带做了迁移,于是单独起 JSON API 或模型网关时,库还是老
|
||
结构——第一次查询直接 `no such column`,而部署的人只会看到服务起不来,
|
||
看不出是因为漏了哪一步。
|
||
"""
|
||
self.initialize("", seed_admin=False)
|
||
|
||
def initialize(self, initial_password: str, *, seed_admin: bool = True) -> bool:
|
||
with self.connect() as db:
|
||
db.executescript(
|
||
"""
|
||
PRAGMA journal_mode = WAL;
|
||
CREATE TABLE IF NOT EXISTS users (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
username TEXT NOT NULL UNIQUE COLLATE NOCASE,
|
||
password_salt TEXT NOT NULL,
|
||
password_digest TEXT NOT NULL,
|
||
role TEXT NOT NULL,
|
||
active INTEGER NOT NULL DEFAULT 1,
|
||
must_change_password INTEGER NOT NULL DEFAULT 1,
|
||
created_at TEXT NOT NULL,
|
||
updated_at TEXT NOT NULL
|
||
);
|
||
CREATE TABLE IF NOT EXISTS auth_tokens (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||
token_digest TEXT NOT NULL UNIQUE,
|
||
kind TEXT NOT NULL,
|
||
csrf_token TEXT NOT NULL,
|
||
device_name TEXT NOT NULL DEFAULT '',
|
||
created_at TEXT NOT NULL,
|
||
expires_at INTEGER NOT NULL,
|
||
last_used_at TEXT NOT NULL
|
||
);
|
||
CREATE TABLE IF NOT EXISTS model_config (
|
||
id INTEGER PRIMARY KEY CHECK(id = 1),
|
||
config_json TEXT NOT NULL,
|
||
version INTEGER NOT NULL DEFAULT 1,
|
||
updated_at TEXT NOT NULL,
|
||
updated_by INTEGER REFERENCES users(id)
|
||
);
|
||
CREATE TABLE IF NOT EXISTS app_release (
|
||
id INTEGER PRIMARY KEY CHECK(id = 1),
|
||
latest_version TEXT NOT NULL,
|
||
download_url TEXT NOT NULL DEFAULT '',
|
||
release_notes TEXT NOT NULL DEFAULT '',
|
||
force_upgrade INTEGER NOT NULL DEFAULT 0,
|
||
updated_at TEXT NOT NULL,
|
||
updated_by INTEGER REFERENCES users(id)
|
||
);
|
||
CREATE TABLE IF NOT EXISTS audit_log (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
user_id INTEGER REFERENCES users(id),
|
||
action TEXT NOT NULL,
|
||
detail TEXT NOT NULL DEFAULT '',
|
||
ip_address TEXT NOT NULL DEFAULT '',
|
||
created_at TEXT NOT NULL
|
||
);
|
||
CREATE TABLE IF NOT EXISTS model_providers (
|
||
id TEXT PRIMARY KEY,
|
||
name TEXT NOT NULL,
|
||
kind TEXT NOT NULL
|
||
CHECK(kind IN ('dify','openai','claude','comfyui')),
|
||
base_url TEXT NOT NULL,
|
||
-- auto:按接口类型补全路径(api.openai.com/v1 → .../v1/chat/completions)
|
||
-- exact:地址原样使用,一个字符不加。给路径不按套路的服务用。
|
||
endpoint_mode TEXT NOT NULL DEFAULT 'auto'
|
||
CHECK(endpoint_mode IN ('auto','exact')),
|
||
api_key_enc TEXT NOT NULL DEFAULT '',
|
||
model TEXT NOT NULL DEFAULT '',
|
||
capabilities TEXT NOT NULL DEFAULT 'text',
|
||
max_tokens INTEGER NOT NULL DEFAULT 500,
|
||
temperature REAL NOT NULL DEFAULT 0.35,
|
||
timeout_ms INTEGER NOT NULL DEFAULT 30000,
|
||
max_inflight INTEGER NOT NULL DEFAULT 32,
|
||
rpm_limit INTEGER NOT NULL DEFAULT 0,
|
||
enabled INTEGER NOT NULL DEFAULT 1,
|
||
health TEXT NOT NULL DEFAULT 'unknown',
|
||
health_checked_at TEXT NOT NULL DEFAULT '',
|
||
created_at TEXT NOT NULL,
|
||
updated_at TEXT NOT NULL,
|
||
updated_by INTEGER REFERENCES users(id)
|
||
);
|
||
CREATE TABLE IF NOT EXISTS model_roles (
|
||
version INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
answer_ids TEXT NOT NULL DEFAULT '',
|
||
judge_id TEXT NOT NULL DEFAULT '',
|
||
vision_id TEXT NOT NULL DEFAULT '',
|
||
fallback_ids TEXT NOT NULL DEFAULT '',
|
||
judge_mode TEXT NOT NULL DEFAULT 'shadow'
|
||
CHECK(judge_mode IN ('shadow','score_only','arbitrate')),
|
||
updated_at TEXT NOT NULL,
|
||
updated_by INTEGER REFERENCES users(id)
|
||
);
|
||
CREATE TABLE IF NOT EXISTS model_calls (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
device_id TEXT NOT NULL DEFAULT '',
|
||
task_id TEXT NOT NULL DEFAULT '',
|
||
roles_version INTEGER NOT NULL DEFAULT 0,
|
||
judge_mode TEXT NOT NULL DEFAULT '',
|
||
chosen TEXT NOT NULL DEFAULT '',
|
||
judge_winner TEXT NOT NULL DEFAULT '',
|
||
judge_score REAL NOT NULL DEFAULT 0,
|
||
judge_risk TEXT NOT NULL DEFAULT '',
|
||
candidates_json TEXT NOT NULL DEFAULT '[]',
|
||
total_ms INTEGER NOT NULL DEFAULT 0,
|
||
customer_text TEXT NOT NULL DEFAULT '',
|
||
reply_text TEXT NOT NULL DEFAULT '',
|
||
review_reason TEXT NOT NULL DEFAULT '',
|
||
purpose TEXT NOT NULL DEFAULT 'chat',
|
||
created_at TEXT NOT NULL
|
||
);
|
||
CREATE INDEX IF NOT EXISTS idx_tokens_digest ON auth_tokens(token_digest);
|
||
CREATE INDEX IF NOT EXISTS idx_audit_created ON audit_log(created_at DESC);
|
||
CREATE INDEX IF NOT EXISTS idx_model_calls_created
|
||
ON model_calls(created_at DESC);
|
||
CREATE INDEX IF NOT EXISTS idx_model_calls_device
|
||
ON model_calls(device_id, created_at DESC);
|
||
"""
|
||
)
|
||
# RBAC 必须在 users 迁移之前装好:迁移后的 users.role 是指向 roles
|
||
# 表的外键,角色行不存在的话整张表建不起来。
|
||
import rbac
|
||
|
||
rbac.seed(db, now_text())
|
||
self._migrate_role_check(db)
|
||
self._add_missing_columns(db)
|
||
if not seed_admin:
|
||
db.commit()
|
||
return False
|
||
created = db.execute("SELECT COUNT(*) FROM users").fetchone()[0] == 0
|
||
if created:
|
||
salt, digest = password_hash(initial_password)
|
||
stamp = now_text()
|
||
db.execute(
|
||
"""INSERT INTO users
|
||
(username,password_salt,password_digest,role,active,must_change_password,created_at,updated_at)
|
||
VALUES (?,?,?,?,1,1,?,?)""",
|
||
("admin", salt, digest, "admin", stamp, stamp),
|
||
)
|
||
if db.execute("SELECT COUNT(*) FROM model_config").fetchone()[0] == 0:
|
||
admin = db.execute("SELECT id FROM users ORDER BY id LIMIT 1").fetchone()
|
||
db.execute(
|
||
"""INSERT INTO model_config
|
||
(id,config_json,version,updated_at,updated_by) VALUES (1,?,1,?,?)""",
|
||
(
|
||
json.dumps(load_initial_config(), ensure_ascii=False),
|
||
now_text(),
|
||
admin["id"] if admin else None,
|
||
),
|
||
)
|
||
if db.execute("SELECT COUNT(*) FROM app_release").fetchone()[0] == 0:
|
||
admin = db.execute("SELECT id FROM users ORDER BY id LIMIT 1").fetchone()
|
||
db.execute(
|
||
"""INSERT INTO app_release
|
||
(id,latest_version,download_url,release_notes,force_upgrade,updated_at,updated_by)
|
||
VALUES (1,?,?,?,0,?,?)""",
|
||
(
|
||
DEFAULT_APP_VERSION,
|
||
"",
|
||
"",
|
||
now_text(),
|
||
admin["id"] if admin else None,
|
||
),
|
||
)
|
||
db.commit()
|
||
return created
|
||
|
||
def authenticate(self, username: str, password: str) -> sqlite3.Row | None:
|
||
with self.connect() as db:
|
||
user = db.execute(
|
||
"SELECT * FROM users WHERE username = ? AND active = 1", (username,)
|
||
).fetchone()
|
||
if user and verify_password(password, user["password_salt"], user["password_digest"]):
|
||
return user
|
||
return None
|
||
|
||
def create_token(
|
||
self, user_id: int, kind: str, device_name: str, lifetime: int
|
||
) -> tuple[str, str]:
|
||
token = secrets.token_urlsafe(36)
|
||
csrf = secrets.token_urlsafe(24)
|
||
stamp = now_text()
|
||
with self.connect() as db:
|
||
db.execute(
|
||
"""INSERT INTO auth_tokens
|
||
(user_id,token_digest,kind,csrf_token,device_name,created_at,expires_at,last_used_at)
|
||
VALUES (?,?,?,?,?,?,?,?)""",
|
||
(
|
||
user_id,
|
||
token_hash(token),
|
||
kind,
|
||
csrf,
|
||
str(device_name or "")[:120],
|
||
stamp,
|
||
int(time.time()) + lifetime,
|
||
stamp,
|
||
),
|
||
)
|
||
db.execute("DELETE FROM auth_tokens WHERE expires_at < ?", (int(time.time()),))
|
||
db.commit()
|
||
return token, csrf
|
||
|
||
def session(self, token: str) -> sqlite3.Row | None:
|
||
if not token:
|
||
return None
|
||
with self.connect() as db:
|
||
row = db.execute(
|
||
"""SELECT u.*, t.id AS token_id, t.kind AS token_kind,
|
||
t.csrf_token, t.expires_at
|
||
FROM auth_tokens t JOIN users u ON u.id=t.user_id
|
||
WHERE t.token_digest=? AND t.expires_at>=? AND u.active=1""",
|
||
(token_hash(token), int(time.time())),
|
||
).fetchone()
|
||
if row:
|
||
db.execute(
|
||
"UPDATE auth_tokens SET last_used_at=? WHERE id=?",
|
||
(now_text(), row["token_id"]),
|
||
)
|
||
db.commit()
|
||
return row
|
||
|
||
def revoke(self, token: str) -> None:
|
||
if not token:
|
||
return
|
||
with self.connect() as db:
|
||
db.execute("DELETE FROM auth_tokens WHERE token_digest=?", (token_hash(token),))
|
||
db.commit()
|
||
|
||
def config(self) -> sqlite3.Row:
|
||
with self.connect() as db:
|
||
return db.execute(
|
||
"""SELECT c.*, u.username AS updated_by_name
|
||
FROM model_config c LEFT JOIN users u ON u.id=c.updated_by WHERE c.id=1"""
|
||
).fetchone()
|
||
|
||
def save_config(self, config: dict[str, Any], user_id: int, ip: str) -> int:
|
||
with self.connect() as db:
|
||
db.execute("BEGIN IMMEDIATE")
|
||
current = db.execute("SELECT version FROM model_config WHERE id=1").fetchone()
|
||
version = int(current["version"]) + 1
|
||
db.execute(
|
||
"""UPDATE model_config SET config_json=?,version=?,updated_at=?,updated_by=?
|
||
WHERE id=1""",
|
||
(json.dumps(config, ensure_ascii=False), version, now_text(), user_id),
|
||
)
|
||
self._audit(db, user_id, "config.update", f"version={version}", ip)
|
||
db.commit()
|
||
return version
|
||
|
||
def audit_entries(self, limit: int = 200) -> list[dict[str, Any]]:
|
||
"""审计日志。谁在什么时候改了什么——出事之后唯一能回溯的东西。"""
|
||
limit = max(1, min(int(limit or 200), 1000))
|
||
with self.connect() as db:
|
||
rows = db.execute(
|
||
"""SELECT a.id, a.action, a.detail, a.ip_address, a.created_at,
|
||
u.username
|
||
FROM audit_log a LEFT JOIN users u ON u.id = a.user_id
|
||
ORDER BY a.id DESC LIMIT ?""",
|
||
(limit,),
|
||
).fetchall()
|
||
return [{key: row[key] for key in row.keys()} for row in rows]
|
||
|
||
def model_call_stats(self, days: int = 7) -> dict[str, Any]:
|
||
"""调用统计。回答的是这三个问题:
|
||
|
||
· 各出口被选中的比例 → 第二个模型到底值不值那一倍成本
|
||
· 裁判分数的分布 → 现有回复的真实水平,够好就没必要上双模型
|
||
· 高风险占比 → 该不该把 score_only 打开、审核阈值定在哪
|
||
|
||
这些是决定要不要花第二份钱的唯一依据,拍脑袋定不出来。
|
||
|
||
只统计 `purpose='chat'`:界面守卫的内部识别调用也会经过网关,条数远多于
|
||
真实对话。把它们算进来,裁判分布量的就成了"布局识别答得准不准",而这张
|
||
表要回答的是"发给客户的回复够不够好"——两件事,不能混在一个平均值里。
|
||
"""
|
||
days = max(1, min(int(days or 7), 90))
|
||
since = _since_text(days)
|
||
with self.connect() as db:
|
||
total = db.execute(
|
||
"SELECT COUNT(*) AS n FROM model_calls "
|
||
"WHERE created_at >= ? AND purpose = 'chat'",
|
||
(since,),
|
||
).fetchone()["n"]
|
||
chosen = db.execute(
|
||
"""SELECT chosen, COUNT(*) AS n FROM model_calls
|
||
WHERE created_at >= ? AND purpose = 'chat' AND chosen != ''
|
||
GROUP BY chosen ORDER BY n DESC""",
|
||
(since,),
|
||
).fetchall()
|
||
risk = db.execute(
|
||
"""SELECT judge_risk, COUNT(*) AS n FROM model_calls
|
||
WHERE created_at >= ? AND purpose = 'chat' AND judge_risk != ''
|
||
GROUP BY judge_risk""",
|
||
(since,),
|
||
).fetchall()
|
||
# 分数分桶:0.0~0.2 / 0.2~0.4 / … 直接在 SQL 里算,别把几万行拉回来
|
||
buckets = db.execute(
|
||
"""SELECT CAST(judge_score * 5 AS INTEGER) AS bucket, COUNT(*) AS n
|
||
FROM model_calls
|
||
WHERE created_at >= ? AND purpose = 'chat' AND judge_winner != ''
|
||
GROUP BY bucket ORDER BY bucket""",
|
||
(since,),
|
||
).fetchall()
|
||
judged = db.execute(
|
||
"""SELECT COUNT(*) AS n, AVG(judge_score) AS avg_score
|
||
FROM model_calls
|
||
WHERE created_at >= ? AND purpose = 'chat' AND judge_winner != ''""",
|
||
(since,),
|
||
).fetchone()
|
||
latency = db.execute(
|
||
"""SELECT AVG(total_ms) AS avg_ms, MAX(total_ms) AS max_ms
|
||
FROM model_calls WHERE created_at >= ? AND purpose = 'chat'""",
|
||
(since,),
|
||
).fetchone()
|
||
return {
|
||
"since": since,
|
||
"total": int(total),
|
||
"judged": int(judged["n"] or 0),
|
||
"avg_score": round(float(judged["avg_score"] or 0.0), 3),
|
||
"avg_ms": int(latency["avg_ms"] or 0),
|
||
"max_ms": int(latency["max_ms"] or 0),
|
||
"chosen": [{"provider": row["chosen"], "count": row["n"]} for row in chosen],
|
||
"risk": {row["judge_risk"]: row["n"] for row in risk},
|
||
"score_buckets": [
|
||
{
|
||
"range": f"{min(int(row['bucket']), 4) * 0.2:.1f}~"
|
||
f"{(min(int(row['bucket']), 4) + 1) * 0.2:.1f}",
|
||
"count": row["n"],
|
||
}
|
||
for row in buckets
|
||
],
|
||
}
|
||
|
||
def list_model_calls(
|
||
self,
|
||
days: int = 7,
|
||
limit: int = 50,
|
||
offset: int = 0,
|
||
q: str = "",
|
||
purpose: str = "chat",
|
||
) -> dict[str, Any]:
|
||
"""按时间倒序列出单次调用,用于定位"这句话到底是怎么被回复的"。
|
||
|
||
`model_call_stats` 只回答"最近风险占比多高"这类聚合问题;出了具体问题
|
||
要查某一条——客户说了什么、模型候选都答了什么、选中的是哪个、有没有
|
||
触发审核、为什么——只能一条条翻,这就是这个方法的用途。
|
||
|
||
默认只看 `chat`:界面守卫每轮轮询都要问一次模型,条数是真实对话的几十
|
||
倍,混在一起这张表就没法用了。`purpose=""` 表示全都要(内部调用同样在
|
||
花钱,需要时得能查)。
|
||
"""
|
||
days = max(1, min(int(days or 7), 90))
|
||
limit = max(1, min(int(limit or 50), 200))
|
||
offset = max(0, int(offset or 0))
|
||
since = _since_text(days)
|
||
clauses = ["created_at >= ?"]
|
||
params: list[Any] = [since]
|
||
wanted = str(purpose or "").strip().lower()
|
||
if wanted:
|
||
clauses.append("purpose = ?")
|
||
params.append(wanted)
|
||
needle = str(q or "").strip()
|
||
if needle:
|
||
clauses.append("(customer_text LIKE ? OR reply_text LIKE ?)")
|
||
like = f"%{needle}%"
|
||
params += [like, like]
|
||
where = " AND ".join(clauses)
|
||
with self.connect() as db:
|
||
total = db.execute(
|
||
f"SELECT COUNT(*) AS n FROM model_calls WHERE {where}", params
|
||
).fetchone()["n"]
|
||
rows = db.execute(
|
||
f"""SELECT id, device_id, task_id, roles_version, judge_mode, chosen,
|
||
judge_winner, judge_score, judge_risk, candidates_json,
|
||
total_ms, customer_text, reply_text, review_reason,
|
||
purpose, created_at
|
||
FROM model_calls WHERE {where}
|
||
ORDER BY id DESC LIMIT ? OFFSET ?""",
|
||
(*params, limit, offset),
|
||
).fetchall()
|
||
items = []
|
||
for row in rows:
|
||
item = {key: row[key] for key in row.keys()}
|
||
try:
|
||
item["candidates"] = json.loads(item.pop("candidates_json") or "[]")
|
||
except ValueError:
|
||
item["candidates"] = []
|
||
item.pop("candidates_json", None)
|
||
items.append(item)
|
||
return {"total": int(total), "items": items}
|
||
|
||
# ── 角色与权限 ────────────────────────────────────────────────────────
|
||
@staticmethod
|
||
def _add_missing_columns(db) -> None:
|
||
"""给已经存在的表补新加的列。
|
||
|
||
用 ADD COLUMN 而不是重建表:重建要关外键、搬数据、改名,任何一步出错都
|
||
可能把用户的模型清单弄丢,而这里只是加一个带默认值的可选列。
|
||
"""
|
||
wanted = {
|
||
"model_providers": {
|
||
"endpoint_mode": "TEXT NOT NULL DEFAULT 'auto'",
|
||
},
|
||
"model_calls": {
|
||
"customer_text": "TEXT NOT NULL DEFAULT ''",
|
||
"reply_text": "TEXT NOT NULL DEFAULT ''",
|
||
"review_reason": "TEXT NOT NULL DEFAULT ''",
|
||
"purpose": "TEXT NOT NULL DEFAULT 'chat'",
|
||
},
|
||
}
|
||
for table, columns in wanted.items():
|
||
exists = {
|
||
row[1] for row in db.execute(f"PRAGMA table_info({table})")
|
||
}
|
||
if not exists: # 表还没建出来,SCHEMA 里已经带上了
|
||
continue
|
||
for name, spec in columns.items():
|
||
if name not in exists:
|
||
db.execute(f"ALTER TABLE {table} ADD COLUMN {name} {spec}")
|
||
Database._ensure_model_calls_task_unique(db)
|
||
Database._backfill_guard_purpose(db)
|
||
db.commit()
|
||
|
||
@staticmethod
|
||
def _backfill_guard_purpose(db) -> None:
|
||
"""把历史上那些"界面识别"调用从客户调用日志里摘出去。
|
||
|
||
`purpose` 是后加的列,老记录一律落成默认的 `chat`,于是调用日志里一眼
|
||
望过去全是布局识别的 JSON、客户消息全空——真要查"某句话为什么这么回"
|
||
反而找不到。这些内部调用有非常明确的指纹:候选文本里带着分类器专用的
|
||
JSON 字段名,正常客服回复不可能出现。只改分类标记,一行都不删。
|
||
"""
|
||
exists = {row[1] for row in db.execute("PRAGMA table_info(model_calls)")}
|
||
if "purpose" not in exists or "candidates_json" not in exists:
|
||
return
|
||
db.execute(
|
||
"""UPDATE model_calls SET purpose = 'guard'
|
||
WHERE purpose = 'chat' AND (
|
||
candidates_json LIKE '%navigation_right_ratio%'
|
||
OR candidates_json LIKE '%reply_capable%'
|
||
)"""
|
||
)
|
||
|
||
@staticmethod
|
||
def _ensure_model_calls_task_unique(db) -> None:
|
||
"""`task_id` 唯一索引要求先清掉历史重复值。
|
||
|
||
老版本拿会话指纹当 task_id:同一个会话里每一轮对话都复用同一个值,
|
||
直接建唯一索引会撞上已经存在的重复数据。这里只清空重复项里较旧那些
|
||
的 task_id(只留最新一行),不删除任何一行——观测数据不因为一次迁移
|
||
丢失,只是把从来就不是真正调用标识的值收回成空。
|
||
"""
|
||
exists = {row[1] for row in db.execute("PRAGMA table_info(model_calls)")}
|
||
if "task_id" not in exists:
|
||
return
|
||
try:
|
||
db.execute(
|
||
"CREATE UNIQUE INDEX IF NOT EXISTS idx_model_calls_task "
|
||
"ON model_calls(task_id) WHERE task_id != ''"
|
||
)
|
||
except sqlite3.IntegrityError:
|
||
db.execute(
|
||
"""UPDATE model_calls SET task_id = ''
|
||
WHERE task_id != '' AND id NOT IN (
|
||
SELECT MAX(id) FROM model_calls
|
||
WHERE task_id != '' GROUP BY task_id
|
||
)"""
|
||
)
|
||
db.execute(
|
||
"CREATE UNIQUE INDEX IF NOT EXISTS idx_model_calls_task "
|
||
"ON model_calls(task_id) WHERE task_id != ''"
|
||
)
|
||
|
||
def _migrate_role_check(self, db) -> None:
|
||
"""把 users.role 上写死三个角色的 CHECK 约束换成外键。
|
||
|
||
原来的 `CHECK(role IN ('admin','operator','viewer'))` 让"新增一个角色"变
|
||
成一次代码改动 + 发版。SQLite 不支持 DROP CONSTRAINT,只能重建表。
|
||
|
||
重建必须**先关掉外键约束**——这是 SQLite 官方流程里的第一步,我上一版漏
|
||
了,结果在真实库上炸了:`auth_tokens.user_id` 引用 `users`,开着外键
|
||
`DROP TABLE users` 直接 FOREIGN KEY constraint failed。更糟的是
|
||
executescript 处于自动提交状态,失败前建好的 `users_new` 会留在库里,
|
||
下次启动又撞上同一个坑。所以这里还要先清掉上一次的残留。
|
||
|
||
`PRAGMA foreign_keys` 在事务内是空操作,必须在 BEGIN 之前设置。
|
||
"""
|
||
row = db.execute(
|
||
"SELECT sql FROM sqlite_master WHERE type='table' AND name='users'"
|
||
).fetchone()
|
||
needs_migration = bool(row) and "CHECK(role IN" in (row["sql"] or "")
|
||
|
||
leftover = db.execute(
|
||
"SELECT name FROM sqlite_master WHERE type='table' AND name='users_new'"
|
||
).fetchone()
|
||
if leftover and not needs_migration:
|
||
# 迁移其实已经成功过,users_new 是更早一次失败留下的垃圾
|
||
db.execute("DROP TABLE users_new")
|
||
db.commit()
|
||
return
|
||
if not needs_migration:
|
||
return
|
||
|
||
# 角色表里没有的角色值会让 INSERT 撞外键。真实库里出现过手工改库、
|
||
# 或早期版本遗留的值,这里统一落到 viewer 而不是让整个迁移失败——
|
||
# 权限收紧是安全的方向,直接崩掉服务不是。
|
||
known = {item[0] for item in db.execute("SELECT code FROM roles")}
|
||
orphans = [
|
||
item[0]
|
||
for item in db.execute("SELECT DISTINCT role FROM users")
|
||
if item[0] not in known
|
||
]
|
||
|
||
db.commit() # PRAGMA 必须在事务外
|
||
db.execute("PRAGMA foreign_keys = OFF")
|
||
try:
|
||
db.execute("BEGIN IMMEDIATE")
|
||
if leftover:
|
||
db.execute("DROP TABLE users_new")
|
||
if orphans:
|
||
db.execute(
|
||
"UPDATE users SET role='viewer' WHERE role NOT IN "
|
||
f"({','.join('?' for _ in known)})",
|
||
tuple(sorted(known)),
|
||
)
|
||
db.execute(
|
||
"""CREATE TABLE users_new (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
username TEXT NOT NULL UNIQUE COLLATE NOCASE,
|
||
password_salt TEXT NOT NULL,
|
||
password_digest TEXT NOT NULL,
|
||
role TEXT NOT NULL REFERENCES roles(code),
|
||
active INTEGER NOT NULL DEFAULT 1,
|
||
must_change_password INTEGER NOT NULL DEFAULT 1,
|
||
created_at TEXT NOT NULL,
|
||
updated_at TEXT NOT NULL
|
||
)"""
|
||
)
|
||
db.execute(
|
||
"""INSERT INTO users_new
|
||
SELECT id,username,password_salt,password_digest,role,active,
|
||
must_change_password,created_at,updated_at FROM users"""
|
||
)
|
||
db.execute("DROP TABLE users")
|
||
db.execute("ALTER TABLE users_new RENAME TO users")
|
||
broken = db.execute("PRAGMA foreign_key_check").fetchall()
|
||
if broken:
|
||
raise sqlite3.IntegrityError(
|
||
f"迁移后外键校验未通过:{broken[:3]}"
|
||
)
|
||
db.commit()
|
||
except Exception:
|
||
db.rollback()
|
||
raise
|
||
finally:
|
||
db.execute("PRAGMA foreign_keys = ON")
|
||
if orphans:
|
||
print(f"[RBAC] {len(orphans)} 个未知角色已降级为 viewer:{orphans}")
|
||
print("[RBAC] users.role 已从写死的 CHECK 迁移为角色表外键")
|
||
|
||
def roles(self) -> list[dict[str, Any]]:
|
||
import rbac
|
||
|
||
with self.connect() as db:
|
||
rows = db.execute(
|
||
"""SELECT r.*, (SELECT COUNT(*) FROM users u WHERE u.role=r.code)
|
||
AS user_count
|
||
FROM roles r ORDER BY r.builtin DESC, r.code"""
|
||
).fetchall()
|
||
return [
|
||
{
|
||
**{key: row[key] for key in row.keys()},
|
||
"builtin": bool(row["builtin"]),
|
||
"permissions": sorted(rbac.permissions_of(db, row["code"])),
|
||
}
|
||
for row in rows
|
||
]
|
||
|
||
def permission_catalog(self) -> list[dict[str, str]]:
|
||
with self.connect() as db:
|
||
rows = db.execute(
|
||
"SELECT code,name,group_name FROM permissions ORDER BY group_name,code"
|
||
).fetchall()
|
||
return [{key: row[key] for key in row.keys()} for row in rows]
|
||
|
||
def permissions_for_user(self, user_id: int) -> set[str]:
|
||
"""一个用户实际拥有的权限码。前端的路由守卫和按钮显隐都以它为准。"""
|
||
import rbac
|
||
|
||
with self.connect() as db:
|
||
row = db.execute(
|
||
"SELECT role, active FROM users WHERE id=?", (int(user_id),)
|
||
).fetchone()
|
||
if row is None or not row["active"]:
|
||
return set()
|
||
return rbac.permissions_of(db, row["role"])
|
||
|
||
def save_role(
|
||
self, code: str, name: str, permissions: list[str], user_id: int, ip: str
|
||
) -> dict[str, Any]:
|
||
import rbac
|
||
|
||
code = str(code or "").strip().lower()
|
||
if not code or not code.replace("_", "").replace("-", "").isalnum():
|
||
raise ValueError("角色编码只能用字母、数字、下划线和连字符")
|
||
with self.connect() as db:
|
||
db.execute("BEGIN IMMEDIATE")
|
||
existing = db.execute("SELECT * FROM roles WHERE code=?", (code,)).fetchone()
|
||
now = now_text()
|
||
if existing is None:
|
||
db.execute(
|
||
"""INSERT INTO roles (code,name,builtin,created_at,updated_at)
|
||
VALUES (?,?,0,?,?)""",
|
||
(code, str(name or code), now, now),
|
||
)
|
||
else:
|
||
db.execute(
|
||
"UPDATE roles SET name=?,updated_at=? WHERE code=?",
|
||
(str(name or existing["name"]), now, code),
|
||
)
|
||
granted = rbac.grant(db, code, permissions, now)
|
||
self._audit(
|
||
db, user_id, "role.save", f"{code} → {len(granted)} 项权限", ip
|
||
)
|
||
db.commit()
|
||
return {"code": code, "permissions": sorted(granted)}
|
||
|
||
def delete_role(self, code: str, user_id: int, ip: str) -> None:
|
||
"""删角色。内置角色和仍有人在用的角色都不许删。"""
|
||
import rbac
|
||
|
||
code = str(code or "").strip().lower()
|
||
if code in rbac.BUILTIN_ROLE_CODES:
|
||
raise ValueError("内置角色不可删除")
|
||
with self.connect() as db:
|
||
db.execute("BEGIN IMMEDIATE")
|
||
used = db.execute(
|
||
"SELECT COUNT(*) AS n FROM users WHERE role=?", (code,)
|
||
).fetchone()["n"]
|
||
if used:
|
||
raise ValueError(f"还有 {used} 个用户在用这个角色,请先改他们的角色")
|
||
db.execute("DELETE FROM roles WHERE code=?", (code,))
|
||
self._audit(db, user_id, "role.delete", code, ip)
|
||
db.commit()
|
||
|
||
# ── 模型清单与角色编排 ────────────────────────────────────────────────
|
||
def _secret_key(self) -> bytes:
|
||
import secret_box
|
||
|
||
cached = getattr(self, "_secret_key_cache", None)
|
||
if cached is None:
|
||
cached = secret_box.load_or_create_key(Path(self.path).parent)
|
||
self._secret_key_cache = cached
|
||
return cached
|
||
|
||
def model_providers(self, *, include_secrets: bool = False) -> list[dict[str, Any]]:
|
||
"""模型清单。默认只返回遮罩后的密钥——接口层永不回显明文。"""
|
||
import secret_box
|
||
|
||
with self.connect() as db:
|
||
rows = db.execute(
|
||
"SELECT * FROM model_providers ORDER BY enabled DESC, name"
|
||
).fetchall()
|
||
items = []
|
||
for row in rows:
|
||
item = {key: row[key] for key in row.keys()}
|
||
encrypted = item.pop("api_key_enc", "")
|
||
plain = ""
|
||
if encrypted:
|
||
try:
|
||
plain = secret_box.decrypt(encrypted, self._secret_key())
|
||
except Exception:
|
||
# 主密钥换过或密文损坏。这一路必须显式标成不可用,而不是
|
||
# 拿空密钥去调用,换回一个看不懂的 401。
|
||
item["health"] = "down"
|
||
item["last_error"] = "密钥解密失败,请重新填写"
|
||
item["enabled"] = bool(item.get("enabled", 1))
|
||
item["timeout"] = float(item.get("timeout_ms", 30000)) / 1000.0
|
||
item.setdefault("endpoint_mode", "auto")
|
||
# 算好的实际请求地址一并给出去。不给的话,"填的地址"和"真正会请求
|
||
# 的地址"之间隔着一层拼接规则,配错了只能等 404 才发现。
|
||
item["endpoint"] = model_protocol.endpoint_url(
|
||
str(item.get("kind") or ""),
|
||
str(item.get("base_url") or ""),
|
||
str(item.get("endpoint_mode") or "auto"),
|
||
)
|
||
if include_secrets:
|
||
item["api_key"] = plain
|
||
else:
|
||
item["api_key_masked"] = secret_box.masked(plain)
|
||
items.append(item)
|
||
return items
|
||
|
||
def save_model_provider(self, item: dict[str, Any], user_id: int, ip: str) -> str:
|
||
"""新增或更新一条模型出口。api_key 留空表示「不改动原密钥」。"""
|
||
import secret_box
|
||
|
||
provider_id = str(item.get("id") or "").strip()
|
||
if not provider_id:
|
||
raise ValueError("模型 ID 不能为空")
|
||
kind = str(item.get("kind") or "").strip().lower()
|
||
if kind not in {"dify", "openai", "claude", "comfyui"}:
|
||
raise ValueError(f"不支持的接口类型:{kind}")
|
||
base_url = str(item.get("base_url") or "").strip()
|
||
if not base_url:
|
||
raise ValueError("接口地址不能为空")
|
||
endpoint_mode = str(item.get("endpoint_mode") or "auto").strip().lower()
|
||
if endpoint_mode not in model_protocol.ENDPOINT_MODES:
|
||
raise ValueError(f"不支持的地址模式:{endpoint_mode}")
|
||
|
||
with self.connect() as db:
|
||
db.execute("BEGIN IMMEDIATE")
|
||
existing = db.execute(
|
||
"SELECT api_key_enc FROM model_providers WHERE id=?", (provider_id,)
|
||
).fetchone()
|
||
raw_key = str(item.get("api_key") or "")
|
||
if raw_key:
|
||
key_enc = secret_box.encrypt(raw_key, self._secret_key())
|
||
elif existing is not None:
|
||
key_enc = existing["api_key_enc"]
|
||
else:
|
||
key_enc = ""
|
||
now = now_text()
|
||
db.execute(
|
||
"""INSERT INTO model_providers
|
||
(id,name,kind,base_url,endpoint_mode,api_key_enc,model,capabilities,
|
||
max_tokens,temperature,timeout_ms,max_inflight,rpm_limit,
|
||
enabled,health,health_checked_at,created_at,updated_at,updated_by)
|
||
VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,'',?,?,?)
|
||
ON CONFLICT(id) DO UPDATE SET
|
||
name=excluded.name, kind=excluded.kind,
|
||
base_url=excluded.base_url,
|
||
endpoint_mode=excluded.endpoint_mode,
|
||
api_key_enc=excluded.api_key_enc,
|
||
model=excluded.model, capabilities=excluded.capabilities,
|
||
max_tokens=excluded.max_tokens, temperature=excluded.temperature,
|
||
timeout_ms=excluded.timeout_ms, max_inflight=excluded.max_inflight,
|
||
rpm_limit=excluded.rpm_limit, enabled=excluded.enabled,
|
||
updated_at=excluded.updated_at, updated_by=excluded.updated_by""",
|
||
(
|
||
provider_id,
|
||
str(item.get("name") or provider_id),
|
||
kind,
|
||
base_url,
|
||
endpoint_mode,
|
||
key_enc,
|
||
str(item.get("model") or ""),
|
||
str(item.get("capabilities") or "text"),
|
||
int(item.get("max_tokens") or 500),
|
||
float(item.get("temperature") or 0.35),
|
||
int(item.get("timeout_ms") or 30000),
|
||
int(item.get("max_inflight") or 32),
|
||
int(item.get("rpm_limit") or 0),
|
||
1 if item.get("enabled", True) else 0,
|
||
str(item.get("health") or "unknown"),
|
||
now,
|
||
now,
|
||
user_id,
|
||
),
|
||
)
|
||
self._audit(db, user_id, "model.provider.save", provider_id, ip)
|
||
db.commit()
|
||
return provider_id
|
||
|
||
def delete_model_provider(self, provider_id: str, user_id: int, ip: str) -> None:
|
||
"""删一条出口。仍被角色编排引用的不许删——否则回复链路会直接断。"""
|
||
provider_id = str(provider_id or "").strip()
|
||
roles = self.model_roles()
|
||
referenced = set(
|
||
part.strip()
|
||
for key in ("answer_ids", "judge_id", "vision_id", "fallback_ids")
|
||
for part in str(roles.get(key) or "").split(",")
|
||
if part.strip()
|
||
)
|
||
if provider_id in referenced:
|
||
raise ValueError("该模型仍被角色编排引用,请先在编排里移除再删除")
|
||
with self.connect() as db:
|
||
db.execute("DELETE FROM model_providers WHERE id=?", (provider_id,))
|
||
self._audit(db, user_id, "model.provider.delete", provider_id, ip)
|
||
db.commit()
|
||
|
||
def model_roles(self) -> dict[str, Any]:
|
||
with self.connect() as db:
|
||
row = db.execute(
|
||
"SELECT * FROM model_roles ORDER BY version DESC LIMIT 1"
|
||
).fetchone()
|
||
if row is None:
|
||
return {
|
||
"version": 0,
|
||
"answer_ids": "",
|
||
"judge_id": "",
|
||
"vision_id": "",
|
||
"fallback_ids": "",
|
||
"judge_mode": "shadow",
|
||
}
|
||
return {key: row[key] for key in row.keys()}
|
||
|
||
def save_model_roles(self, roles: dict[str, Any], user_id: int, ip: str) -> int:
|
||
"""写一版新编排。编排是追加版本而不是原地改——回滚时直接指回旧版本。"""
|
||
mode = str(roles.get("judge_mode") or "shadow").lower()
|
||
if mode not in {"shadow", "score_only", "arbitrate"}:
|
||
raise ValueError(f"不支持的裁判模式:{mode}")
|
||
known = {item["id"] for item in self.model_providers()}
|
||
answer_ids = [
|
||
part.strip()
|
||
for part in str(roles.get("answer_ids") or "").split(",")
|
||
if part.strip()
|
||
]
|
||
if not answer_ids:
|
||
raise ValueError("至少要指定一个答题模型")
|
||
for key in answer_ids + [roles.get("judge_id"), roles.get("vision_id")]:
|
||
key = str(key or "").strip()
|
||
if key and key not in known:
|
||
raise ValueError(f"模型清单里没有 {key}")
|
||
with self.connect() as db:
|
||
db.execute("BEGIN IMMEDIATE")
|
||
cursor = db.execute(
|
||
"""INSERT INTO model_roles
|
||
(answer_ids,judge_id,vision_id,fallback_ids,judge_mode,
|
||
updated_at,updated_by)
|
||
VALUES (?,?,?,?,?,?,?)""",
|
||
(
|
||
",".join(answer_ids),
|
||
str(roles.get("judge_id") or "").strip(),
|
||
str(roles.get("vision_id") or "").strip(),
|
||
str(roles.get("fallback_ids") or "").strip(),
|
||
mode,
|
||
now_text(),
|
||
user_id,
|
||
),
|
||
)
|
||
version = int(cursor.lastrowid)
|
||
self._audit(
|
||
db, user_id, "model.roles.save", f"version={version} mode={mode}", ip
|
||
)
|
||
db.commit()
|
||
return version
|
||
|
||
def log_model_call(self, record: dict[str, Any]) -> None:
|
||
"""记一次编排调用。写失败绝不能影响回复——这是观测,不是业务。
|
||
|
||
桌面端和模型网关可能对同一次调用各报一次(网关看得到候选和裁判,桌面端
|
||
才知道这条回复触发没触发审核规则)。两边共用 `task_id`:谁先落盘谁建行,
|
||
后到的那边只把自己独有的 `review_reason` 补上去,不会覆盖对方已经写好
|
||
的候选/裁判数据,最终一次调用只留一行。`task_id` 留空(本地无网关路径
|
||
目前就是这样)时不去重,按老行为直接插入。
|
||
"""
|
||
try:
|
||
judge = record.get("judge") or {}
|
||
with self.connect() as db:
|
||
db.execute(
|
||
"""INSERT INTO model_calls
|
||
(device_id,task_id,roles_version,judge_mode,chosen,
|
||
judge_winner,judge_score,judge_risk,candidates_json,
|
||
total_ms,customer_text,reply_text,review_reason,
|
||
purpose,created_at)
|
||
VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)
|
||
ON CONFLICT(task_id) WHERE task_id != '' DO UPDATE SET
|
||
review_reason=excluded.review_reason""",
|
||
(
|
||
str(record.get("device_id") or ""),
|
||
str(record.get("task_id") or ""),
|
||
int(record.get("roles_version") or 0),
|
||
str(record.get("judge_mode") or ""),
|
||
str(record.get("chosen") or ""),
|
||
str(judge.get("winner") or ""),
|
||
float(judge.get("score") or 0.0),
|
||
str(judge.get("risk") or ""),
|
||
json.dumps(record.get("candidates") or [], ensure_ascii=False),
|
||
int(record.get("total_ms") or 0),
|
||
str(record.get("customer_text") or ""),
|
||
str(record.get("reply_text") or ""),
|
||
str(record.get("review_reason") or ""),
|
||
# 桌面端补报的都是真实对话:内部识别调用不经过这条路。
|
||
str(record.get("purpose") or "chat"),
|
||
now_text(),
|
||
),
|
||
)
|
||
db.commit()
|
||
except Exception as exc:
|
||
print(f"[模型调用日志] 落库失败(不影响回复): {exc}")
|
||
|
||
def release(self) -> sqlite3.Row:
|
||
with self.connect() as db:
|
||
return db.execute(
|
||
"""SELECT r.*, u.username AS updated_by_name
|
||
FROM app_release r LEFT JOIN users u ON u.id=r.updated_by WHERE r.id=1"""
|
||
).fetchone()
|
||
|
||
def save_release(
|
||
self,
|
||
latest_version: str,
|
||
download_url: str,
|
||
release_notes: str,
|
||
force_upgrade: bool,
|
||
user_id: int,
|
||
ip: str,
|
||
) -> None:
|
||
with self.connect() as db:
|
||
db.execute(
|
||
"""UPDATE app_release SET latest_version=?,download_url=?,release_notes=?,
|
||
force_upgrade=?,updated_at=?,updated_by=? WHERE id=1""",
|
||
(
|
||
latest_version,
|
||
download_url,
|
||
release_notes,
|
||
int(force_upgrade),
|
||
now_text(),
|
||
user_id,
|
||
),
|
||
)
|
||
self._audit(
|
||
db,
|
||
user_id,
|
||
"release.update",
|
||
f"version={latest_version}, force={int(force_upgrade)}",
|
||
ip,
|
||
)
|
||
db.commit()
|
||
|
||
@staticmethod
|
||
def _audit(
|
||
db: sqlite3.Connection, user_id: int | None, action: str, detail: str, ip: str
|
||
) -> None:
|
||
db.execute(
|
||
"INSERT INTO audit_log(user_id,action,detail,ip_address,created_at) VALUES(?,?,?,?,?)",
|
||
(user_id, action, detail[:500], ip[:80], now_text()),
|
||
)
|
||
|
||
def audit(self, user_id: int | None, action: str, detail: str, ip: str) -> None:
|
||
with self.connect() as db:
|
||
self._audit(db, user_id, action, detail, ip)
|
||
db.commit()
|
||
|
||
def list_users(self) -> list[sqlite3.Row]:
|
||
with self.connect() as db:
|
||
return list(db.execute("SELECT * FROM users ORDER BY id"))
|
||
|
||
def recent_audit(self, limit: int = 12) -> list[sqlite3.Row]:
|
||
with self.connect() as db:
|
||
return list(
|
||
db.execute(
|
||
"""SELECT a.*, u.username FROM audit_log a
|
||
LEFT JOIN users u ON u.id=a.user_id ORDER BY a.id DESC LIMIT ?""",
|
||
(limit,),
|
||
)
|
||
)
|
||
|
||
def create_user(
|
||
self, username: str, password: str, role: str, actor_id: int, ip: str
|
||
) -> None:
|
||
salt, digest = password_hash(password)
|
||
stamp = now_text()
|
||
with self.connect() as db:
|
||
db.execute(
|
||
"""INSERT INTO users
|
||
(username,password_salt,password_digest,role,active,must_change_password,created_at,updated_at)
|
||
VALUES(?,?,?,?,1,1,?,?)""",
|
||
(username, salt, digest, role, stamp, stamp),
|
||
)
|
||
self._audit(db, actor_id, "user.create", f"username={username}, role={role}", ip)
|
||
db.commit()
|
||
|
||
def update_user(
|
||
self,
|
||
user_id: int,
|
||
role: str,
|
||
active: bool,
|
||
new_password: str,
|
||
actor_id: int,
|
||
ip: str,
|
||
) -> None:
|
||
with self.connect() as db:
|
||
target = db.execute("SELECT * FROM users WHERE id=?", (user_id,)).fetchone()
|
||
if not target:
|
||
raise ValueError("用户不存在")
|
||
if user_id == actor_id and not active:
|
||
raise ValueError("不能停用当前登录账号")
|
||
if target["role"] == "admin" and (role != "admin" or not active):
|
||
count = db.execute(
|
||
"SELECT COUNT(*) FROM users WHERE role='admin' AND active=1"
|
||
).fetchone()[0]
|
||
if count <= 1:
|
||
raise ValueError("系统至少需要保留一个启用的管理员")
|
||
values: list[Any] = [role, int(active), now_text()]
|
||
sql = "UPDATE users SET role=?,active=?,updated_at=?"
|
||
if new_password:
|
||
salt, digest = password_hash(new_password)
|
||
sql += ",password_salt=?,password_digest=?,must_change_password=1"
|
||
values.extend([salt, digest])
|
||
sql += " WHERE id=?"
|
||
values.append(user_id)
|
||
db.execute(sql, values)
|
||
if not active or new_password:
|
||
db.execute("DELETE FROM auth_tokens WHERE user_id=?", (user_id,))
|
||
self._audit(
|
||
db,
|
||
actor_id,
|
||
"user.update",
|
||
f"username={target['username']}, role={role}, active={int(active)}, password_reset={bool(new_password)}",
|
||
ip,
|
||
)
|
||
db.commit()
|
||
|
||
def change_password(self, user_id: int, current: str, new_password: str, ip: str) -> None:
|
||
with self.connect() as db:
|
||
user = db.execute("SELECT * FROM users WHERE id=?", (user_id,)).fetchone()
|
||
if not user or not verify_password(
|
||
current, user["password_salt"], user["password_digest"]
|
||
):
|
||
raise ValueError("当前密码不正确")
|
||
salt, digest = password_hash(new_password)
|
||
db.execute(
|
||
"""UPDATE users SET password_salt=?,password_digest=?,must_change_password=0,
|
||
updated_at=? WHERE id=?""",
|
||
(salt, digest, now_text(), user_id),
|
||
)
|
||
self._audit(db, user_id, "password.change", "", ip)
|
||
db.commit()
|
||
|
||
def reset_admin_password(self, password: str) -> None:
|
||
salt, digest = password_hash(password)
|
||
with self.connect() as db:
|
||
result = db.execute(
|
||
"""UPDATE users SET password_salt=?,password_digest=?,active=1,
|
||
must_change_password=1,updated_at=? WHERE username='admin'""",
|
||
(salt, digest, now_text()),
|
||
)
|
||
if result.rowcount == 0:
|
||
stamp = now_text()
|
||
db.execute(
|
||
"""INSERT INTO users
|
||
(username,password_salt,password_digest,role,active,must_change_password,created_at,updated_at)
|
||
VALUES('admin',?,?,'admin',1,1,?,?)""",
|
||
(salt, digest, stamp, stamp),
|
||
)
|
||
db.execute(
|
||
"DELETE FROM auth_tokens WHERE user_id=(SELECT id FROM users WHERE username='admin')"
|
||
)
|
||
db.commit()
|
||
|
||
|
||
class LoginLimiter:
|
||
def __init__(self):
|
||
self._attempts: dict[str, deque[float]] = defaultdict(deque)
|
||
self._lock = threading.Lock()
|
||
|
||
def allowed(self, key: str) -> bool:
|
||
cutoff = time.time() - 300
|
||
with self._lock:
|
||
items = self._attempts[key]
|
||
while items and items[0] < cutoff:
|
||
items.popleft()
|
||
return len(items) < 8
|
||
|
||
def failure(self, key: str) -> None:
|
||
with self._lock:
|
||
self._attempts[key].append(time.time())
|
||
|
||
def success(self, key: str) -> None:
|
||
with self._lock:
|
||
self._attempts.pop(key, None)
|
||
|
||
|
||
LOGIN_LIMITER = LoginLimiter()
|
||
|
||
|
||
def _model_endpoint(api_base: str, provider: str, mode: str = "auto") -> str:
|
||
"""连通性测试要请求的地址。
|
||
|
||
`mode="exact"` 时地址原样使用——和真正调用模型时走的是同一条规则
|
||
(`model_protocol.endpoint_url`)。两边规则不一致的话,测试会去戳一个和实际
|
||
调用不同的地址:测通了照样用不了,或者反过来,最没用的那种测试。
|
||
"""
|
||
base = str(api_base or "").strip().rstrip("/")
|
||
parsed = urllib.parse.urlparse(base)
|
||
if parsed.scheme not in ("http", "https") or not parsed.netloc:
|
||
raise ValueError("API 地址必须是完整的 http 或 https 地址")
|
||
if parsed.query or parsed.fragment:
|
||
raise ValueError("API 地址不能包含查询参数或片段")
|
||
try:
|
||
parsed.port
|
||
except ValueError as exc:
|
||
raise ValueError("API 地址中的端口无效") from exc
|
||
if str(mode or "auto").strip().lower() == "exact":
|
||
return base
|
||
path = (parsed.path or "").rstrip("/")
|
||
lower_path = path.lower()
|
||
if provider == "dify":
|
||
if lower_path.endswith(("/chat-messages", "/completion-messages")):
|
||
return base
|
||
if lower_path.endswith("/v1"):
|
||
return f"{base}/chat-messages"
|
||
return f"{base}/v1/chat-messages"
|
||
if provider == "comfyui":
|
||
if lower_path.endswith("/system_stats"):
|
||
return base
|
||
return f"{base}/system_stats"
|
||
if lower_path.endswith("/chat/completions"):
|
||
return base
|
||
if re.match(r"^/v1/.+", path):
|
||
return base
|
||
return f"{base}/chat/completions"
|
||
|
||
|
||
def model_test_config(values: dict[str, Any], current: dict[str, Any]) -> dict[str, Any]:
|
||
"""合并页面临时值与已保存密钥,不写入数据库。"""
|
||
api_base_value = (
|
||
values.get("AI_API_BASE") if "AI_API_BASE" in values else current.get("AI_API_BASE")
|
||
)
|
||
model_value = values.get("AI_MODEL") if "AI_MODEL" in values else current.get("AI_MODEL")
|
||
api_base = str(api_base_value or "").strip()
|
||
model = str(model_value or "").strip()
|
||
merged_provider = {
|
||
"AI_PROVIDER_TYPE": values.get(
|
||
"AI_PROVIDER_TYPE", current.get("AI_PROVIDER_TYPE")
|
||
),
|
||
"AI_API_BASE": api_base,
|
||
}
|
||
provider = provider_type(merged_provider)
|
||
if "AI_PROVIDER_TYPE" in values:
|
||
requested_provider = str(values.get("AI_PROVIDER_TYPE") or "").strip().lower()
|
||
if requested_provider not in PROVIDER_TYPES:
|
||
raise ValueError("服务类型无效")
|
||
provider = requested_provider
|
||
supplied_key = str(values.get("AI_API_KEY") or "").strip()
|
||
api_key = supplied_key or str(current.get("AI_API_KEY") or "").strip()
|
||
raw_timeout = values.get("AI_TIMEOUT", current.get("AI_TIMEOUT", 30))
|
||
try:
|
||
timeout = int(raw_timeout)
|
||
except (TypeError, ValueError) as exc:
|
||
raise ValueError("请求超时必须是整数") from exc
|
||
timeout = min(60, max(5, timeout))
|
||
endpoint_mode = str(values.get("AI_ENDPOINT_MODE") or "auto").strip().lower()
|
||
endpoint = _model_endpoint(api_base, provider, endpoint_mode)
|
||
if provider == "dify" and not api_key:
|
||
raise ValueError("Dify 连通性测试需要 API Key")
|
||
if provider == "openai" and not model:
|
||
raise ValueError("OpenAI 兼容接口的模型名称不能为空")
|
||
return {
|
||
"api_base": api_base,
|
||
"api_key": api_key,
|
||
"model": model,
|
||
"timeout": timeout,
|
||
"endpoint": endpoint,
|
||
"provider_type": provider,
|
||
"endpoint_mode": endpoint_mode,
|
||
}
|
||
|
||
|
||
def _safe_endpoint_label(endpoint: str) -> str:
|
||
parsed = urllib.parse.urlparse(endpoint)
|
||
host = parsed.hostname or ""
|
||
if parsed.port:
|
||
host = f"{host}:{parsed.port}"
|
||
return urllib.parse.urlunparse((parsed.scheme, host, parsed.path, "", "", ""))
|
||
|
||
|
||
def _remote_error_detail(raw: bytes, api_key: str) -> str:
|
||
text = raw.decode("utf-8", errors="replace")[:1000].strip()
|
||
try:
|
||
data = json.loads(text)
|
||
error = data.get("error") if isinstance(data, dict) else None
|
||
if isinstance(error, dict):
|
||
text = str(error.get("message") or error.get("code") or text)
|
||
elif isinstance(data, dict):
|
||
text = str(data.get("message") or data.get("detail") or text)
|
||
except (TypeError, ValueError):
|
||
pass
|
||
if api_key:
|
||
text = text.replace(api_key, "[已隐藏]")
|
||
return " ".join(text.split())[:300]
|
||
|
||
|
||
def _model_answer(data: dict[str, Any], provider: str) -> str:
|
||
if provider == "dify":
|
||
answer = data.get("answer")
|
||
if answer is None and isinstance(data.get("data"), dict):
|
||
answer = data["data"].get("answer")
|
||
return str(answer or "").strip()
|
||
if provider == "comfyui":
|
||
return "ComfyUI system_stats 可用" if data else ""
|
||
choices = data.get("choices")
|
||
if not isinstance(choices, list) or not choices:
|
||
return ""
|
||
message = choices[0].get("message") if isinstance(choices[0], dict) else None
|
||
content = message.get("content") if isinstance(message, dict) else ""
|
||
if isinstance(content, list):
|
||
content = " ".join(
|
||
str(item.get("text") or "") for item in content if isinstance(item, dict)
|
||
)
|
||
return str(content or "").strip()
|
||
|
||
|
||
def _perform_http_request(
|
||
endpoint: str,
|
||
*,
|
||
method: str,
|
||
headers: dict[str, str],
|
||
payload: dict[str, Any] | None,
|
||
timeout: int,
|
||
) -> tuple[int, bytes]:
|
||
"""直接使用 http.client,避免部分精简 Python 缺少 urllib HTTPSHandler。"""
|
||
parsed = urllib.parse.urlparse(endpoint)
|
||
host = parsed.hostname
|
||
if not host:
|
||
raise ValueError("API 地址缺少主机名")
|
||
path = urllib.parse.urlunparse(("", "", parsed.path or "/", "", parsed.query, ""))
|
||
if parsed.scheme == "https":
|
||
connection_class = getattr(http.client, "HTTPSConnection", None)
|
||
if connection_class is None:
|
||
raise OSError("当前后端 Python 环境缺少 HTTPS/SSL 支持")
|
||
elif parsed.scheme == "http":
|
||
connection_class = http.client.HTTPConnection
|
||
else:
|
||
raise ValueError(f"不支持的 URL 协议:{parsed.scheme or '空'}")
|
||
connection = connection_class(host, parsed.port, timeout=timeout)
|
||
body = None if payload is None else json.dumps(payload, ensure_ascii=False).encode("utf-8")
|
||
try:
|
||
connection.request(method, path, body=body, headers=headers)
|
||
response = connection.getresponse()
|
||
raw = response.read(1_000_001)
|
||
return int(response.status), raw
|
||
finally:
|
||
connection.close()
|
||
|
||
|
||
def test_model_connection(config: dict[str, Any]) -> dict[str, Any]:
|
||
"""向模型发出一个最小请求,返回不含密钥的诊断结果。"""
|
||
endpoint = str(config["endpoint"])
|
||
api_key = str(config.get("api_key") or "")
|
||
provider_type_value = str(config.get("provider_type") or "openai")
|
||
provider = PROVIDER_TYPES.get(provider_type_value, provider_type_value)
|
||
if provider_type_value == "dify":
|
||
payload = {
|
||
"inputs": {},
|
||
"query": "连通性测试:请只回复 OK。",
|
||
"response_mode": "blocking",
|
||
"user": "zhen-ai-backend-test",
|
||
}
|
||
method = "POST"
|
||
elif provider_type_value == "comfyui":
|
||
payload = None
|
||
method = "GET"
|
||
else:
|
||
payload = {
|
||
"model": config["model"],
|
||
"messages": [{"role": "user", "content": "连通性测试:请只回复 OK。"}],
|
||
"stream": False,
|
||
}
|
||
method = "POST"
|
||
headers = {
|
||
"Content-Type": "application/json",
|
||
"Accept": "application/json",
|
||
"User-Agent": "ZhenAI-Backend-Connectivity-Test/1.0",
|
||
}
|
||
if api_key:
|
||
headers["Authorization"] = f"Bearer {api_key}"
|
||
started = time.perf_counter()
|
||
try:
|
||
status, raw = _perform_http_request(
|
||
endpoint,
|
||
method=method,
|
||
headers=headers,
|
||
payload=payload,
|
||
timeout=int(config["timeout"]),
|
||
)
|
||
latency_ms = round((time.perf_counter() - started) * 1000)
|
||
if len(raw) > 1_000_000:
|
||
raise ValueError("模型响应过大,已停止读取")
|
||
if not 200 <= status < 300:
|
||
detail = _remote_error_detail(raw, api_key)
|
||
labels = {
|
||
400: "请求参数不被模型服务接受",
|
||
401: "API Key 无效或缺少鉴权",
|
||
403: "当前 API Key 没有访问权限",
|
||
404: "接口地址或模型名称不存在",
|
||
429: "请求受限、余额不足或调用频率过高",
|
||
}
|
||
message = labels.get(status, f"模型服务返回 HTTP {status}")
|
||
if detail:
|
||
message += f":{detail}"
|
||
return {
|
||
"ok": False,
|
||
"provider": provider,
|
||
"model": str(config.get("model") or "-"),
|
||
"endpoint": _safe_endpoint_label(endpoint),
|
||
"http_status": status,
|
||
"latency_ms": latency_ms,
|
||
"message": message,
|
||
}
|
||
try:
|
||
data = json.loads(raw.decode("utf-8"))
|
||
except (UnicodeDecodeError, ValueError) as exc:
|
||
raise ValueError("模型已响应,但返回内容不是有效 JSON") from exc
|
||
if not isinstance(data, dict):
|
||
raise ValueError("模型已响应,但返回 JSON 不是对象")
|
||
answer = _model_answer(data, provider_type_value)
|
||
if not answer:
|
||
raise ValueError("模型已响应,但未返回可识别的回复内容")
|
||
if api_key:
|
||
answer = answer.replace(api_key, "[已隐藏]")
|
||
return {
|
||
"ok": True,
|
||
"provider": provider,
|
||
"model": str(config.get("model") or "-"),
|
||
"endpoint": _safe_endpoint_label(endpoint),
|
||
"http_status": status,
|
||
"latency_ms": latency_ms,
|
||
"message": f"连接成功,模型回复:{answer[:120]}",
|
||
}
|
||
except (TimeoutError, OSError, http.client.HTTPException, ValueError) as exc:
|
||
latency_ms = round((time.perf_counter() - started) * 1000)
|
||
reason = getattr(exc, "reason", exc)
|
||
message = "连接超时" if isinstance(reason, TimeoutError) else str(reason)
|
||
if api_key:
|
||
message = message.replace(api_key, "[已隐藏]")
|
||
return {
|
||
"ok": False,
|
||
"provider": provider,
|
||
"model": str(config.get("model") or "-"),
|
||
"endpoint": _safe_endpoint_label(endpoint),
|
||
"http_status": None,
|
||
"latency_ms": latency_ms,
|
||
"message": f"连接失败:{' '.join(message.split())[:300]}",
|
||
}
|
||
|
||
|
||
BASE_CSS = """
|
||
:root{--bg:#f3f7f5;--surface:#fff;--ink:#17251f;--muted:#67786f;--line:#dce7e1;
|
||
--accent:#0d9871;--deep:#102b21;--danger:#c94352;--soft:#e1f4ec}
|
||
*{box-sizing:border-box}body{margin:0;background:var(--bg);color:var(--ink);font-family:"Microsoft YaHei UI","Segoe UI",sans-serif}
|
||
a{color:var(--accent);text-decoration:none}.shell{min-height:100vh;display:grid;grid-template-columns:250px 1fr}
|
||
aside{background:var(--deep);color:white;padding:30px 24px;position:sticky;top:0;height:100vh}.brand{font-size:22px;font-weight:800;letter-spacing:.04em}
|
||
.brand small{display:block;margin-top:8px;color:#9fc0b3;font-size:12px;font-weight:500}.userbox{margin-top:44px;padding:16px;background:#17392c;border:1px solid #2b5143;border-radius:14px}
|
||
.role{display:inline-block;margin-top:8px;padding:4px 9px;background:#275643;border-radius:999px;color:#cce9de;font-size:12px}
|
||
nav{margin-top:30px;display:grid;gap:8px}nav a{color:#c6d8d0;padding:10px 12px;border-radius:9px}nav a:hover{background:#1c4435;color:white}
|
||
main{padding:34px;max-width:1450px;width:100%;margin:0 auto}.top{display:flex;justify-content:space-between;align-items:end;margin-bottom:24px}.top h1{font-size:30px;margin:0 0 7px}.muted{color:var(--muted);font-size:13px}
|
||
.grid{display:grid;grid-template-columns:1.35fr .9fr;gap:20px;align-items:start}.card{background:var(--surface);border:1px solid var(--line);border-radius:17px;padding:22px;margin-bottom:20px;box-shadow:0 10px 30px rgba(17,51,38,.035)}
|
||
.card h2{font-size:18px;margin:0 0 5px}.cardhead{display:flex;justify-content:space-between;gap:15px;align-items:start;margin-bottom:18px}.version{font:600 12px Consolas,monospace;color:var(--accent);background:var(--soft);padding:6px 9px;border-radius:8px}
|
||
.formgrid{display:grid;grid-template-columns:1fr 1fr;gap:14px}.full{grid-column:1/-1}label{display:block;color:#53675d;font-size:13px;margin-bottom:6px}input,select,textarea{width:100%;border:1px solid #ceddd5;background:#fbfcfb;border-radius:9px;padding:10px 12px;color:var(--ink);font:inherit;outline:none}
|
||
input:focus,select:focus,textarea:focus{border-color:var(--accent);box-shadow:0 0 0 3px #dff3eb}textarea{min-height:150px;font:13px Consolas,monospace;resize:vertical}.switches{display:grid;grid-template-columns:1fr 1fr;gap:10px 18px}.check{display:flex;align-items:center;gap:9px;color:var(--ink)}.check input{width:18px;height:18px;accent-color:var(--accent)}
|
||
.actions{display:flex;justify-content:flex-end;margin-top:18px;gap:10px}button,.button{border:0;border-radius:9px;padding:10px 17px;font:600 14px inherit;cursor:pointer;background:var(--accent);color:white}button:hover{filter:brightness(.93)}button.secondary{background:#edf3f0;color:#28483b}.danger{color:var(--danger)}
|
||
.flash{padding:12px 15px;border-radius:10px;margin-bottom:18px;background:var(--soft);color:#087656;border:1px solid #c9eadc}.flash.error{background:#fbeaec;color:var(--danger);border-color:#f2cbd1}
|
||
.user{display:grid;grid-template-columns:1.1fr .8fr .6fr 1.1fr auto;gap:8px;align-items:end;padding:12px 0;border-top:1px solid #edf2ef}.user:first-of-type{border-top:0}.username{font-weight:700;padding:11px 0}.user input,.user select{padding:8px 9px}.tiny{font-size:11px;color:#829189}
|
||
.audit{display:grid;gap:12px}.event{border-left:3px solid #b8ddce;padding-left:11px}.event b{font-size:13px}.event div{font-size:11px;color:#829189;margin-top:3px}
|
||
.loginwrap{min-height:100vh;display:grid;place-items:center;padding:24px}.login{width:min(430px,100%);background:white;border:1px solid var(--line);border-radius:22px;padding:34px;box-shadow:0 25px 70px rgba(12,48,34,.12)}.login h1{margin:0 0 8px;font-size:27px}.login form{display:grid;gap:15px;margin-top:26px}.login button{margin-top:5px;padding:12px}.notice{background:#fff7e8;border:1px solid #f0d6a5;color:#8a5a0a;padding:13px;border-radius:10px;font-size:13px;margin-bottom:18px}
|
||
@media(max-width:980px){.shell{grid-template-columns:1fr}aside{height:auto;position:static}.grid{grid-template-columns:1fr}.user{grid-template-columns:1fr 1fr}.user .actions{grid-column:1/-1}.formgrid{grid-template-columns:1fr}.full{grid-column:auto}main{padding:20px}}
|
||
"""
|
||
|
||
|
||
def public_server_url(host: str, port: int) -> str:
|
||
display_host = "127.0.0.1" if host in ("0.0.0.0", "::", "") else host
|
||
if ":" in display_host and not display_host.startswith("["):
|
||
display_host = f"[{display_host}]"
|
||
return f"http://{display_host}:{port}"
|
||
|
||
|
||
def write_runtime_info(
|
||
path: Path,
|
||
host: str,
|
||
port: int,
|
||
*,
|
||
local_sync_token: str = "",
|
||
) -> dict[str, Any]:
|
||
"""发布实际监听端口,供同项目桌面端自动发现。"""
|
||
info = {
|
||
"pid": os.getpid(),
|
||
"host": host,
|
||
"port": int(port),
|
||
"server_url": public_server_url(host, port),
|
||
"local_sync_token": local_sync_token,
|
||
"started_at": now_text(),
|
||
}
|
||
path = path.resolve()
|
||
path.parent.mkdir(parents=True, exist_ok=True)
|
||
temporary = path.with_suffix(path.suffix + ".tmp")
|
||
temporary.write_text(json.dumps(info, ensure_ascii=False, indent=2), encoding="utf-8")
|
||
os.replace(temporary, path)
|
||
return info
|
||
|
||
|
||
def clear_runtime_info(path: Path, port: int) -> None:
|
||
"""只清理由当前进程写入的发现文件,避免影响另一个后台实例。"""
|
||
path = path.resolve()
|
||
try:
|
||
info = json.loads(path.read_text(encoding="utf-8"))
|
||
except (OSError, ValueError, TypeError):
|
||
return
|
||
if info.get("pid") == os.getpid() and info.get("port") == int(port):
|
||
try:
|
||
path.unlink()
|
||
except FileNotFoundError:
|
||
pass
|
||
|
||
|
||
_REVIEW_RULE_ID_PATTERN = re.compile(r"[^a-z0-9]+")
|
||
|
||
|
||
def _slugify_rule_id(label: str, taken: set[str]) -> str:
|
||
"""从规则名造一个可以当 id 用的短串,撞了就加个序号。
|
||
|
||
前端的规则编辑器只让人填"名称"和"关键词",不该逼着运营去想一个英文 id。
|
||
这里自动造、自动去重,界面上就当它不存在。
|
||
"""
|
||
base = _REVIEW_RULE_ID_PATTERN.sub("-", label.strip().lower()).strip("-") or "rule"
|
||
candidate = base
|
||
suffix = 2
|
||
while candidate in taken:
|
||
candidate = f"{base}-{suffix}"
|
||
suffix += 1
|
||
return candidate
|
||
|
||
|
||
def validate_review_rules(raw: Any) -> list[dict[str, Any]]:
|
||
"""校验一份「选择性审核规则」清单。
|
||
|
||
形状见 `CONFIG_DEFAULTS["AI_REVIEW_RULES"]` 上面那段注释。这里管的是"存进去
|
||
的必须是能被 `wechat_bot._reply_needs_review` 安全消费的形状"——不是每一种
|
||
错误都值得单独报,但空关键词、非法类型这些会让规则悄悄失效或者让桌面端在
|
||
读配置时炸掉,必须挡在保存这一步。
|
||
"""
|
||
if not isinstance(raw, list):
|
||
raise ValueError("审核规则 JSON 根节点必须是数组")
|
||
if len(raw) > 50:
|
||
raise ValueError("审核规则最多 50 条")
|
||
seen_ids: set[str] = set()
|
||
cleaned: list[dict[str, Any]] = []
|
||
for index, item in enumerate(raw, start=1):
|
||
if not isinstance(item, dict):
|
||
raise ValueError(f"第 {index} 条审核规则格式不对,应该是一个对象")
|
||
label = str(item.get("label") or "").strip()
|
||
if not label:
|
||
raise ValueError(f"第 {index} 条审核规则没有填名称")
|
||
if len(label) > 40:
|
||
raise ValueError(f"第 {index} 条审核规则名称过长(上限 40 字)")
|
||
keywords_raw = item.get("keywords")
|
||
if not isinstance(keywords_raw, list):
|
||
raise ValueError(f"规则「{label}」的关键词必须是数组")
|
||
keywords = [str(word).strip() for word in keywords_raw if str(word).strip()]
|
||
if not keywords:
|
||
raise ValueError(f"规则「{label}」至少要有一个关键词,否则永远不会命中")
|
||
if len(keywords) > 30:
|
||
raise ValueError(f"规则「{label}」的关键词最多 30 个")
|
||
for word in keywords:
|
||
if len(word) > 40:
|
||
raise ValueError(f"规则「{label}」里有关键词超过 40 字,像是填错了")
|
||
rule_id = str(item.get("id") or "").strip() or _slugify_rule_id(label, seen_ids)
|
||
if rule_id in seen_ids:
|
||
rule_id = _slugify_rule_id(label, seen_ids)
|
||
seen_ids.add(rule_id)
|
||
cleaned.append({
|
||
"id": rule_id,
|
||
"label": label,
|
||
"keywords": keywords,
|
||
"enabled": bool(item.get("enabled", True)),
|
||
})
|
||
return cleaned
|
||
|
||
|
||
def validate_config_form(form: dict[str, str], current: dict[str, Any]) -> dict[str, Any]:
|
||
"""校验并合并一份桌面端配置。
|
||
|
||
模型连接参数已经搬去「模型清单 + 角色编排」,这里不再接收也不再校验。库里
|
||
老记录带着的那几个字段原样保留(见 RETIRED_CONFIG_KEYS)——删掉是删用户
|
||
数据,而且万一要回滚就没得可回。
|
||
"""
|
||
config = {key: current.get(key) for key in CONFIG_KEYS}
|
||
for key in RETIRED_CONFIG_KEYS:
|
||
if key in current:
|
||
config[key] = current[key]
|
||
for key in BOOL_KEYS:
|
||
config[key] = form.get(key) == "1"
|
||
for key in ("AI_AGENT_NAME", "AI_HOSPITAL_NAME"):
|
||
config[key] = form.get(key, "").strip()
|
||
if not config["AI_AGENT_NAME"] or not config["AI_HOSPITAL_NAME"]:
|
||
raise ValueError("客服名称和机构名称不能为空")
|
||
limits = {
|
||
"AI_CONTEXT_MAX_ROUNDS": (1, 50),
|
||
"AI_MCP_MAX_ROUNDS": (1, 20),
|
||
}
|
||
for key, (minimum, maximum) in limits.items():
|
||
try:
|
||
value = int(form.get(key, ""))
|
||
except ValueError as exc:
|
||
raise ValueError(f"{key} 必须是整数") from exc
|
||
if not minimum <= value <= maximum:
|
||
raise ValueError(f"{key} 必须在 {minimum}-{maximum} 之间")
|
||
config[key] = value
|
||
try:
|
||
servers = json.loads(form.get("AI_MCP_SERVERS", "[]") or "[]")
|
||
except ValueError as exc:
|
||
raise ValueError("MCP 服务器不是有效 JSON") from exc
|
||
if not isinstance(servers, list):
|
||
raise ValueError("MCP 服务器 JSON 根节点必须是数组")
|
||
config["AI_MCP_SERVERS"] = servers
|
||
|
||
gateway_url = form.get("AI_GATEWAY_URL", "").strip()
|
||
if gateway_url and not gateway_url.startswith(("http://", "https://", "/")):
|
||
# 这里必须拦住。放过去的话会被拼成
|
||
# https://你的域名/gw.example.com/v1/answer 这种没意义的地址,
|
||
# 而失败要等到桌面端下次发模型请求才暴露出来。
|
||
raise ValueError(
|
||
"模型网关地址要么写完整地址(http:// 或 https:// 开头),"
|
||
"要么写以 / 开头的路径;留空表示按后台地址自动推算"
|
||
)
|
||
if len(gateway_url) > 500:
|
||
raise ValueError("模型网关地址过长")
|
||
config["AI_GATEWAY_URL"] = gateway_url
|
||
|
||
try:
|
||
rules_raw = json.loads(form.get("AI_REVIEW_RULES", "[]") or "[]")
|
||
except ValueError as exc:
|
||
raise ValueError("审核规则不是有效 JSON") from exc
|
||
config["AI_REVIEW_RULES"] = validate_review_rules(rules_raw)
|
||
|
||
return config
|
||
|
||
|
||
def validate_release_form(form: dict[str, str]) -> dict[str, Any]:
|
||
version = form.get("latest_version", "").strip()
|
||
if not APP_VERSION_PATTERN.fullmatch(version):
|
||
raise ValueError("版本号格式应为 1.0.0,可选填写 v 前缀")
|
||
if version.lower().startswith("v"):
|
||
version = version[1:]
|
||
download_url = form.get("download_url", "").strip()
|
||
if download_url:
|
||
parsed = urllib.parse.urlparse(download_url)
|
||
if parsed.scheme not in ("http", "https") or not parsed.netloc:
|
||
raise ValueError("升级下载地址必须是完整的 http 或 https 地址")
|
||
force_upgrade = form.get("force_upgrade") == "1"
|
||
if force_upgrade and not download_url:
|
||
raise ValueError("开启强制升级前必须填写升级下载地址")
|
||
release_notes = form.get("release_notes", "").strip()
|
||
if len(download_url) > 1000:
|
||
raise ValueError("升级下载地址过长")
|
||
if len(release_notes) > 4000:
|
||
raise ValueError("更新说明不能超过 4000 字")
|
||
return {
|
||
"latest_version": version,
|
||
"download_url": download_url,
|
||
"release_notes": release_notes,
|
||
"force_upgrade": force_upgrade,
|
||
}
|
||
|
||
|