更新bug

This commit is contained in:
Your Name
2026-07-31 11:48:16 +08:00
parent f913a57529
commit f22cc1a70d
109 changed files with 37586 additions and 927 deletions
+382 -4
View File
@@ -16,6 +16,7 @@ import getpass
import hashlib
import hmac
import html
import http.client
import ipaddress
import json
import os
@@ -56,10 +57,13 @@ APP_VERSION_PATTERN = re.compile(
PBKDF2_ITERATIONS = 310_000
CONFIG_KEYS = (
"AI_ENABLED",
"AI_DEVELOPMENT_MODE",
"AI_PROVIDER_TYPE",
"AI_API_BASE",
"AI_API_KEY",
"AI_MODEL",
"AI_USE_VISION",
"AI_UI_GUARD_ENABLED",
"AI_CONTEXT_ENABLED",
"AI_CONTEXT_MAX_ROUNDS",
"AI_COUNTER_INSULT_ENABLED",
@@ -74,11 +78,18 @@ CONFIG_KEYS = (
)
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(
@@ -122,6 +133,23 @@ 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"
def load_initial_config() -> dict[str, Any]:
path = SCRIPT_DIR / "ai_settings.json"
try:
@@ -130,10 +158,13 @@ def load_initial_config() -> dict[str, Any]:
saved = {}
defaults: dict[str, Any] = {
"AI_ENABLED": True,
"AI_DEVELOPMENT_MODE": False,
"AI_PROVIDER_TYPE": "openai",
"AI_API_BASE": "",
"AI_API_KEY": "",
"AI_MODEL": "",
"AI_USE_VISION": False,
"AI_UI_GUARD_ENABLED": True,
"AI_CONTEXT_ENABLED": True,
"AI_CONTEXT_MAX_ROUNDS": 5,
"AI_COUNTER_INSULT_ENABLED": False,
@@ -148,6 +179,7 @@ def load_initial_config() -> dict[str, Any]:
}
if isinstance(saved, dict):
defaults.update({key: saved[key] for key in CONFIG_KEYS if key in saved})
defaults["AI_PROVIDER_TYPE"] = provider_type(defaults)
return defaults
@@ -534,6 +566,281 @@ class LoginLimiter:
LOGIN_LIMITER = LoginLimiter()
def _model_endpoint(api_base: str, provider: str) -> str:
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
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 = _model_endpoint(api_base, provider)
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,
}
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]}",
}
def model_test_result_page(result: dict[str, Any]) -> str:
ok = bool(result.get("ok"))
title = "模型连接成功" if ok else "模型连接失败"
tone = "flash" if ok else "flash error"
http_status = result.get("http_status")
status_text = str(http_status) if http_status is not None else "未建立 HTTP 响应"
body = f"""
<div class='loginwrap'><section class='login'>
<div class='brand' style='color:var(--accent)'>ZHEN AI ADMIN</div>
<h1>{title}</h1>
<div class='{tone}'>{html.escape(str(result.get('message') or ''))}</div>
<div class='formgrid'>
<div><label>协议</label><div>{html.escape(str(result.get('provider') or ''))}</div></div>
<div><label>耗时</label><div>{int(result.get('latency_ms') or 0)} ms</div></div>
<div class='full'><label>请求地址</label><div>{html.escape(str(result.get('endpoint') or ''))}</div></div>
<div><label>模型</label><div>{html.escape(str(result.get('model') or ''))}</div></div>
<div><label>HTTP 状态</label><div>{html.escape(status_text)}</div></div>
</div>
<div class='actions'><a class='button' href='/'>返回配置后台</a></div>
<div class='tiny'>测试不会保存页面配置,也不会显示 API Key。</div>
</section></div>"""
return page(title, body)
BASE_CSS = """
:root{--bg:#f3f7f5;--surface:#fff;--ink:#17251f;--muted:#67786f;--line:#dce7e1;
--accent:#0d9871;--deep:#102b21;--danger:#c94352;--soft:#e1f4ec}
@@ -777,7 +1084,9 @@ class AdminHandler(BaseHTTPRequestHandler):
return None
return user, token
def require_api_auth(self) -> sqlite3.Row | None:
def require_api_auth(
self, *, roles: tuple[str, ...] | None = None
) -> sqlite3.Row | None:
user, _ = self.auth(True)
if not user:
self.json_response(HTTPStatus.UNAUTHORIZED, {"error": "登录已失效,请重新登录"})
@@ -787,6 +1096,9 @@ class AdminHandler(BaseHTTPRequestHandler):
HTTPStatus.FORBIDDEN, {"error": "请先在后台网页修改初始密码"}
)
return None
if roles and user["role"] not in roles:
self.json_response(HTTPStatus.FORBIDDEN, {"error": "当前角色没有执行此操作的权限"})
return None
return user
def local_sync_authorized(self) -> bool:
@@ -870,6 +1182,8 @@ class AdminHandler(BaseHTTPRequestHandler):
self.web_logout()
elif path == "/admin/config":
self.web_save_config()
elif path == "/admin/model/test":
self.web_test_model()
elif path == "/admin/release":
self.web_save_release()
elif path == "/admin/users/create":
@@ -882,6 +1196,8 @@ class AdminHandler(BaseHTTPRequestHandler):
self.api_login()
elif path == "/api/v1/auth/logout":
self.api_logout()
elif path == "/api/v1/model/test":
self.api_test_model()
else:
self.json_response(HTTPStatus.NOT_FOUND, {"error": "接口不存在"})
except ValueError as exc:
@@ -1003,6 +1319,52 @@ class AdminHandler(BaseHTTPRequestHandler):
version = self.db.save_config(config, user["id"], self.client_ip)
self.redirect("/?message=" + urllib.parse.quote(f"配置已发布为 v{version}"))
def web_test_model(self) -> None:
auth = self.require_web_auth(roles=("admin", "operator"))
if not auth:
return
user, _ = auth
form = self.form_body()
if not self.csrf_ok(user, form):
raise ValueError("页面已过期,请刷新后重试")
current = json.loads(self.db.config()["config_json"])
try:
result = test_model_connection(model_test_config(form, current))
except ValueError as exc:
result = {
"ok": False,
"provider": "-",
"model": form.get("AI_MODEL", ""),
"endpoint": "",
"http_status": None,
"latency_ms": 0,
"message": str(exc),
}
self.db.audit(
user["id"],
"model.test",
f"ok={int(bool(result['ok']))}, model={str(result.get('model') or '')[:80]}, "
f"endpoint={str(result.get('endpoint') or '')[:200]}, http={result.get('http_status')}",
self.client_ip,
)
self.html_response(HTTPStatus.OK, model_test_result_page(result))
def api_test_model(self) -> None:
user = self.require_api_auth(roles=("admin", "operator"))
if not user:
return
current = json.loads(self.db.config()["config_json"])
result = test_model_connection(model_test_config(self.json_body(), current))
self.db.audit(
user["id"],
"model.test.api",
f"ok={int(bool(result['ok']))}, model={str(result.get('model') or '')[:80]}, "
f"endpoint={str(result.get('endpoint') or '')[:200]}, http={result.get('http_status')}",
self.client_ip,
)
status = HTTPStatus.OK if result["ok"] else HTTPStatus.BAD_GATEWAY
self.json_response(status, result)
def web_save_release(self) -> None:
auth = self.require_web_auth(roles=("admin", "operator"))
if not auth:
@@ -1159,11 +1521,20 @@ class AdminHandler(BaseHTTPRequestHandler):
esc = lambda key: html.escape(str(config.get(key, "")), quote=True)
checked = lambda key: " checked" if config.get(key) else ""
disabled = " disabled" if not can_edit else ""
selected_provider = provider_type(config)
provider_options = "".join(
f"<option value='{key}'{' selected' if key == selected_provider else ''}>"
f"{html.escape(label)}</option>"
for key, label in PROVIDER_TYPES.items()
)
mcp = html.escape(
json.dumps(config.get("AI_MCP_SERVERS", []), ensure_ascii=False, indent=2)
)
submit = (
"<div class='actions'><button type='submit'>保存并发布配置</button></div>"
"<div class='actions'>"
"<button class='secondary' type='submit' formaction='/admin/model/test' "
"formtarget='_blank'>测试模型连通性</button>"
"<button type='submit'>保存并发布配置</button></div>"
if can_edit
else "<div class='notice'>当前为只读角色,可查看配置但不能修改。</div>"
)
@@ -1173,12 +1544,15 @@ class AdminHandler(BaseHTTPRequestHandler):
<div class='switches'>
<label class='check'><input type='checkbox' name='AI_ENABLED' value='1'{checked('AI_ENABLED')}{disabled}>启用 AI 回复</label>
<label class='check'><input type='checkbox' name='AI_CONTEXT_ENABLED' value='1'{checked('AI_CONTEXT_ENABLED')}{disabled}>启用会话上下文</label>
<label class='check'><input type='checkbox' name='AI_USE_VISION' value='1'{checked('AI_USE_VISION')}{disabled}>用视觉模式</label>
<label class='check'><input type='checkbox' name='AI_USE_VISION' value='1'{checked('AI_USE_VISION')}{disabled}>始终使用视觉模式(媒体消息自动启用)</label>
<label class='check'><input type='checkbox' name='AI_UI_GUARD_ENABLED' value='1'{checked('AI_UI_GUARD_ENABLED')}{disabled}>启用 AI 页面守护(异常时自动恢复)</label>
<label class='check'><input type='checkbox' name='AI_COUNTER_INSULT_ENABLED' value='1'{checked('AI_COUNTER_INSULT_ENABLED')}{disabled}>启用反辱骂策略</label>
<label class='check'><input type='checkbox' name='AI_MCP_ENABLED' value='1'{checked('AI_MCP_ENABLED')}{disabled}>启用 MCP 工具</label>
<label class='check'><input type='checkbox' name='AI_DEVELOPMENT_MODE' value='1'{checked('AI_DEVELOPMENT_MODE')}{disabled}>开启开发模式(显示脱敏诊断)</label>
</div><div style='height:24px'></div>
<div class='cardhead'><div><h2>模型与身份</h2><div class='muted'>API Key 留空表示保持当前值;页面永不回显密钥。</div></div></div>
<div class='cardhead'><div><h2>模型与身份</h2><div class='muted'>API Key 留空表示使用已保存的值;连通性测试使用页面当前值,但不会保存或回显密钥。</div></div></div>
<div class='formgrid'>
<div class='full'><label>服务类型</label><select name='AI_PROVIDER_TYPE'{disabled}>{provider_options}</select></div>
<div><label>API 地址</label><input name='AI_API_BASE' value='{esc('AI_API_BASE')}' required{disabled}></div>
<div><label>模型名称</label><input name='AI_MODEL' value='{esc('AI_MODEL')}'{disabled}></div>
<div><label>API Key</label><input type='password' name='AI_API_KEY' placeholder='已保存;留空不修改'{disabled}></div>
@@ -1261,6 +1635,10 @@ def validate_config_form(form: dict[str, str], current: dict[str, Any]) -> dict[
config = {key: current.get(key) for key in CONFIG_KEYS}
for key in BOOL_KEYS:
config[key] = form.get(key) == "1"
provider = form.get("AI_PROVIDER_TYPE", "").strip().lower()
if provider not in PROVIDER_TYPES:
raise ValueError("服务类型无效")
config["AI_PROVIDER_TYPE"] = provider
for key in ("AI_API_BASE", "AI_MODEL", "AI_AGENT_NAME", "AI_HOSPITAL_NAME"):
config[key] = form.get(key, "").strip()
api_key = form.get("AI_API_KEY", "").strip()