This commit is contained in:
Your Name
2026-07-29 09:34:02 +08:00
parent 0ff8943ee2
commit f913a57529
54 changed files with 2453 additions and 378 deletions
+203 -19
View File
@@ -19,6 +19,7 @@ import html
import ipaddress
import json
import os
import re
import secrets
import sqlite3
import sys
@@ -41,8 +42,17 @@ 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
CONFIG_KEYS = (
"AI_ENABLED",
@@ -71,15 +81,30 @@ BOOL_KEYS = {
}
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 password_hash(password: str, salt: bytes | None = None) -> tuple[str, str]:
salt = salt or secrets.token_bytes(16)
digest = hashlib.pbkdf2_hmac(
"sha256", password.encode("utf-8"), salt, PBKDF2_ITERATIONS
)
digest = pbkdf2_sha256(password.encode("utf-8"), salt)
return base64.b64encode(salt).decode("ascii"), base64.b64encode(digest).decode("ascii")
@@ -89,9 +114,7 @@ def verify_password(password: str, salt_text: str, digest_text: str) -> bool:
expected = base64.b64decode(digest_text)
except Exception:
return False
actual = hashlib.pbkdf2_hmac(
"sha256", password.encode("utf-8"), salt, PBKDF2_ITERATIONS
)
actual = pbkdf2_sha256(password.encode("utf-8"), salt)
return hmac.compare_digest(actual, expected)
@@ -185,6 +208,15 @@ class Database:
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),
@@ -218,6 +250,20 @@ class Database:
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
@@ -303,6 +349,44 @@ class Database:
db.commit()
return version
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
@@ -717,6 +801,31 @@ class AdminHandler(BaseHTTPRequestHandler):
return False
return is_loopback and hmac.compare_digest(supplied, expected)
def desktop_sync_authorized(self) -> bool:
supplied = self.headers.get("X-Desktop-Sync-Key", "")
return bool(
supplied
and DESKTOP_SYNC_KEY
and hmac.compare_digest(supplied, DESKTOP_SYNC_KEY)
)
def config_payload(self) -> dict[str, Any]:
row = self.db.config()
release = self.db.release()
return {
"version": row["version"],
"updated_at": row["updated_at"],
"updated_by": row["updated_by_name"] or "system",
"config": json.loads(row["config_json"]),
"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"],
},
}
@staticmethod
def csrf_ok(user: sqlite3.Row, form: dict[str, str]) -> bool:
return hmac.compare_digest(str(user["csrf_token"]), str(form.get("csrf", "")))
@@ -743,16 +852,12 @@ class AdminHandler(BaseHTTPRequestHandler):
local_sync = self.local_sync_authorized()
user = None if local_sync else self.require_api_auth()
if local_sync or user:
row = self.db.config()
self.json_response(
HTTPStatus.OK,
{
"version": row["version"],
"updated_at": row["updated_at"],
"updated_by": row["updated_by_name"] or "system",
"config": json.loads(row["config_json"]),
},
)
self.json_response(HTTPStatus.OK, self.config_payload())
elif path == "/api/v1/desktop/config":
if self.desktop_sync_authorized():
self.json_response(HTTPStatus.OK, self.config_payload())
else:
self.json_response(HTTPStatus.UNAUTHORIZED, {"error": "桌面端同步凭证无效"})
else:
self.json_response(HTTPStatus.NOT_FOUND, {"error": "页面不存在"})
@@ -765,6 +870,8 @@ class AdminHandler(BaseHTTPRequestHandler):
self.web_logout()
elif path == "/admin/config":
self.web_save_config()
elif path == "/admin/release":
self.web_save_release()
elif path == "/admin/users/create":
self.web_create_user()
elif path == "/admin/users/update":
@@ -896,6 +1003,28 @@ 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_save_release(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("页面已过期,请刷新后重试")
release = validate_release_form(form)
self.db.save_release(
release["latest_version"],
release["download_url"],
release["release_notes"],
release["force_upgrade"],
user["id"],
self.client_ip,
)
self.redirect(
"/?message="
+ urllib.parse.quote(f"桌面端版本策略已更新为 v{release['latest_version']}")
)
def web_create_user(self) -> None:
auth = self.require_web_auth(roles=("admin",))
if not auth:
@@ -984,15 +1113,17 @@ class AdminHandler(BaseHTTPRequestHandler):
self.html_response(HTTPStatus.OK, page("首次登录 · 配置管理后台", content))
return
row = self.db.config()
release = self.db.release()
config = json.loads(row["config_json"])
config_card = self.config_card(user, csrf, row, config)
release_card = self.release_card(user, csrf, release)
users_card = self.users_card(user, csrf)
audit_card = self.audit_card()
content = f"""
<div class='shell'>{self.sidebar(user, csrf)}<main>
<div class='top'><div><h1>模型配置中心</h1><div class='muted'>保存后,已登录桌面端会在启动或定时同步时自动应用。</div></div>
<div class='version'>CURRENT · v{int(row['version'])}</div></div>
{flash}<div class='grid'><section>{config_card}</section><section>{users_card}{password_card}{audit_card}</section></div>
<div class='version'>应用 v{html.escape(release['latest_version'])} · 配置 v{int(row['version'])}</div></div>
{flash}<div class='grid'><section>{config_card}</section><section>{release_card}{users_card}{password_card}{audit_card}</section></div>
</main></div>"""
self.html_response(HTTPStatus.OK, page("模型配置中心", content))
@@ -1001,7 +1132,7 @@ class AdminHandler(BaseHTTPRequestHandler):
return f"""
<aside><div class='brand'>ZHEN AI ADMIN<small>企微客服统一配置中心</small></div>
<div class='userbox'><b>{html.escape(user['username'])}</b><br><span class='role'>{ROLE_LABELS[user['role']]}</span></div>
<nav><a href='/'>模型配置</a><a href='#users'>用户与角色</a><a href='#security'>账号安全</a></nav>
<nav><a href='/'>模型配置</a><a href='#release'>版本升级</a><a href='#users'>用户与角色</a><a href='#security'>账号安全</a></nav>
<form method='post' action='/logout' style='position:absolute;bottom:25px;left:24px;right:24px'>
<input type='hidden' name='csrf' value='{csrf}'><button class='secondary' style='width:100%'>退出登录</button></form></aside>"""
@@ -1063,6 +1194,32 @@ class AdminHandler(BaseHTTPRequestHandler):
<div style='max-width:240px;margin-top:12px'><label>单次最多工具轮数</label><input type='number' min='1' max='20' name='AI_MCP_MAX_ROUNDS' value='{esc('AI_MCP_MAX_ROUNDS')}'{disabled}></div>
{submit}</form><div class='tiny'>最后更新:{html.escape(row['updated_at'])} · {html.escape(row['updated_by_name'] or 'system')}</div></section>"""
@staticmethod
def release_card(user: sqlite3.Row, csrf: str, release: sqlite3.Row) -> str:
can_edit = user["role"] in ("admin", "operator")
disabled = " disabled" if not can_edit else ""
checked = " checked" if release["force_upgrade"] else ""
latest = html.escape(release["latest_version"], quote=True)
download_url = html.escape(release["download_url"], quote=True)
notes = html.escape(release["release_notes"])
submit = (
"<div class='actions'><button type='submit'>保存版本策略</button></div>"
if can_edit
else "<div class='notice'>当前为只读角色,不能修改版本策略。</div>"
)
return f"""
<section class='card' id='release'><div class='cardhead'><div><h2>桌面端版本升级</h2>
<div class='muted'>软件每次打开都会检查这里的版本策略。</div></div><span class='version'>v{latest}</span></div>
<form method='post' action='/admin/release'><input type='hidden' name='csrf' value='{csrf}'>
<div class='formgrid'>
<div class='full'><label>最新版本号</label><input name='latest_version' value='{latest}' placeholder='例如 1.0.1' required{disabled}></div>
<div class='full'><label>升级下载地址</label><input type='url' name='download_url' value='{download_url}' placeholder='https://...'{disabled}></div>
<div class='full'><label>更新说明</label><textarea name='release_notes' style='min-height:92px'{disabled}>{notes}</textarea></div>
<div class='full'><label class='check'><input type='checkbox' name='force_upgrade' value='1'{checked}{disabled}>强制升级(旧版本只能升级或退出)</label></div>
</div>
<div class='notice'>开启强制升级前,请确认下载地址可正常打开。版本号与桌面端不一致时会立即生效。</div>
{submit}</form><div class='tiny'>最后更新:{html.escape(release['updated_at'])} · {html.escape(release['updated_by_name'] or 'system')}</div></section>"""
def users_card(self, user: sqlite3.Row, csrf: str) -> str:
if user["role"] != "admin":
return f"""<section class='card' id='users'><h2>用户与角色</h2><div class='muted' style='margin-top:8px'>仅管理员可以管理登录账号。</div></section>"""
@@ -1144,6 +1301,33 @@ def validate_config_form(form: dict[str, str], current: dict[str, Any]) -> dict[
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,
}
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description="企微客服助手配置管理后台")
parser.add_argument("--host", default=DEFAULT_HOST, help="监听地址,默认 127.0.0.1")