1225 lines
55 KiB
Python
1225 lines
55 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""管理 API 与 RBAC。
|
||
|
||
重点不是"接口能返回数据",而是**权限真的拦得住**。前端隐藏按钮只是体验,
|
||
后端不拦就等于没有权限系统——任何人拿 curl 都能绕过前端。
|
||
|
||
所以每一条写接口都有一个"没权限的角色去调"的用例。
|
||
"""
|
||
|
||
import json
|
||
import shutil
|
||
import sqlite3
|
||
import tempfile
|
||
from pathlib import Path
|
||
from unittest import TestCase, main, mock
|
||
|
||
from fastapi.testclient import TestClient
|
||
|
||
import admin_api
|
||
from test_support import local_state_redirect
|
||
import admin_backend as ab
|
||
|
||
STRONG = "Zhen@2026Ops"
|
||
|
||
|
||
class _Base(TestCase):
|
||
def setUp(self):
|
||
self.root = Path(tempfile.mkdtemp())
|
||
self.db = ab.Database(self.root / "t.db")
|
||
self.db.initialize("Admin@123456")
|
||
uid = self.db.authenticate("admin", "Admin@123456")["id"]
|
||
self.db.change_password(uid, "Admin@123456", STRONG, "1.1.1.1")
|
||
self.client = TestClient(admin_api.create_app(self.root / "t.db"))
|
||
|
||
def login(self, username="admin", password=STRONG):
|
||
r = self.client.post(
|
||
"/api/v2/auth/login", json={"username": username, "password": password}
|
||
)
|
||
self.assertEqual(r.status_code, 200, r.text)
|
||
return {"Authorization": f"Bearer {r.json()['access_token']}"}
|
||
|
||
def make_user(self, username, role):
|
||
admin = self.login()
|
||
r = self.client.post(
|
||
"/api/v2/users",
|
||
headers=admin,
|
||
json={"username": username, "password": STRONG, "role": role},
|
||
)
|
||
self.assertEqual(r.status_code, 200, r.text)
|
||
# 新建的用户首登要改密,这里直接把标记清掉,测的是权限不是改密流程
|
||
row = self.db.authenticate(username, STRONG)
|
||
self.db.change_password(row["id"], STRONG, STRONG + "x", "1.1.1.1")
|
||
return self.login(username, STRONG + "x")
|
||
|
||
|
||
class AuthTest(_Base):
|
||
def test_login_returns_the_permission_list(self):
|
||
r = self.client.post(
|
||
"/api/v2/auth/login", json={"username": "admin", "password": STRONG}
|
||
)
|
||
payload = r.json()["user"]
|
||
self.assertEqual(payload["role"], "admin")
|
||
self.assertIn("model:write", payload["permissions"])
|
||
|
||
def test_a_wrong_password_is_rejected(self):
|
||
r = self.client.post(
|
||
"/api/v2/auth/login", json={"username": "admin", "password": "nope"}
|
||
)
|
||
self.assertEqual(r.status_code, 401)
|
||
|
||
def test_no_token_means_401_not_403(self):
|
||
self.assertEqual(self.client.get("/api/v2/me").status_code, 401)
|
||
|
||
def test_a_garbage_token_is_rejected(self):
|
||
r = self.client.get("/api/v2/me", headers={"Authorization": "Bearer not-a-real-token"})
|
||
self.assertEqual(r.status_code, 401)
|
||
|
||
def test_logout_actually_kills_the_token(self):
|
||
headers = self.login()
|
||
self.client.post("/api/v2/auth/logout", headers=headers)
|
||
self.assertEqual(self.client.get("/api/v2/me", headers=headers).status_code, 401)
|
||
|
||
def test_first_login_can_still_reach_the_password_endpoint(self):
|
||
"""老实现首登一律 403,连改密接口都调不通,只能去网页后台改。"""
|
||
self.db.create_user("newbie", STRONG, "viewer", 1, "1.1.1.1")
|
||
r = self.client.post(
|
||
"/api/v2/auth/login", json={"username": "newbie", "password": STRONG}
|
||
)
|
||
self.assertEqual(r.status_code, 200)
|
||
self.assertTrue(r.json()["user"]["must_change_password"])
|
||
headers = {"Authorization": f"Bearer {r.json()['access_token']}"}
|
||
changed = self.client.post(
|
||
"/api/v2/me/password",
|
||
headers=headers,
|
||
json={"current_password": STRONG, "new_password": "Newbie@2026x"},
|
||
)
|
||
self.assertEqual(changed.status_code, 200, changed.text)
|
||
|
||
def test_a_weak_new_password_is_refused(self):
|
||
headers = self.login()
|
||
r = self.client.post(
|
||
"/api/v2/me/password",
|
||
headers=headers,
|
||
json={"current_password": STRONG, "new_password": "123456"},
|
||
)
|
||
self.assertEqual(r.status_code, 400)
|
||
|
||
|
||
class PermissionEnforcementTest(_Base):
|
||
"""前端隐藏按钮只是体验;后端拦不住就等于没有权限系统。"""
|
||
|
||
def test_viewer_can_read_models_but_not_change_them(self):
|
||
headers = self.make_user("v1", "viewer")
|
||
self.assertEqual(self.client.get("/api/v2/models", headers=headers).status_code, 200)
|
||
r = self.client.post(
|
||
"/api/v2/models",
|
||
headers=headers,
|
||
json={"id": "x", "kind": "openai", "base_url": "https://u/v1"},
|
||
)
|
||
self.assertEqual(r.status_code, 403)
|
||
self.assertIn("model:write", r.json()["detail"])
|
||
|
||
def test_operator_can_change_models_but_not_touch_roles(self):
|
||
headers = self.make_user("o1", "operator")
|
||
ok = self.client.post(
|
||
"/api/v2/models",
|
||
headers=headers,
|
||
json={"id": "x", "kind": "openai", "base_url": "https://u/v1", "api_key": "k"},
|
||
)
|
||
self.assertEqual(ok.status_code, 200, ok.text)
|
||
denied = self.client.post(
|
||
"/api/v2/roles", headers=headers, json={"code": "custom", "permissions": []}
|
||
)
|
||
self.assertEqual(denied.status_code, 403)
|
||
|
||
def test_viewer_cannot_manage_users(self):
|
||
headers = self.make_user("v2", "viewer")
|
||
r = self.client.post(
|
||
"/api/v2/users",
|
||
headers=headers,
|
||
json={"username": "who", "password": STRONG, "role": "viewer"},
|
||
)
|
||
self.assertEqual(r.status_code, 403)
|
||
|
||
def test_viewer_cannot_read_the_audit_log(self):
|
||
headers = self.make_user("v3", "viewer")
|
||
self.assertEqual(self.client.get("/api/v2/audit", headers=headers).status_code, 403)
|
||
|
||
def test_the_error_says_which_permission_is_missing(self):
|
||
headers = self.make_user("v4", "viewer")
|
||
r = self.client.delete("/api/v2/models/x", headers=headers)
|
||
self.assertEqual(r.status_code, 403)
|
||
self.assertIn("model:write", r.json()["detail"])
|
||
|
||
def test_revoking_a_permission_takes_effect_immediately(self):
|
||
"""权限存在库里,不在代码里——改完不用重启,也不用重新登录。"""
|
||
headers = self.make_user("o2", "operator")
|
||
admin = self.login()
|
||
before = self.client.post(
|
||
"/api/v2/models",
|
||
headers=headers,
|
||
json={"id": "y", "kind": "openai", "base_url": "https://u/v1", "api_key": "k"},
|
||
)
|
||
self.assertEqual(before.status_code, 200)
|
||
|
||
remaining = sorted(
|
||
set(next(r["permissions"] for r in self.db.roles() if r["code"] == "operator"))
|
||
- {"model:write"}
|
||
)
|
||
self.client.post(
|
||
"/api/v2/roles",
|
||
headers=admin,
|
||
json={"code": "operator", "name": "配置员", "permissions": remaining},
|
||
)
|
||
after = self.client.post(
|
||
"/api/v2/models",
|
||
headers=headers, # 同一个 token,没有重新登录
|
||
json={"id": "z", "kind": "openai", "base_url": "https://u/v1", "api_key": "k"},
|
||
)
|
||
self.assertEqual(after.status_code, 403, "收回权限后旧 token 就该立刻失效")
|
||
|
||
|
||
class RoleManagementTest(_Base):
|
||
def test_a_custom_role_can_be_created_and_used(self):
|
||
headers = self.login()
|
||
r = self.client.post(
|
||
"/api/v2/roles",
|
||
headers=headers,
|
||
json={"code": "reviewer", "name": "审核员",
|
||
"permissions": ["review:read", "review:write", "model:read"]},
|
||
)
|
||
self.assertEqual(r.status_code, 200, r.text)
|
||
codes = {item["code"] for item in r.json()["roles"]}
|
||
self.assertIn("reviewer", codes)
|
||
# 老实现里 users.role 上有写死的 CHECK,自定义角色根本存不进去
|
||
created = self.client.post(
|
||
"/api/v2/users",
|
||
headers=headers,
|
||
json={"username": "shen", "password": STRONG, "role": "reviewer"},
|
||
)
|
||
self.assertEqual(created.status_code, 200, created.text)
|
||
|
||
def test_an_unknown_permission_code_is_refused(self):
|
||
headers = self.login()
|
||
r = self.client.post(
|
||
"/api/v2/roles",
|
||
headers=headers,
|
||
json={"code": "bad", "permissions": ["model:write", "宇宙:毁灭"]},
|
||
)
|
||
self.assertEqual(r.status_code, 400)
|
||
self.assertIn("宇宙:毁灭", r.json()["detail"])
|
||
|
||
def test_the_admin_role_cannot_be_stripped_of_its_powers(self):
|
||
"""把 admin 削权之后就没人能再加回来了,只能改库。"""
|
||
headers = self.login()
|
||
r = self.client.post(
|
||
"/api/v2/roles",
|
||
headers=headers,
|
||
json={"code": "admin", "name": "管理员", "permissions": ["config:read"]},
|
||
)
|
||
self.assertEqual(r.status_code, 400)
|
||
self.assertIn("锁在门外", r.json()["detail"])
|
||
|
||
def test_builtin_roles_cannot_be_deleted(self):
|
||
headers = self.login()
|
||
r = self.client.delete("/api/v2/roles/operator", headers=headers)
|
||
self.assertEqual(r.status_code, 400)
|
||
|
||
def test_a_role_still_in_use_cannot_be_deleted(self):
|
||
headers = self.login()
|
||
self.client.post(
|
||
"/api/v2/roles", headers=headers,
|
||
json={"code": "temp", "permissions": ["model:read"]})
|
||
self.client.post(
|
||
"/api/v2/users", headers=headers,
|
||
json={"username": "t1", "password": STRONG, "role": "temp"})
|
||
r = self.client.delete("/api/v2/roles/temp", headers=headers)
|
||
self.assertEqual(r.status_code, 400)
|
||
self.assertIn("用户", r.json()["detail"])
|
||
|
||
def test_granting_replaces_rather_than_appends(self):
|
||
"""界面是一组勾选框,提交的是最终状态。增量语义会让"取消勾选"静默失效。"""
|
||
headers = self.login()
|
||
self.client.post(
|
||
"/api/v2/roles", headers=headers,
|
||
json={"code": "r1", "permissions": ["model:read", "stats:read"]})
|
||
r = self.client.post(
|
||
"/api/v2/roles", headers=headers,
|
||
json={"code": "r1", "permissions": ["model:read"]})
|
||
got = next(x for x in r.json()["roles"] if x["code"] == "r1")
|
||
self.assertEqual(got["permissions"], ["model:read"])
|
||
|
||
def test_assigning_a_nonexistent_role_is_refused(self):
|
||
headers = self.login()
|
||
r = self.client.post(
|
||
"/api/v2/users", headers=headers,
|
||
json={"username": "ghost", "password": STRONG, "role": "没有这个角色"})
|
||
self.assertEqual(r.status_code, 400)
|
||
|
||
|
||
class UserApiTest(_Base):
|
||
def test_the_listing_never_leaks_password_material(self):
|
||
headers = self.login()
|
||
body = self.client.get("/api/v2/users", headers=headers).text
|
||
self.assertNotIn("password_salt", body)
|
||
self.assertNotIn("password_digest", body)
|
||
|
||
def test_a_weak_password_is_refused_on_create(self):
|
||
headers = self.login()
|
||
r = self.client.post(
|
||
"/api/v2/users", headers=headers,
|
||
json={"username": "weak", "password": "123", "role": "viewer"})
|
||
self.assertEqual(r.status_code, 400)
|
||
|
||
|
||
class StatsTest(_Base):
|
||
def test_model_call_stats_answers_the_cost_question(self):
|
||
for index in range(6):
|
||
self.db.log_model_call({
|
||
"device_id": "d1", "task_id": f"t{index}", "roles_version": 1,
|
||
"judge_mode": "shadow", "chosen": "模型A" if index % 2 else "模型B",
|
||
"judge": {"winner": "A", "score": 0.4 + index * 0.1,
|
||
"risk": "high" if index == 0 else "low"},
|
||
"candidates": [], "total_ms": 4000 + index * 100,
|
||
})
|
||
headers = self.login()
|
||
data = self.client.get("/api/v2/stats/model-calls?days=7", headers=headers).json()
|
||
self.assertEqual(data["total"], 6)
|
||
self.assertEqual(data["judged"], 6)
|
||
self.assertGreater(data["avg_score"], 0)
|
||
self.assertEqual({row["provider"] for row in data["chosen"]}, {"模型A", "模型B"})
|
||
self.assertEqual(data["risk"]["high"], 1)
|
||
self.assertTrue(data["score_buckets"])
|
||
self.assertGreater(data["avg_ms"], 0)
|
||
|
||
def test_the_call_log_lists_individual_calls_for_tracing(self):
|
||
"""聚合统计答不了"这一条到底怎么回的"——这个接口按条翻,供排查问题用。"""
|
||
self.db.log_model_call({
|
||
"device_id": "d1", "task_id": "t1", "roles_version": 1,
|
||
"judge_mode": "arbitrate", "chosen": "模型A",
|
||
"judge": {"winner": "A", "score": 0.9, "risk": "low"},
|
||
"candidates": [{"provider": "模型A", "text": "建议您注意休息"}],
|
||
"total_ms": 900, "customer_text": "我这几天总是失眠",
|
||
"reply_text": "建议您注意休息", "review_reason": "",
|
||
})
|
||
self.db.log_model_call({
|
||
"device_id": "d1", "task_id": "t2", "roles_version": 1,
|
||
"judge_mode": "arbitrate", "chosen": "模型A",
|
||
"judge": {"winner": "A", "score": 0.3, "risk": "high"},
|
||
"candidates": [], "total_ms": 500, "customer_text": "我这是不是得了糖尿病",
|
||
"reply_text": "建议就医", "review_reason": "命中审核规则「诊断」",
|
||
})
|
||
headers = self.login()
|
||
data = self.client.get(
|
||
"/api/v2/stats/model-calls/log?days=7", headers=headers
|
||
).json()
|
||
self.assertEqual(data["total"], 2)
|
||
first = data["items"][0] # id 倒序,t2 是后插入的那条
|
||
self.assertEqual(first["task_id"], "t2")
|
||
self.assertEqual(first["customer_text"], "我这是不是得了糖尿病")
|
||
self.assertEqual(first["review_reason"], "命中审核规则「诊断」")
|
||
self.assertEqual(first["candidates"], [])
|
||
|
||
def test_the_call_log_search_filters_by_keyword(self):
|
||
self.db.log_model_call({
|
||
"device_id": "d1", "task_id": "t1", "customer_text": "咨询糖尿病用药",
|
||
"reply_text": "请遵医嘱", "judge": {}, "candidates": [],
|
||
})
|
||
self.db.log_model_call({
|
||
"device_id": "d1", "task_id": "t2", "customer_text": "今天天气不错",
|
||
"reply_text": "是呀", "judge": {}, "candidates": [],
|
||
})
|
||
headers = self.login()
|
||
data = self.client.get(
|
||
"/api/v2/stats/model-calls/log?days=7&q=糖尿病", headers=headers
|
||
).json()
|
||
self.assertEqual(data["total"], 1)
|
||
self.assertEqual(data["items"][0]["task_id"], "t1")
|
||
|
||
def test_the_desktop_can_report_a_model_call_to_the_new_api(self):
|
||
"""8766 得能收调用留痕,否则桌面端改指过来就会静悄悄丢掉全部调用记录。
|
||
|
||
这是老后台(8765)退役前必须补齐的最后一块:配置同步早就有 v2 了,
|
||
只有这条还卡在 v1。少了它,"改个地址就切过来"会变成"切过来之后再也
|
||
查不到某句话是怎么回的"。
|
||
"""
|
||
resp = self.client.post(
|
||
"/api/v2/model/calls",
|
||
headers={"X-Desktop-Sync-Key": ab.DESKTOP_SYNC_KEY},
|
||
json={
|
||
"device_id": "desk-1", "task_id": "v2-1", "chosen": "模型A",
|
||
"judge": {"winner": "A", "score": 0.8, "risk": "low"},
|
||
"candidates": [{"provider": "模型A", "text": "在的"}],
|
||
"total_ms": 700, "customer_text": "在吗", "reply_text": "在的",
|
||
"review_reason": "",
|
||
},
|
||
)
|
||
self.assertEqual(resp.status_code, 200, resp.text)
|
||
found = self.db.list_model_calls(days=7)
|
||
self.assertEqual(found["total"], 1)
|
||
self.assertEqual(found["items"][0]["customer_text"], "在吗")
|
||
|
||
def test_reporting_a_call_without_the_sync_key_is_refused(self):
|
||
resp = self.client.post("/api/v2/model/calls", json={"task_id": "x"})
|
||
self.assertEqual(resp.status_code, 401)
|
||
|
||
def test_a_malformed_call_report_never_500s(self):
|
||
"""观测数据出问题不能让桌面端以为回复流程失败了。"""
|
||
resp = self.client.post(
|
||
"/api/v2/model/calls",
|
||
headers={"X-Desktop-Sync-Key": ab.DESKTOP_SYNC_KEY},
|
||
content=b"not json",
|
||
)
|
||
self.assertEqual(resp.status_code, 200)
|
||
|
||
def test_the_call_log_is_gated_behind_stats_read(self):
|
||
self.client.post(
|
||
"/api/v2/roles", headers=self.login(),
|
||
json={"code": "no-stats", "permissions": ["model:read"]})
|
||
headers = self.make_user("viewer-only", "no-stats")
|
||
resp = self.client.get("/api/v2/stats/model-calls/log", headers=headers)
|
||
self.assertEqual(resp.status_code, 403)
|
||
|
||
def test_the_audit_log_records_who_changed_what(self):
|
||
headers = self.login()
|
||
self.client.post(
|
||
"/api/v2/roles", headers=headers,
|
||
json={"code": "audited", "permissions": ["model:read"]})
|
||
entries = self.client.get("/api/v2/audit", headers=headers).json()["entries"]
|
||
actions = {row["action"] for row in entries}
|
||
self.assertIn("role.save", actions)
|
||
self.assertTrue(any(row["username"] == "admin" for row in entries))
|
||
|
||
|
||
class RoleCheckMigrationTest(TestCase):
|
||
"""把 users.role 从写死的 CHECK 迁到外键。
|
||
|
||
这段在真实库上炸过一次:SQLite 重建表必须先关外键,我漏了,
|
||
`auth_tokens.user_id` 引用 users,`DROP TABLE users` 直接
|
||
FOREIGN KEY constraint failed;更糟的是 executescript 处于自动提交状态,
|
||
半成品 `users_new` 留在库里,下次启动又撞同一个坑。
|
||
"""
|
||
|
||
OLD_SCHEMA = """
|
||
CREATE TABLE 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 CHECK(role IN ('admin','operator','viewer')),
|
||
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 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
|
||
);
|
||
"""
|
||
|
||
def _legacy_db(self, *, role="operator", leftover=False):
|
||
"""造一个改版前的库:老 users 表 + 有引用它的登录令牌。"""
|
||
root = Path(tempfile.mkdtemp())
|
||
path = root / "t.db"
|
||
con = sqlite3.connect(path)
|
||
con.executescript(self.OLD_SCHEMA)
|
||
con.execute(
|
||
"""INSERT INTO users (username,password_salt,password_digest,role,
|
||
created_at,updated_at)
|
||
VALUES ('laoban','s','d',?,'2026-01-01','2026-01-01')""",
|
||
(role,),
|
||
)
|
||
con.execute(
|
||
"""INSERT INTO auth_tokens (user_id,token_digest,kind,csrf_token,
|
||
created_at,expires_at,last_used_at)
|
||
VALUES (1,'digest','api','csrf','2026-01-01',9999999999,'2026-01-01')"""
|
||
)
|
||
if leftover:
|
||
# 模拟上一次失败留下的半成品
|
||
con.execute("CREATE TABLE users_new (id INTEGER PRIMARY KEY)")
|
||
con.execute("INSERT INTO users_new VALUES (1)")
|
||
con.commit()
|
||
con.close()
|
||
return path
|
||
|
||
def _inspect(self, path):
|
||
con = sqlite3.connect(path)
|
||
con.row_factory = sqlite3.Row
|
||
try:
|
||
tables = {r[0] for r in con.execute(
|
||
"SELECT name FROM sqlite_master WHERE type='table'")}
|
||
sql = con.execute(
|
||
"SELECT sql FROM sqlite_master WHERE name='users'").fetchone()[0]
|
||
con.execute("PRAGMA foreign_keys = ON")
|
||
return {
|
||
"has_leftover": "users_new" in tables,
|
||
"still_checked": "CHECK(role IN" in sql,
|
||
"users": [(r["username"], r["role"]) for r in
|
||
con.execute("SELECT username,role FROM users")],
|
||
"tokens": con.execute(
|
||
"SELECT COUNT(*) FROM auth_tokens").fetchone()[0],
|
||
"fk_broken": con.execute("PRAGMA foreign_key_check").fetchall(),
|
||
}
|
||
finally:
|
||
con.close()
|
||
|
||
def test_a_legacy_database_migrates_without_losing_anything(self):
|
||
path = self._legacy_db()
|
||
ab.Database(path).initialize("Admin@123456")
|
||
after = self._inspect(path)
|
||
self.assertFalse(after["still_checked"], "写死的 CHECK 应该已经没了")
|
||
self.assertFalse(after["has_leftover"])
|
||
self.assertIn(("laoban", "operator"), after["users"])
|
||
self.assertEqual(after["tokens"], 1, "登录令牌不能在迁移中丢掉")
|
||
self.assertEqual(after["fk_broken"], [])
|
||
|
||
def test_leftovers_from_a_failed_run_are_cleaned_up(self):
|
||
path = self._legacy_db(leftover=True)
|
||
ab.Database(path).initialize("Admin@123456")
|
||
after = self._inspect(path)
|
||
self.assertFalse(after["has_leftover"])
|
||
self.assertFalse(after["still_checked"])
|
||
self.assertEqual(after["tokens"], 1)
|
||
|
||
def test_an_unknown_role_is_downgraded_instead_of_crashing(self):
|
||
"""手工改库或早期版本会留下角色表里没有的值。收紧权限,不是崩掉服务。"""
|
||
path = self._legacy_db(role="admin")
|
||
con = sqlite3.connect(path)
|
||
# ignore_check_constraints 让我们能塞进一个 CHECK 不允许的值,
|
||
# 模拟"库被手工改过"或早期版本遗留
|
||
con.execute("PRAGMA ignore_check_constraints = ON")
|
||
con.execute(
|
||
"""INSERT INTO users (username,password_salt,password_digest,role,
|
||
created_at,updated_at)
|
||
VALUES ('guiji','s','d','早就删掉的角色','2026-01-01','2026-01-01')""")
|
||
con.commit()
|
||
con.close()
|
||
|
||
with mock.patch("builtins.print"):
|
||
ab.Database(path).initialize("Admin@123456")
|
||
after = self._inspect(path)
|
||
self.assertFalse(after["still_checked"])
|
||
self.assertIn(("guiji", "viewer"), after["users"])
|
||
self.assertIn(("laoban", "admin"), after["users"], "正常角色不该被动")
|
||
self.assertEqual(after["fk_broken"], [])
|
||
|
||
def test_running_it_twice_is_harmless(self):
|
||
path = self._legacy_db()
|
||
ab.Database(path).initialize("Admin@123456")
|
||
ab.Database(path).initialize("Admin@123456")
|
||
after = self._inspect(path)
|
||
self.assertFalse(after["has_leftover"])
|
||
self.assertEqual(after["tokens"], 1)
|
||
|
||
|
||
class _MigratedBase(_Base):
|
||
"""从老网页后台搬过来的那几块功能,共用的脚手架。"""
|
||
|
||
def setUp(self):
|
||
super().setUp()
|
||
self.admin = self.login()
|
||
|
||
def _with_perms(self, username: str, codes: list[str]) -> dict:
|
||
"""造一个只有指定权限码的用户,用来验证权限确实是按码判的。"""
|
||
role = f"r_{username}"
|
||
resp = self.client.post(
|
||
"/api/v2/roles",
|
||
headers=self.admin,
|
||
json={"code": role, "name": username, "permissions": codes},
|
||
)
|
||
self.assertEqual(resp.status_code, 200, resp.text)
|
||
return self.make_user(username, role)
|
||
|
||
|
||
class ConfigApiTest(_MigratedBase):
|
||
"""桌面端配置。这一组是从老网页后台整体搬过来的(能力开关 / 模型与身份 / MCP)。
|
||
|
||
校验走的是和老后台同一个 `validate_config_form`——这些测试同时也在保证两边
|
||
标准一致,不会出现"网页后台拦得住、新后台放得过"。
|
||
"""
|
||
|
||
def _config(self) -> dict:
|
||
body = self.client.get("/api/v2/config", headers=self.admin).json()["config"]
|
||
body.update({"AI_AGENT_NAME": "贴心管家", "AI_HOSPITAL_NAME": "甄养堂"})
|
||
return body
|
||
|
||
def test_model_settings_are_gone_from_the_config(self) -> None:
|
||
"""服务类型 / API 地址 / 密钥 / 模型名 / 温度 / tokens / 超时都搬去了角色编排。
|
||
|
||
同一件事有两个地方能配,迟早不一致,而且出问题时没人说得清以哪边为准。
|
||
密钥尤其:留在这份配置里就必须明文下发到每一台客户端。
|
||
"""
|
||
data = self.client.get("/api/v2/config", headers=self.admin).json()
|
||
for key in ab.RETIRED_CONFIG_KEYS:
|
||
with self.subTest(key=key):
|
||
self.assertNotIn(key, data["config"])
|
||
self.assertEqual(set(data["config"]), set(ab.CONFIG_KEYS))
|
||
self.assertIn("角色编排", data["model_settings_moved_to"])
|
||
|
||
def test_no_api_key_value_is_ever_returned(self) -> None:
|
||
"""库里可能还留着老记录的密钥,但一个字节都不能出接口。
|
||
|
||
断言的是**值**不出现,不是字段名——`retired_keys` 里就带着 "AI_API_KEY"
|
||
这个名字,那是给前端渲染指路说明用的,本身不是秘密。
|
||
"""
|
||
secret = "sk-left-over-from-the-old-schema"
|
||
with ab.Database(self.root / "t.db").connect() as db:
|
||
row = db.execute("SELECT config_json FROM model_config WHERE id=1").fetchone()
|
||
stored = json.loads(row["config_json"])
|
||
stored["AI_API_KEY"] = secret
|
||
db.execute(
|
||
"UPDATE model_config SET config_json=? WHERE id=1",
|
||
(json.dumps(stored, ensure_ascii=False),),
|
||
)
|
||
db.commit()
|
||
|
||
resp = self.client.get("/api/v2/config", headers=self.admin)
|
||
self.assertEqual(resp.status_code, 200)
|
||
self.assertNotIn(secret, resp.text)
|
||
self.assertNotIn("AI_API_KEY", resp.json()["config"])
|
||
|
||
# 桌面端同步同样不该再下发它——这正是"密钥不出后端"的那一步
|
||
sync = self.client.get(
|
||
"/api/v2/desktop/config",
|
||
headers={"X-Desktop-Sync-Key": ab.DESKTOP_SYNC_KEY},
|
||
)
|
||
self.assertNotIn(secret, sync.text)
|
||
|
||
def test_smuggling_a_retired_field_back_in_does_not_take_effect(self) -> None:
|
||
"""老客户端或手工调接口的人可能还带着这些字段。
|
||
|
||
必须被忽略,而不是悄悄写进库里——写进去就等于绕过了整个改版,而且下次
|
||
有人看这份配置会以为模型还是在这儿配的。
|
||
"""
|
||
before = json.loads(self.db.config()["config_json"]).get("AI_API_KEY")
|
||
resp = self.client.post(
|
||
"/api/v2/config",
|
||
headers=self.admin,
|
||
json={
|
||
**self._config(),
|
||
"AI_API_KEY": "sk-sneaky",
|
||
"AI_API_BASE": "https://evil.example/v1",
|
||
"AI_TEMPERATURE": 1.9,
|
||
},
|
||
)
|
||
self.assertEqual(resp.status_code, 200, resp.text)
|
||
stored = json.loads(self.db.config()["config_json"])
|
||
self.assertEqual(stored.get("AI_API_KEY"), before, "老值原样保留,不被覆盖")
|
||
self.assertNotEqual(stored.get("AI_API_BASE"), "https://evil.example/v1")
|
||
|
||
def test_saving_bumps_the_version_so_clients_notice(self) -> None:
|
||
before = self.client.get("/api/v2/config", headers=self.admin).json()["version"]
|
||
resp = self.client.post(
|
||
"/api/v2/config", headers=self.admin, json=self._config()
|
||
)
|
||
self.assertEqual(resp.status_code, 200, resp.text)
|
||
self.assertEqual(resp.json()["version"], before + 1)
|
||
|
||
def test_invalid_values_are_refused_with_a_usable_reason(self) -> None:
|
||
cases = [
|
||
({"AI_CONTEXT_MAX_ROUNDS": 0}, "AI_CONTEXT_MAX_ROUNDS"),
|
||
({"AI_CONTEXT_MAX_ROUNDS": 51}, "AI_CONTEXT_MAX_ROUNDS"),
|
||
({"AI_MCP_MAX_ROUNDS": 99}, "AI_MCP_MAX_ROUNDS"),
|
||
({"AI_AGENT_NAME": ""}, "客服名称"),
|
||
({"AI_HOSPITAL_NAME": ""}, "机构名称"),
|
||
]
|
||
for override, expected in cases:
|
||
with self.subTest(**override):
|
||
resp = self.client.post(
|
||
"/api/v2/config",
|
||
headers=self.admin,
|
||
json={**self._config(), **override},
|
||
)
|
||
self.assertEqual(resp.status_code, 400, resp.text)
|
||
self.assertIn(expected, resp.json()["detail"])
|
||
|
||
def test_the_config_payload_matches_the_legacy_key_set(self) -> None:
|
||
"""字段少一个,桌面端就少读一个设置——而且是静悄悄地少。"""
|
||
data = self.client.get("/api/v2/config", headers=self.admin).json()
|
||
self.assertEqual(set(data["config"]), set(ab.CONFIG_KEYS))
|
||
|
||
def test_config_read_and_write_are_separate_permissions(self) -> None:
|
||
"""排查线上问题的人要能看配置,但不该能改。"""
|
||
viewer = self._with_perms("cfgviewer", ["config:read"])
|
||
self.assertEqual(
|
||
self.client.get("/api/v2/config", headers=viewer).status_code, 200
|
||
)
|
||
self.assertEqual(
|
||
self.client.post(
|
||
"/api/v2/config", headers=viewer, json=self._config()
|
||
).status_code,
|
||
403,
|
||
)
|
||
|
||
|
||
class ReleaseApiTest(_MigratedBase):
|
||
_BLANK = {
|
||
"latest_version": "1.0.0",
|
||
"download_url": "",
|
||
"release_notes": "",
|
||
"force_upgrade": False,
|
||
}
|
||
|
||
def test_publishing_a_version_records_it(self) -> None:
|
||
resp = self.client.post(
|
||
"/api/v2/release",
|
||
headers=self.admin,
|
||
json={
|
||
"latest_version": "2.5.0",
|
||
"download_url": "https://dl.example.com/setup.exe",
|
||
"release_notes": "修了审核放行",
|
||
"force_upgrade": True,
|
||
},
|
||
)
|
||
self.assertEqual(resp.status_code, 200, resp.text)
|
||
current = self.client.get("/api/v2/release", headers=self.admin).json()
|
||
self.assertEqual(current["latest_version"], "2.5.0")
|
||
self.assertTrue(current["force_upgrade"])
|
||
self.assertEqual(current["updated_by"], "admin")
|
||
|
||
def test_a_v_prefix_is_stripped_not_rejected(self) -> None:
|
||
"""人习惯写 v2.5.0。存进去必须是 2.5.0,否则客户端比对版本号永远不相等。"""
|
||
resp = self.client.post(
|
||
"/api/v2/release",
|
||
headers=self.admin,
|
||
json={**self._BLANK, "latest_version": "v2.5.0"},
|
||
)
|
||
self.assertEqual(resp.status_code, 200, resp.text)
|
||
self.assertEqual(resp.json()["latest_version"], "2.5.0")
|
||
|
||
def test_force_upgrade_without_a_download_url_is_refused(self) -> None:
|
||
"""开了强制升级又没给下载地址,等于把所有客户端堵死在原地。"""
|
||
resp = self.client.post(
|
||
"/api/v2/release",
|
||
headers=self.admin,
|
||
json={**self._BLANK, "latest_version": "2.6.0", "force_upgrade": True},
|
||
)
|
||
self.assertEqual(resp.status_code, 400)
|
||
self.assertIn("下载地址", resp.json()["detail"])
|
||
|
||
def test_bad_versions_and_urls_are_refused(self) -> None:
|
||
cases = [
|
||
({"latest_version": "随便写"}, "版本号"),
|
||
({"latest_version": "2.7.0", "download_url": "ftp://x/y"}, "http"),
|
||
({"latest_version": "2.8.0", "release_notes": "字" * 4001}, "4000"),
|
||
]
|
||
for override, expected in cases:
|
||
with self.subTest(**override):
|
||
resp = self.client.post(
|
||
"/api/v2/release", headers=self.admin, json={**self._BLANK, **override}
|
||
)
|
||
self.assertEqual(resp.status_code, 400, resp.text)
|
||
self.assertIn(expected, resp.json()["detail"])
|
||
|
||
def test_reading_needs_config_read_but_writing_needs_release_write(self) -> None:
|
||
viewer = self._with_perms("relviewer", ["config:read"])
|
||
self.assertEqual(
|
||
self.client.get("/api/v2/release", headers=viewer).status_code, 200
|
||
)
|
||
self.assertEqual(
|
||
self.client.post(
|
||
"/api/v2/release",
|
||
headers=viewer,
|
||
json={**self._BLANK, "latest_version": "3.0.0"},
|
||
).status_code,
|
||
403,
|
||
)
|
||
|
||
|
||
class ModelTestApiTest(_MigratedBase):
|
||
def _provider(self, **extra) -> None:
|
||
body = {
|
||
"id": "p1",
|
||
"name": "主答题",
|
||
"kind": "openai",
|
||
# 9 号端口没人监听,连接会被立刻拒绝——不需要等超时
|
||
"base_url": "http://127.0.0.1:9/v1",
|
||
"api_key": "sk-secret-in-db",
|
||
"model": "gpt-4o-mini",
|
||
"capabilities": "text",
|
||
}
|
||
body.update(extra)
|
||
resp = self.client.post("/api/v2/models", headers=self.admin, json=body)
|
||
self.assertEqual(resp.status_code, 200, resp.text)
|
||
|
||
def test_testing_a_catalog_model_never_needs_the_key_from_the_browser(self) -> None:
|
||
"""前端只送 id,密钥由后端从库里解密——明文不在网络上多走一趟。"""
|
||
self._provider()
|
||
resp = self.client.post(
|
||
"/api/v2/models/test",
|
||
headers=self.admin,
|
||
json={"provider_id": "p1", "timeout_seconds": 5},
|
||
)
|
||
self.assertEqual(resp.status_code, 200, resp.text)
|
||
self.assertFalse(resp.json()["result"]["ok"])
|
||
self.assertNotIn("sk-secret-in-db", resp.text)
|
||
|
||
def test_an_unreachable_endpoint_reports_200_with_a_failed_result(self) -> None:
|
||
"""测试跑成功了、结论是"连不上"——这是 200,不是 5xx。
|
||
|
||
返 5xx 会被前端的通用错误拦截器接管,弹一句"服务器内部错误",
|
||
把真正有用的诊断信息盖掉。
|
||
"""
|
||
self._provider()
|
||
resp = self.client.post(
|
||
"/api/v2/models/test",
|
||
headers=self.admin,
|
||
json={"provider_id": "p1", "timeout_seconds": 5},
|
||
)
|
||
self.assertEqual(resp.status_code, 200)
|
||
result = resp.json()["result"]
|
||
self.assertFalse(result["ok"])
|
||
self.assertTrue(result["message"], "失败必须给出原因,不能是空字符串")
|
||
self.assertTrue(result["endpoint"], "要说清是往哪个地址发的")
|
||
|
||
def test_a_missing_model_is_404(self) -> None:
|
||
resp = self.client.post(
|
||
"/api/v2/models/test", headers=self.admin, json={"provider_id": "没有这个"}
|
||
)
|
||
self.assertEqual(resp.status_code, 404)
|
||
|
||
def test_claude_says_why_instead_of_pretending_to_test(self) -> None:
|
||
"""老后台的测试器只认 openai/dify/comfyui。硬套 claude 只会得到误导性的 404。"""
|
||
self._provider(id="c1", kind="claude", base_url="https://api.anthropic.com")
|
||
resp = self.client.post(
|
||
"/api/v2/models/test", headers=self.admin, json={"provider_id": "c1"}
|
||
)
|
||
self.assertEqual(resp.status_code, 400)
|
||
self.assertIn("Claude", resp.json()["detail"])
|
||
|
||
def test_testing_requires_write_because_it_spends_money(self) -> None:
|
||
"""只读用户能看模型清单,但不该能拿着公司的密钥往外发请求。"""
|
||
self._provider()
|
||
reader = self._with_perms("mreader", ["model:read"])
|
||
resp = self.client.post(
|
||
"/api/v2/models/test", headers=reader, json={"provider_id": "p1"}
|
||
)
|
||
self.assertEqual(resp.status_code, 403)
|
||
|
||
def test_a_test_is_written_to_the_audit_log(self) -> None:
|
||
self._provider()
|
||
self.client.post(
|
||
"/api/v2/models/test",
|
||
headers=self.admin,
|
||
json={"provider_id": "p1", "timeout_seconds": 5},
|
||
)
|
||
actions = [
|
||
row["action"]
|
||
for row in self.client.get("/api/v2/audit", headers=self.admin).json()["entries"]
|
||
]
|
||
self.assertIn("model.test", actions)
|
||
|
||
|
||
class DesktopSyncApiTest(_MigratedBase):
|
||
"""桌面端同步。
|
||
|
||
老后台的 `/api/v1/desktop/config` 继续可用,已经装出去的客户端不受影响;
|
||
这里等价提供一份,好让 8765 将来能退役。返回结构必须和老的完全一致——差一
|
||
个键,客户端就少读一样东西,而且是静悄悄地少。
|
||
"""
|
||
|
||
_SHAPE = {
|
||
"version", "updated_at", "updated_by", "config", "models", "roles",
|
||
"release",
|
||
# 网关地址由后台算好下发——桌面端只配一个后台地址,其余自动
|
||
"gateway",
|
||
}
|
||
|
||
def test_a_wrong_sync_key_is_rejected(self) -> None:
|
||
self.assertEqual(self.client.get("/api/v2/desktop/config").status_code, 401)
|
||
self.assertEqual(
|
||
self.client.get(
|
||
"/api/v2/desktop/config", headers={"X-Desktop-Sync-Key": "wrong"}
|
||
).status_code,
|
||
401,
|
||
)
|
||
|
||
def test_the_payload_has_the_same_shape_as_the_legacy_endpoint(self) -> None:
|
||
resp = self.client.get(
|
||
"/api/v2/desktop/config",
|
||
headers={"X-Desktop-Sync-Key": ab.DESKTOP_SYNC_KEY},
|
||
)
|
||
self.assertEqual(resp.status_code, 200, resp.text)
|
||
data = resp.json()
|
||
self.assertEqual(set(data), self._SHAPE)
|
||
self.assertEqual(
|
||
set(data["release"]),
|
||
{"latest_version", "download_url", "release_notes", "force_upgrade", "updated_at"},
|
||
)
|
||
self.assertEqual(set(data["gateway"]), {"enabled", "url"})
|
||
self.assertTrue(
|
||
data["gateway"]["url"].startswith("http"),
|
||
"下发的网关地址必须能直接用,不能是空串——空串会让客户端以为没配网关",
|
||
)
|
||
|
||
def test_model_keys_go_out_masked_never_in_the_clear(self) -> None:
|
||
"""桌面端拿到的是遮罩值;真正的调用要走模型网关。"""
|
||
self.client.post(
|
||
"/api/v2/models",
|
||
headers=self.admin,
|
||
json={
|
||
"id": "p1",
|
||
"name": "主",
|
||
"kind": "openai",
|
||
"base_url": "https://x/v1",
|
||
"api_key": "sk-never-ship-this",
|
||
"model": "m",
|
||
"capabilities": "text",
|
||
},
|
||
)
|
||
resp = self.client.get(
|
||
"/api/v2/desktop/config",
|
||
headers={"X-Desktop-Sync-Key": ab.DESKTOP_SYNC_KEY},
|
||
)
|
||
self.assertNotIn("sk-never-ship-this", resp.text)
|
||
for item in resp.json()["models"]:
|
||
self.assertNotIn("api_key", item)
|
||
self.assertIn("api_key_masked", item)
|
||
|
||
def test_a_published_release_reaches_the_desktop_payload(self) -> None:
|
||
"""后台点了发布,客户端下一次同步就该看到——这条链断了没人会立刻发现。"""
|
||
self.client.post(
|
||
"/api/v2/release",
|
||
headers=self.admin,
|
||
json={
|
||
"latest_version": "4.1.0",
|
||
"download_url": "https://dl.example.com/a.exe",
|
||
"release_notes": "note",
|
||
"force_upgrade": True,
|
||
},
|
||
)
|
||
data = self.client.get(
|
||
"/api/v2/desktop/config",
|
||
headers={"X-Desktop-Sync-Key": ab.DESKTOP_SYNC_KEY},
|
||
).json()
|
||
self.assertEqual(data["release"]["latest_version"], "4.1.0")
|
||
self.assertTrue(data["release"]["force_upgrade"])
|
||
|
||
|
||
|
||
class ReviewRuleValidationTest(TestCase):
|
||
"""`validate_review_rules`:选择性审核规则存进库之前的最后一道关卡。
|
||
|
||
这一层校验错了,后果不是接口报错这么轻——是某条规则悄悄失效(关键词是空的,
|
||
永远不会命中)或者桌面端读配置直接炸掉(形状不对)。宁可在保存这一步多问
|
||
几句,也不要让坏数据流到能自动发消息给真实客户的那台机器上。
|
||
"""
|
||
|
||
def test_a_normal_rule_round_trips(self) -> None:
|
||
cleaned = ab.validate_review_rules([
|
||
{"label": "诊断", "keywords": ["确诊", "诊断"], "enabled": True},
|
||
])
|
||
self.assertEqual(len(cleaned), 1)
|
||
self.assertEqual(cleaned[0]["label"], "诊断")
|
||
self.assertEqual(cleaned[0]["keywords"], ["确诊", "诊断"])
|
||
self.assertTrue(cleaned[0]["enabled"])
|
||
self.assertTrue(cleaned[0]["id"])
|
||
|
||
def test_the_root_must_be_a_list(self) -> None:
|
||
with self.assertRaises(ValueError) as caught:
|
||
ab.validate_review_rules({"label": "诊断"})
|
||
self.assertIn("数组", str(caught.exception))
|
||
|
||
def test_an_item_that_is_not_an_object_is_refused(self) -> None:
|
||
with self.assertRaises(ValueError):
|
||
ab.validate_review_rules(["诊断"])
|
||
|
||
def test_a_rule_without_a_label_is_refused(self) -> None:
|
||
with self.assertRaises(ValueError) as caught:
|
||
ab.validate_review_rules([{"keywords": ["确诊"]}])
|
||
self.assertIn("名称", str(caught.exception))
|
||
|
||
def test_a_rule_with_no_keywords_is_refused(self) -> None:
|
||
"""空关键词的规则不是"温和一点",是"这条规则永远不会命中"——必须当错误处理。"""
|
||
with self.assertRaises(ValueError) as caught:
|
||
ab.validate_review_rules([{"label": "诊断", "keywords": []}])
|
||
self.assertIn("诊断", str(caught.exception))
|
||
self.assertIn("关键词", str(caught.exception))
|
||
|
||
def test_keywords_that_are_all_blank_strings_still_count_as_empty(self) -> None:
|
||
with self.assertRaises(ValueError):
|
||
ab.validate_review_rules([{"label": "诊断", "keywords": [" ", ""]}])
|
||
|
||
def test_keywords_must_be_a_list(self) -> None:
|
||
with self.assertRaises(ValueError):
|
||
ab.validate_review_rules([{"label": "诊断", "keywords": "确诊"}])
|
||
|
||
def test_missing_ids_are_auto_generated_and_deduplicated(self) -> None:
|
||
cleaned = ab.validate_review_rules([
|
||
{"label": "诊断", "keywords": ["a"]},
|
||
{"label": "诊断", "keywords": ["b"]},
|
||
])
|
||
ids = [rule["id"] for rule in cleaned]
|
||
self.assertEqual(len(ids), len(set(ids)), "两条同名规则不能得到同一个 id")
|
||
|
||
def test_enabled_defaults_to_true_when_omitted(self) -> None:
|
||
cleaned = ab.validate_review_rules([{"label": "诊断", "keywords": ["a"]}])
|
||
self.assertTrue(cleaned[0]["enabled"])
|
||
|
||
def test_enabled_false_is_respected(self) -> None:
|
||
cleaned = ab.validate_review_rules(
|
||
[{"label": "诊断", "keywords": ["a"], "enabled": False}]
|
||
)
|
||
self.assertFalse(cleaned[0]["enabled"])
|
||
|
||
def test_too_many_rules_is_refused(self) -> None:
|
||
rules = [{"label": f"规则{i}", "keywords": ["x"]} for i in range(51)]
|
||
with self.assertRaises(ValueError) as caught:
|
||
ab.validate_review_rules(rules)
|
||
self.assertIn("50", str(caught.exception))
|
||
|
||
def test_too_many_keywords_on_one_rule_is_refused(self) -> None:
|
||
with self.assertRaises(ValueError):
|
||
ab.validate_review_rules(
|
||
[{"label": "诊断", "keywords": [f"词{i}" for i in range(31)]}]
|
||
)
|
||
|
||
def test_an_absurdly_long_label_is_refused(self) -> None:
|
||
with self.assertRaises(ValueError):
|
||
ab.validate_review_rules([{"label": "长" * 41, "keywords": ["a"]}])
|
||
|
||
|
||
class ReviewRuleConfigApiTest(_MigratedBase):
|
||
"""审核规则走的是和 MCP 服务器一样的配置同步管道,端到端确认接得上。"""
|
||
|
||
def _base_config(self) -> dict:
|
||
body = self.client.get("/api/v2/config", headers=self.admin).json()["config"]
|
||
body.update({"AI_AGENT_NAME": "贴心管家", "AI_HOSPITAL_NAME": "甄养堂"})
|
||
return body
|
||
|
||
def test_the_seeded_defaults_come_with_three_starter_rules(self) -> None:
|
||
"""全新的库不该是空的——刚上线这个功能时,起步规则不能一条都没有。"""
|
||
config = self.client.get("/api/v2/config", headers=self.admin).json()["config"]
|
||
labels = {rule["label"] for rule in config["AI_REVIEW_RULES"]}
|
||
self.assertEqual(labels, {"诊断", "用药调整", "投诉退款"})
|
||
|
||
def test_saving_a_rule_set_round_trips_through_the_api(self) -> None:
|
||
payload = {
|
||
**self._base_config(),
|
||
"AI_REVIEW_RULES": [
|
||
{"label": "测试规则", "keywords": ["敏感词"], "enabled": True},
|
||
],
|
||
}
|
||
resp = self.client.post("/api/v2/config", headers=self.admin, json=payload)
|
||
self.assertEqual(resp.status_code, 200, resp.text)
|
||
fetched = self.client.get("/api/v2/config", headers=self.admin).json()["config"]
|
||
self.assertEqual(len(fetched["AI_REVIEW_RULES"]), 1)
|
||
self.assertEqual(fetched["AI_REVIEW_RULES"][0]["label"], "测试规则")
|
||
|
||
def test_an_invalid_rule_is_rejected_with_a_readable_reason(self) -> None:
|
||
payload = {**self._base_config(), "AI_REVIEW_RULES": [{"label": "", "keywords": []}]}
|
||
resp = self.client.post("/api/v2/config", headers=self.admin, json=payload)
|
||
self.assertEqual(resp.status_code, 400)
|
||
self.assertIn("名称", resp.json()["detail"])
|
||
|
||
def test_the_desktop_sync_endpoints_ship_the_rules_too(self) -> None:
|
||
payload = {
|
||
**self._base_config(),
|
||
"AI_REVIEW_RULES": [{"label": "同步测试", "keywords": ["x"], "enabled": True}],
|
||
}
|
||
self.client.post("/api/v2/config", headers=self.admin, json=payload)
|
||
resp = self.client.get(
|
||
"/api/v2/desktop/config",
|
||
headers={"X-Desktop-Sync-Key": ab.DESKTOP_SYNC_KEY},
|
||
)
|
||
rules = resp.json()["config"]["AI_REVIEW_RULES"]
|
||
self.assertEqual([r["label"] for r in rules], ["同步测试"])
|
||
|
||
class ConfigNullPoisoningTest(_MigratedBase):
|
||
"""库里存的 config_json 可能带着 null 或缺键——这不是假设,是踩出来的真事。
|
||
|
||
`ai_settings.json` 曾经把未设置的开关存成 JSON `null`,那份文件当初被用来
|
||
初始化这张表,null 就这样原样搬进了 `model_config.config_json`。后来
|
||
`AI_GATEWAY_URL` 加进 CONFIG_KEYS 时,建库更早的那些记录里压根没有这一项。
|
||
|
||
两种情况过去都被 `{key: stored.get(key) for key in CONFIG_KEYS}` 直接透传给
|
||
前端:`.get()` 对"键不存在"和"值是 None"给出的都是 None。前端把 None 存进
|
||
JS 表单变成 `null`,用户点保存,POST 出去的还是 null,撞上接口的类型
|
||
校验——`Input should be a valid boolean, input: null`,一个从字面上完全看
|
||
不出病根在哪的错误。
|
||
"""
|
||
|
||
def _poison(self, **overrides) -> None:
|
||
"""把库里的 config_json 换成一份"老结构":缺键 + null 都占一个。"""
|
||
payload = {
|
||
"AI_ENABLED": True,
|
||
"AI_DEVELOPMENT_MODE": None,
|
||
"AI_UI_GUARD_ENABLED": None,
|
||
"AI_USE_VISION": False,
|
||
"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": [],
|
||
# 故意不写 AI_GATEWAY_URL:老记录里压根没有这一列
|
||
}
|
||
payload.update(overrides)
|
||
with sqlite3.connect(self.db.path) as con:
|
||
con.execute(
|
||
"UPDATE model_config SET config_json=? WHERE id=1",
|
||
(json.dumps(payload),),
|
||
)
|
||
con.commit()
|
||
|
||
def test_get_config_never_leaks_null_for_a_missing_or_null_field(self) -> None:
|
||
self._poison()
|
||
resp = self.client.get("/api/v2/config", headers=self.admin)
|
||
self.assertEqual(resp.status_code, 200)
|
||
config = resp.json()["config"]
|
||
self.assertEqual(
|
||
{k for k, v in config.items() if v is None},
|
||
set(),
|
||
"GET 不该再让任何字段以 null 的样子出现在响应里",
|
||
)
|
||
# 缺键和显式 null 都要落到各自的出厂默认值,而不是随便一个假值
|
||
self.assertIs(config["AI_DEVELOPMENT_MODE"], ab.CONFIG_DEFAULTS["AI_DEVELOPMENT_MODE"])
|
||
self.assertIs(config["AI_UI_GUARD_ENABLED"], ab.CONFIG_DEFAULTS["AI_UI_GUARD_ENABLED"])
|
||
self.assertEqual(config["AI_GATEWAY_URL"], ab.CONFIG_DEFAULTS["AI_GATEWAY_URL"])
|
||
|
||
def test_round_tripping_what_get_returns_never_422s(self) -> None:
|
||
"""这就是这个 bug 真实发生的顺序:前端 GET 拿表单初值,原样 POST 回去。"""
|
||
self._poison()
|
||
fetched = self.client.get("/api/v2/config", headers=self.admin).json()["config"]
|
||
resp = self.client.post("/api/v2/config", headers=self.admin, json=fetched)
|
||
self.assertEqual(resp.status_code, 200, resp.text)
|
||
|
||
def test_a_stale_cached_page_that_still_sends_null_is_tolerated(self) -> None:
|
||
"""浏览器里开着修复之前拉取的旧页面,表单里还揣着当时的 null。
|
||
|
||
这条防的不是"数据脏",是"客户端代码旧"——GET 修好之后,已经打开的标签页
|
||
不会自动重新拉取,用户下一次点保存,发出去的可能仍然是 null。
|
||
"""
|
||
good = self.client.get("/api/v2/config", headers=self.admin).json()["config"]
|
||
stale = dict(good)
|
||
stale.update({
|
||
"AI_DEVELOPMENT_MODE": None,
|
||
"AI_UI_GUARD_ENABLED": None,
|
||
"AI_GATEWAY_URL": None,
|
||
})
|
||
resp = self.client.post("/api/v2/config", headers=self.admin, json=stale)
|
||
self.assertEqual(resp.status_code, 200, resp.text)
|
||
stored = json.loads(self.db.config()["config_json"])
|
||
# null 被换成默认值写进去,不是被当成"什么都没传"而悄悄丢弃
|
||
self.assertEqual(stored["AI_DEVELOPMENT_MODE"], ab.CONFIG_DEFAULTS["AI_DEVELOPMENT_MODE"])
|
||
self.assertEqual(stored["AI_GATEWAY_URL"], ab.CONFIG_DEFAULTS["AI_GATEWAY_URL"])
|
||
|
||
def test_the_legacy_desktop_sync_endpoint_is_covered_too(self) -> None:
|
||
"""下发给桌面端的那一份,也走同一段"透传 CONFIG_KEYS"的代码。
|
||
|
||
老后台退役后这段搬到了模块级的 `desktop_config_payload`,接口和数据层
|
||
共用同一份实现——这里直接对着它测,不用拉起 HTTP 服务。
|
||
"""
|
||
self._poison()
|
||
payload = ab.desktop_config_payload(self.db)
|
||
self.assertEqual(
|
||
{k for k, v in payload["config"].items() if v is None},
|
||
set(),
|
||
)
|
||
|
||
def test_the_v2_desktop_sync_endpoint_is_covered_too(self) -> None:
|
||
self._poison()
|
||
resp = self.client.get(
|
||
"/api/v2/desktop/config",
|
||
headers={"X-Desktop-Sync-Key": ab.DESKTOP_SYNC_KEY},
|
||
)
|
||
self.assertEqual(resp.status_code, 200)
|
||
config = resp.json()["config"]
|
||
self.assertEqual({k for k, v in config.items() if v is None}, set())
|
||
|
||
def test_a_deliberately_set_false_is_not_mistaken_for_unset(self) -> None:
|
||
"""补默认值的判据必须是"None",不能是"假值"——否则关掉的开关会被强制打开。"""
|
||
self._poison(AI_ENABLED=False, AI_MCP_SERVERS=[])
|
||
config = self.client.get("/api/v2/config", headers=self.admin).json()["config"]
|
||
self.assertFalse(config["AI_ENABLED"], "用户明确关掉的开关不该被默认值救活")
|
||
self.assertEqual(config["AI_MCP_SERVERS"], [])
|
||
|
||
|
||
class StaticMountTest(TestCase):
|
||
"""前端构建产物存在时,API 服务必须照样能起来。
|
||
|
||
这条是真踩过才补的:`create_app` 原先只判断 dist 目录在不在,然后写死挂载
|
||
`dist/assets`。仓库里一直没有构建产物,所以那个分支从来没执行过;等真的跑了
|
||
一次 `vite build`(输出的是 js/ 和 css/,没有 assets/),StaticFiles 在构造时
|
||
直接抛 RuntimeError——启动即崩,连日志都来不及打。
|
||
"""
|
||
|
||
def setUp(self) -> None:
|
||
self.root = Path(tempfile.mkdtemp())
|
||
self.addCleanup(shutil.rmtree, self.root, ignore_errors=True)
|
||
self.dist = self.root / "dist"
|
||
|
||
def _client(self) -> TestClient:
|
||
with mock.patch.object(admin_api, "_frontend_dist", return_value=self.dist):
|
||
app = admin_api.create_app(self.root / "t.db")
|
||
return TestClient(app)
|
||
|
||
def _write(self, rel: str, body: str = "x") -> None:
|
||
target = self.dist / rel
|
||
target.parent.mkdir(parents=True, exist_ok=True)
|
||
target.write_text(body, encoding="utf-8")
|
||
|
||
def test_a_build_without_an_assets_folder_still_boots(self) -> None:
|
||
"""vite 默认输出 js/ 和 css/,没有 assets/。这种产物必须能起来。"""
|
||
self._write("index.html", "<html></html>")
|
||
self._write("js/app.js", "//")
|
||
client = self._client()
|
||
self.assertEqual(client.get("/api/v2/health").status_code, 200)
|
||
self.assertEqual(client.get("/js/app.js").status_code, 200)
|
||
|
||
def test_an_empty_dist_directory_is_ignored(self) -> None:
|
||
"""只有一个空 dist 目录时不挂 SPA 兜底路由。
|
||
|
||
挂了的话每个未知路径都会去回一个不存在的 index.html,报出来的是
|
||
FileNotFoundError,比老老实实 404 难查得多。
|
||
"""
|
||
self.dist.mkdir()
|
||
client = self._client()
|
||
self.assertEqual(client.get("/api/v2/health").status_code, 200)
|
||
self.assertEqual(client.get("/whatever").status_code, 404)
|
||
|
||
def test_all_of_the_usual_asset_folders_get_mounted(self) -> None:
|
||
self._write("index.html", "<html></html>")
|
||
for name in ("assets", "js", "css"):
|
||
self._write(f"{name}/f.txt", name)
|
||
client = self._client()
|
||
for name in ("assets", "js", "css"):
|
||
with self.subTest(folder=name):
|
||
resp = client.get(f"/{name}/f.txt")
|
||
self.assertEqual(resp.status_code, 200)
|
||
self.assertEqual(resp.text, name)
|
||
|
||
def test_an_unknown_api_path_is_404_not_the_spa_page(self) -> None:
|
||
"""兜底路由不能把打错的接口地址伪装成一个正常页面。"""
|
||
self._write("index.html", "<html></html>")
|
||
client = self._client()
|
||
self.assertEqual(client.get("/api/v2/nope").status_code, 404)
|
||
# 前端路由是浏览器端的,服务器上没有这个文件,但刷新必须还能出页面
|
||
self.assertEqual(client.get("/system/roles").status_code, 200)
|
||
|
||
|
||
|
||
def setUpModule():
|
||
# 别让测试读到开发机上的真实桌面端配置——配过模型网关的机器会让
|
||
# `ai_chat.current_provider()` 整体改走网关分支,一大片无关测试跟着变行为。
|
||
local_state_redirect.start()
|
||
|
||
|
||
def tearDownModule():
|
||
local_state_redirect.stop()
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|