Files
kefu/wechat_rpa/test_model_router.py
T
2026-08-27 14:04:28 +08:00

1081 lines
48 KiB
Python

# -*- coding: utf-8 -*-
"""多模型出口、编排与后台模型清单的测试。
三件事必须钉死:
1. Provider 是显式参数,不再靠改全局——改全局在这个多线程进程里必然串号;
2. 编排的四条兜底路径(全挂 / 只回一路 / 裁判挂 / 影子模式)都不能让客户等不到回复;
3. 模型密钥落库必须是密文,接口永不回显明文。
"""
import json
import sqlite3
import tempfile
import threading
import time
from pathlib import Path
from unittest import TestCase, main, mock
from test_support import local_state_redirect
import ai_chat
import model_router as mr
import secret_box
from ai_chat import Provider
def _p(**kw) -> Provider:
base = dict(kind="openai", base_url="https://x/v1", api_key="k", model="m")
base.update(kw)
return Provider(**base)
class ProviderResolutionTest(TestCase):
def test_none_falls_back_to_the_global_config(self):
with mock.patch.object(ai_chat.ai_config, "AI_API_BASE", "https://g/v1"), \
mock.patch.object(ai_chat.ai_config, "AI_MODEL", "global-model"):
self.assertEqual(ai_chat._resolve(None).model, "global-model")
def test_an_explicit_provider_never_reads_globals(self):
# 这条是整个重构的意义:并发问两个模型时不能互相污染
with mock.patch.object(ai_chat.ai_config, "AI_MODEL", "global-model"):
self.assertEqual(ai_chat._resolve(_p(model="mine")).model, "mine")
def test_provider_is_immutable_so_threads_cannot_corrupt_it(self):
provider = _p()
with self.assertRaises(Exception):
provider.model = "changed"
def test_kind_is_detected_from_the_url_when_not_declared(self):
self.assertEqual(
ai_chat._provider_type(_p(kind="auto", base_url="http://x/v1/chat-messages")),
"dify",
)
self.assertEqual(
ai_chat._provider_type(_p(kind="auto", base_url="https://api.anthropic.com")),
"claude",
)
self.assertEqual(
ai_chat._provider_type(_p(kind="auto", base_url="https://api.deepseek.com")),
"openai",
)
def test_each_kind_lands_on_its_own_endpoint(self):
self.assertEqual(
ai_chat._completions_url(_p(kind="claude", base_url="https://api.anthropic.com")),
"https://api.anthropic.com/v1/messages",
)
self.assertEqual(
ai_chat._completions_url(_p(kind="dify", base_url="http://ai/v1")),
"http://ai/v1/chat-messages",
)
self.assertEqual(
ai_chat._completions_url(_p(kind="openai", base_url="https://api.x.com")),
"https://api.x.com/chat/completions",
)
def test_claude_uses_its_own_auth_headers(self):
headers = ai_chat._headers(_p(kind="claude", api_key="sk-ant"))
self.assertEqual(headers["x-api-key"], "sk-ant")
self.assertIn("anthropic-version", headers)
self.assertNotIn("Authorization", headers)
def test_other_kinds_keep_bearer_auth(self):
headers = ai_chat._headers(_p(kind="openai", api_key="sk-1"))
self.assertEqual(headers["Authorization"], "Bearer sk-1")
self.assertNotIn("x-api-key", headers)
class ClaudeShapeTest(TestCase):
"""Anthropic 的三处不兼容,照搬 OpenAI 的 payload 必然 400。"""
def test_system_is_lifted_out_of_the_message_list(self):
system, turns = ai_chat._claude_messages([
{"role": "system", "content": "你是客服"},
{"role": "user", "content": "在吗"},
])
self.assertEqual(system, "你是客服")
self.assertEqual([t["role"] for t in turns], ["user"])
def test_adjacent_same_role_turns_are_merged(self):
_system, turns = ai_chat._claude_messages([
{"role": "user", "content": "在吗"},
{"role": "user", "content": "还在吗"},
])
self.assertEqual(len(turns), 1)
self.assertEqual(len(turns[0]["content"]), 2)
def test_the_first_turn_must_be_user(self):
_system, turns = ai_chat._claude_messages([
{"role": "assistant", "content": "您好"},
{"role": "user", "content": "在吗"},
])
self.assertEqual(turns[0]["role"], "user")
def test_an_empty_conversation_still_produces_a_valid_body(self):
_system, turns = ai_chat._claude_messages([])
self.assertEqual(turns[0]["role"], "user")
def test_images_become_base64_source_blocks(self):
_system, turns = ai_chat._claude_messages([
{"role": "user", "content": [
{"type": "text", "text": "看图"},
{"type": "image_url",
"image_url": {"url": "data:image/png;base64,QUJD"}},
]},
])
block = turns[0]["content"][1]
self.assertEqual(block["type"], "image")
self.assertEqual(block["source"]["media_type"], "image/png")
self.assertEqual(block["source"]["data"], "QUJD")
def test_max_tokens_is_always_sent(self):
captured = {}
class _Resp:
ok = True
status_code = 200
@staticmethod
def json():
return {"content": [{"type": "text", "text": "好的"}]}
def fake_post(url, **kwargs):
captured.update(kwargs.get("json") or {})
return _Resp()
with mock.patch.object(ai_chat, "_post_with_retry", side_effect=fake_post):
out = ai_chat._claude_completion(
[{"role": "user", "content": "在吗"}],
_p(kind="claude", max_tokens=0),
)
# Anthropic 的 max_tokens 是必填且必须 >= 1,给 0 会直接 400
self.assertGreaterEqual(captured["max_tokens"], 1)
self.assertEqual(out["content"], "好的")
class VerdictParsingTest(TestCase):
def test_reads_winner_score_and_risk(self):
v = mr.parse_verdict('{"winner":"B","score":0.9,"risk":"low","reason":"更短"}')
self.assertEqual((v.winner, v.risk), ("B", "low"))
self.assertAlmostEqual(v.score, 0.9)
self.assertTrue(v.participated)
def test_an_unknown_winner_is_unusable(self):
self.assertIsNone(mr.parse_verdict('{"winner":"C","score":0.9}'))
def test_a_non_json_answer_is_unusable(self):
self.assertIsNone(mr.parse_verdict("我觉得 A 更好"))
def test_score_is_clamped_and_junk_risk_becomes_unknown(self):
v = mr.parse_verdict('{"winner":"A","score":9,"risk":"爆炸"}')
self.assertEqual(v.score, 1.0)
self.assertEqual(v.risk, "unknown")
class OrchestrationTest(TestCase):
"""四条兜底路径,任何一条断了客户就等不到回复。"""
def setUp(self):
self.a = _p(name="模型A")
self.b = _p(name="模型B")
self.judge = _p(name="裁判")
# 影子评审的并发闸是进程级的,上一条用例挂住的线程会把它耗光,
# 后面的用例就永远看不到后台线程。每条用例先复位。
mr._SHADOW_LIMIT = threading.Semaphore(2)
def _run(self, replies, verdict_json, mode):
def fake_reply(**kw):
value = replies.get(kw["provider"].name)
if isinstance(value, Exception):
raise value
return value
def fake_chat(messages, tools=None, provider=None):
if isinstance(verdict_json, Exception):
raise verdict_json
return {"content": verdict_json}
with mock.patch.object(ai_chat, "get_ai_reply", side_effect=fake_reply), \
mock.patch.object(ai_chat, "_chat_completion", side_effect=fake_chat), \
mock.patch("builtins.print"):
return mr.answer(
chat_text="一天吃几次",
answer_providers=[self.a, self.b],
judge_provider=self.judge,
judge_mode=mode,
)
_WIN_B = '{"winner":"B","score":0.9,"risk":"low","reason":"B 更短"}'
def test_shadow_mode_never_changes_what_gets_sent(self):
out = self._run({"模型A": "甲说", "模型B": "乙说"}, self._WIN_B, "shadow")
self.assertEqual(out["reply"], "甲说")
def test_shadow_mode_does_not_make_the_customer_wait_for_the_judge(self):
"""影子模式下评审结果不改变任何东西,就绝不该挡在回复路径上。"""
released = threading.Event()
def slow_judge(messages, tools=None, provider=None):
released.wait(timeout=2)
return {"content": self._WIN_B}
with mock.patch.object(ai_chat, "get_ai_reply", return_value="甲说"), mock.patch.object(ai_chat, "_chat_completion", side_effect=slow_judge), mock.patch("builtins.print"):
started = time.monotonic()
out = mr.answer(
chat_text="x",
answer_providers=[self.a, self.b],
judge_provider=self.judge,
judge_mode="shadow",
)
elapsed = time.monotonic() - started
released.set()
self.assertEqual(out["reply"], "甲说")
self.assertLess(elapsed, 0.5, "裁判还没返回,回复就该已经拿到了")
# 同步返回里没有评分——它是后台补上来的
self.assertFalse(out["judge"]["participated"])
def test_the_shadow_verdict_arrives_through_the_callback(self):
landed = threading.Event()
received = {}
def capture(later):
received.update(later)
landed.set()
with mock.patch.object(ai_chat, "get_ai_reply", return_value="甲说"), mock.patch.object(
ai_chat, "_chat_completion",
return_value={"content": self._WIN_B}), mock.patch("builtins.print"):
mr.answer(
chat_text="x",
answer_providers=[self.a, self.b],
judge_provider=self.judge,
judge_mode="shadow",
on_verdict=capture,
)
self.assertTrue(landed.wait(timeout=5), "后台评审没有回调")
self.assertTrue(received["judge"]["participated"])
self.assertAlmostEqual(received["judge"]["score"], 0.9)
self.assertEqual(received["reply"], "甲说")
def test_the_shadow_judge_never_delays_process_exit(self):
# 线程池会在 atexit 里 join;守护线程不会。观测数据不该拦着程序关闭。
with mock.patch.object(ai_chat, "get_ai_reply", return_value="甲说"), mock.patch.object(
ai_chat, "_chat_completion",
side_effect=lambda *a, **k: time.sleep(30)), mock.patch("builtins.print"):
mr.answer(
chat_text="x",
answer_providers=[self.a],
judge_provider=self.judge,
judge_mode="shadow",
)
workers = [t for t in threading.enumerate() if t.name == "ai-shadow-judge"]
self.assertTrue(workers)
self.assertTrue(all(t.daemon for t in workers))
def test_score_only_mode_also_keeps_the_primary_reply(self):
out = self._run({"模型A": "甲说", "模型B": "乙说"}, self._WIN_B, "score_only")
self.assertEqual(out["reply"], "甲说")
self.assertAlmostEqual(out["judge"]["score"], 0.9)
def test_arbitrate_mode_sends_the_winner(self):
out = self._run({"模型A": "甲说", "模型B": "乙说"}, self._WIN_B, "arbitrate")
self.assertEqual(out["reply"], "乙说")
self.assertEqual(out["chosen"], "模型B")
def test_a_dead_judge_falls_back_to_the_primary(self):
out = self._run(
{"模型A": "甲说", "模型B": "乙说"}, RuntimeError("裁判挂了"), "arbitrate"
)
self.assertEqual(out["reply"], "甲说")
self.assertFalse(out["judge"]["participated"])
def test_unparseable_judgement_falls_back_to_the_primary(self):
out = self._run({"模型A": "甲说", "模型B": "乙说"}, "随便说说", "arbitrate")
self.assertEqual(out["reply"], "甲说")
self.assertFalse(out["judge"]["participated"])
def test_one_dead_candidate_still_answers(self):
out = self._run(
{"模型A": RuntimeError("超时"), "模型B": "乙说"}, self._WIN_B, "arbitrate"
)
self.assertEqual(out["reply"], "乙说")
def test_all_dead_returns_empty_rather_than_a_canned_line(self):
out = self._run(
{"模型A": RuntimeError("超时"), "模型B": RuntimeError("超时")},
self._WIN_B,
"arbitrate",
)
self.assertEqual(out["reply"], "")
self.assertTrue(all(c["error"] for c in out["candidates"]))
def test_a_blank_reply_counts_as_a_failure(self):
out = self._run({"模型A": " ", "模型B": "乙说"}, self._WIN_B, "arbitrate")
self.assertEqual(out["reply"], "乙说")
def test_every_candidate_is_recorded_for_audit(self):
out = self._run({"模型A": "甲说", "模型B": "乙说"}, self._WIN_B, "arbitrate")
names = sorted(c["provider"] for c in out["candidates"])
self.assertEqual(names, ["模型A", "模型B"])
self.assertTrue(all(c["latency_ms"] >= 0 for c in out["candidates"]))
def test_candidates_are_asked_in_parallel_not_one_after_another(self):
import threading
import time
live = []
peak = []
lock = threading.Lock()
def slow(**kw):
with lock:
live.append(1)
peak.append(len(live))
time.sleep(0.15)
with lock:
live.pop()
return "答案"
with mock.patch.object(ai_chat, "get_ai_reply", side_effect=slow), \
mock.patch("builtins.print"):
started = time.monotonic()
mr.answer(chat_text="x", answer_providers=[self.a, self.b])
elapsed = time.monotonic() - started
self.assertEqual(max(peak), 2, "两路必须同时在飞,串行会把延迟翻倍")
self.assertLess(elapsed, 0.28)
def test_no_provider_configured_falls_back_to_the_global_one(self):
with mock.patch.object(ai_chat, "get_ai_reply", return_value="全局答"), \
mock.patch("builtins.print"):
out = mr.answer(chat_text="x")
self.assertEqual(out["reply"], "全局答")
def test_an_unknown_mode_degrades_to_shadow(self):
out = self._run({"模型A": "甲说", "模型B": "乙说"}, self._WIN_B, "乱填的")
self.assertEqual(out["judge_mode"], "shadow")
self.assertEqual(out["reply"], "甲说")
class RolesResolutionTest(TestCase):
CATALOG = [
{"id": "d1", "name": "Dify", "kind": "dify", "base_url": "http://a/v1",
"api_key": "k", "enabled": True},
{"id": "c1", "name": "Claude", "kind": "claude",
"base_url": "https://api.anthropic.com", "api_key": "k", "enabled": True},
{"id": "j1", "name": "裁判", "kind": "openai", "base_url": "https://o/v1",
"api_key": "k", "enabled": True},
{"id": "off", "name": "停用的", "kind": "openai", "base_url": "https://o/v1",
"api_key": "k", "enabled": False},
{"id": "img", "name": "文生图", "kind": "comfyui",
"base_url": "http://127.0.0.1:8188", "enabled": True},
]
def test_builds_answer_providers_and_judge(self):
answers, judge, mode = mr.providers_from_roles(
self.CATALOG,
{"answer_ids": "d1,c1", "judge_id": "j1", "judge_mode": "arbitrate"},
)
self.assertEqual([p.name for p in answers], ["Dify", "Claude"])
self.assertEqual(judge.name, "裁判")
self.assertEqual(mode, "arbitrate")
def test_a_disabled_model_is_skipped_not_fatal(self):
answers, _judge, _mode = mr.providers_from_roles(
self.CATALOG, {"answer_ids": "d1,off", "judge_id": "j1"}
)
self.assertEqual([p.name for p in answers], ["Dify"])
def test_a_missing_model_degrades_instead_of_breaking_everything(self):
# 后台删掉一个模型不该让整条回复链路停摆
answers, judge, _mode = mr.providers_from_roles(
self.CATALOG, {"answer_ids": "d1,已删除", "judge_id": "也没了"}
)
self.assertEqual([p.name for p in answers], ["Dify"])
self.assertIsNone(judge)
def test_comfyui_can_never_be_an_answering_candidate(self):
# ComfyUI 是文生图工作流引擎,当不了聊天模型
answers, judge, _mode = mr.providers_from_roles(
self.CATALOG, {"answer_ids": "img,d1", "judge_id": "img"}
)
self.assertEqual([p.name for p in answers], ["Dify"])
self.assertIsNone(judge)
def test_an_unknown_judge_mode_degrades_to_shadow(self):
_a, _j, mode = mr.providers_from_roles(
self.CATALOG, {"answer_ids": "d1", "judge_mode": "全自动"}
)
self.assertEqual(mode, "shadow")
class SecretBoxTest(TestCase):
def setUp(self):
self.root = Path(tempfile.mkdtemp())
self.key = secret_box.load_or_create_key(self.root)
def test_round_trip(self):
token = secret_box.encrypt("sk-ant-XYZ", self.key)
self.assertNotIn("sk-ant-XYZ", token)
self.assertEqual(secret_box.decrypt(token, self.key), "sk-ant-XYZ")
def test_a_tampered_ciphertext_raises_instead_of_returning_garbage(self):
token = secret_box.encrypt("sk-ant-XYZ", self.key)
broken = ("A" if token[0] != "A" else "B") + token[1:]
with self.assertRaises(secret_box.SecretBoxError):
secret_box.decrypt(broken, self.key)
def test_a_wrong_master_key_cannot_decrypt(self):
other = Path(tempfile.mkdtemp())
token = secret_box.encrypt("sk-ant-XYZ", self.key)
with self.assertRaises(secret_box.SecretBoxError):
secret_box.decrypt(token, secret_box.load_or_create_key(other))
def test_the_key_is_stable_across_calls(self):
self.assertEqual(secret_box.load_or_create_key(self.root), self.key)
def test_empty_stays_empty(self):
self.assertEqual(secret_box.encrypt("", self.key), "")
self.assertEqual(secret_box.decrypt("", self.key), "")
def test_mask_never_shows_the_whole_key(self):
self.assertNotIn("sk-ant-XYZ-1234", secret_box.masked("sk-ant-XYZ-1234"))
self.assertEqual(secret_box.masked("short"), "*****")
class ModelCatalogTest(TestCase):
def setUp(self):
import admin_backend
self.root = Path(tempfile.mkdtemp())
self.db = admin_backend.Database(self.root / "t.db")
self.db.initialize("InitialAdmin123")
self.uid = self.db.authenticate("admin", "InitialAdmin123")["id"]
self.db.save_model_provider(
{"id": "d1", "name": "Dify", "kind": "dify", "base_url": "http://a/v1",
"api_key": "app-SECRET"}, self.uid, "1.1.1.1")
self.db.save_model_provider(
{"id": "j1", "name": "裁判", "kind": "openai",
"base_url": "https://o/v1", "api_key": "sk-JUDGE",
"model": "gpt-4o-mini"}, self.uid, "1.1.1.1")
def _raw_key(self, provider_id):
con = sqlite3.connect(self.root / "t.db")
try:
return con.execute(
"SELECT api_key_enc FROM model_providers WHERE id=?", (provider_id,)
).fetchone()[0]
finally:
con.close()
def test_keys_are_encrypted_at_rest(self):
self.assertNotIn("app-SECRET", self._raw_key("d1"))
def test_the_listing_never_returns_a_plaintext_key(self):
for item in self.db.model_providers():
self.assertNotIn("api_key", item)
self.assertNotIn("api_key_enc", item)
self.assertIn("api_key_masked", item)
def test_secrets_are_available_only_on_explicit_request(self):
found = {i["id"]: i["api_key"] for i in self.db.model_providers(include_secrets=True)}
self.assertEqual(found["d1"], "app-SECRET")
def test_saving_without_a_key_keeps_the_existing_one(self):
self.db.save_model_provider(
{"id": "d1", "name": "改名了", "kind": "dify", "base_url": "http://a/v1"},
self.uid, "1.1.1.1")
found = {i["id"]: i["api_key"] for i in self.db.model_providers(include_secrets=True)}
self.assertEqual(found["d1"], "app-SECRET")
def test_an_unsupported_kind_is_rejected(self):
with self.assertRaises(ValueError):
self.db.save_model_provider(
{"id": "x", "kind": "gemini", "base_url": "http://x"}, self.uid, "1.1.1.1")
def test_a_referenced_model_cannot_be_deleted(self):
self.db.save_model_roles(
{"answer_ids": "d1", "judge_id": "j1"}, self.uid, "1.1.1.1")
with self.assertRaises(ValueError):
self.db.delete_model_provider("d1", self.uid, "1.1.1.1")
def test_an_unreferenced_model_can_be_deleted(self):
self.db.save_model_provider(
{"id": "tmp", "kind": "openai", "base_url": "http://x"}, self.uid, "1.1.1.1")
self.db.delete_model_provider("tmp", self.uid, "1.1.1.1")
self.assertNotIn("tmp", {i["id"] for i in self.db.model_providers()})
def test_roles_are_versioned_so_a_rollback_is_just_an_older_version(self):
first = self.db.save_model_roles(
{"answer_ids": "d1", "judge_id": "j1", "judge_mode": "shadow"},
self.uid, "1.1.1.1")
second = self.db.save_model_roles(
{"answer_ids": "d1,j1", "judge_id": "j1", "judge_mode": "arbitrate"},
self.uid, "1.1.1.1")
self.assertGreater(second, first)
self.assertEqual(self.db.model_roles()["judge_mode"], "arbitrate")
def test_roles_reject_unknown_models_and_empty_answers(self):
with self.assertRaises(ValueError):
self.db.save_model_roles({"answer_ids": "没有的"}, self.uid, "1.1.1.1")
with self.assertRaises(ValueError):
self.db.save_model_roles({"answer_ids": ""}, self.uid, "1.1.1.1")
def test_roles_reject_an_unknown_judge_mode(self):
with self.assertRaises(ValueError):
self.db.save_model_roles(
{"answer_ids": "d1", "judge_mode": "全自动"}, self.uid, "1.1.1.1")
def test_call_logging_never_raises_into_the_reply_path(self):
# 观测数据落库失败绝不能让回复流程以为自己失败了
with mock.patch.object(type(self.db), "connect", side_effect=RuntimeError("盘满了")), \
mock.patch("builtins.print"):
self.db.log_model_call({"device_id": "d", "judge": {}})
def test_the_catalog_rides_along_with_the_config_push(self):
self.db.save_model_roles({"answer_ids": "d1", "judge_id": "j1"}, self.uid, "1.1.1.1")
import admin_backend
# 老网页后台退役后,这份结构由模块级的 `desktop_config_payload` 统一给出
payload = admin_backend.desktop_config_payload(self.db)
self.assertIn("models", payload)
self.assertIn("roles", payload)
self.assertEqual(payload["roles"]["answer_ids"], "d1")
# 下发的清单里绝不能出现明文密钥
self.assertNotIn("app-SECRET", str(payload))
class BotPlanTest(TestCase):
"""桌面端怎么决定"本轮问谁"。"""
def _bot(self, plan=None, judge_enabled=False, judge_mode="shadow"):
from wechat_bot import WeChatBot
bot = WeChatBot.__new__(WeChatBot)
bot._backend_model_plan = plan or {"models": [], "roles": {}}
self._judge_enabled = judge_enabled
self._judge_mode = judge_mode
return bot
def _plan(self, bot):
with mock.patch.object(
ai_chat.ai_config, "AI_JUDGE_ENABLED", self._judge_enabled, create=True
), mock.patch.object(
ai_chat.ai_config, "AI_JUDGE_MODE", self._judge_mode, create=True
), mock.patch("builtins.print"):
return bot._model_plan()
def test_no_backend_plan_uses_the_local_provider_alone(self):
answers, judge, mode = self._plan(self._bot())
self.assertEqual(len(answers), 1)
self.assertIsNone(judge)
self.assertEqual(mode, "shadow")
def test_masked_backend_keys_are_not_usable_so_local_stays(self):
# 后台下发的密钥是遮罩值——密钥不出后端。没有网关之前只能用本机出口。
plan = {
"models": [
{"id": "d1", "name": "Dify", "kind": "dify",
"base_url": "http://a/v1", "api_key_masked": "app-***", "enabled": True},
],
"roles": {"answer_ids": "d1", "judge_id": "", "judge_mode": "arbitrate"},
}
answers, _judge, mode = self._plan(self._bot(plan))
self.assertEqual(len(answers), 1)
self.assertEqual(mode, "shadow", "拿不到可用密钥时不能贸然进入仲裁模式")
def test_backend_providers_with_real_keys_are_adopted(self):
plan = {
"models": [
{"id": "d1", "name": "Dify", "kind": "dify", "base_url": "http://a/v1",
"api_key": "app-REAL", "enabled": True},
{"id": "c1", "name": "Claude", "kind": "claude",
"base_url": "https://api.anthropic.com", "api_key": "sk-REAL",
"enabled": True},
{"id": "j1", "name": "裁判", "kind": "openai", "base_url": "https://o/v1",
"api_key": "sk-J", "enabled": True},
],
"roles": {"answer_ids": "d1,c1", "judge_id": "j1", "judge_mode": "arbitrate"},
}
answers, judge, mode = self._plan(self._bot(plan))
self.assertEqual([p.name for p in answers], ["Dify", "Claude"])
self.assertEqual(judge.name, "裁判")
self.assertEqual(mode, "arbitrate")
def test_the_local_provider_can_serve_as_judge_when_enabled(self):
# 没有专门裁判出口时,用本机这一个照样能量出绝对分——这是影子期的全部意义
answers, judge, mode = self._plan(self._bot(judge_enabled=True))
self.assertEqual(len(answers), 1)
self.assertIsNotNone(judge)
self.assertEqual(mode, "shadow")
def test_a_broken_plan_degrades_instead_of_raising(self):
answers, judge, _mode = self._plan(self._bot({"models": "坏数据", "roles": None}))
self.assertEqual(len(answers), 1)
self.assertIsNone(judge)
def test_single_provider_without_judge_keeps_the_original_call_path(self):
"""没配多模型也没配裁判时,一点额外开销都不该引入。"""
from wechat_bot import WeChatBot
bot = WeChatBot.__new__(WeChatBot)
bot._backend_model_plan = {"models": [], "roles": {}}
bot._call_model_with_observer = mock.Mock(
side_effect=lambda cb, *a, **kw: cb(*a, **kw)
)
with mock.patch.object(
ai_chat.ai_config, "AI_JUDGE_ENABLED", False, create=True
), mock.patch.object(ai_chat, "get_ai_reply", return_value="直接回答") as direct, mock.patch.object(mr, "answer") as orchestrated, mock.patch("builtins.print"):
out = bot._orchestrated_reply(b"fp", chat_text="在吗")
self.assertEqual(out, "直接回答")
direct.assert_called_once()
orchestrated.assert_not_called()
def test_the_judge_verdict_is_kept_on_the_task_for_audit(self):
from wechat_bot import WeChatBot
bot = WeChatBot.__new__(WeChatBot)
state = {}
bot._pending_reply_state = mock.Mock(return_value=state)
bot._persist_pending_replies = mock.Mock(return_value=True)
bot._backend_model_plan = {"models": [], "roles": {}}
with mock.patch("builtins.print"):
bot._record_model_call(b"fp", {
"chosen": "模型B",
"candidates": [{"provider": "模型A"}, {"provider": "模型B"}],
"judge": {"winner": "B", "score": 0.9, "risk": "low",
"reason": "更短", "participated": False},
"total_ms": 5200,
})
self.assertEqual(state["model_chosen"], "模型B")
self.assertAlmostEqual(state["judge_score"], 0.9)
self.assertEqual(state["judge_risk"], "low")
self.assertEqual(len(state["model_candidates"]), 2)
def test_a_failed_report_never_breaks_the_reply(self):
from wechat_bot import WeChatBot
bot = WeChatBot.__new__(WeChatBot)
bot._pending_reply_state = mock.Mock(return_value=None)
bot._backend_model_plan = {"models": [], "roles": {}}
import backend_client
with mock.patch.object(
backend_client, "report_model_call", side_effect=RuntimeError("断网")
), mock.patch("builtins.print"):
bot._record_model_call(b"fp", {
"judge": {"participated": True, "score": 0.5}, "candidates": []
}) # 不抛异常即为通过
class PinnedServerUrlTest(TestCase):
"""明确指定过的后台地址,不许被自动发现悄悄改掉。
同机后台的端口是动态的,所以自动发现要能覆盖保存的地址——但只能覆盖"没人
明确指定过"的那种。覆盖掉用户填的地址,表现是登录成功、随后每个请求 401,
而界面上完全看不出地址被换过。
"""
def setUp(self):
import backend_client
self.bc = backend_client
self.root = Path(tempfile.mkdtemp())
self._saved = (backend_client.CONNECTION_FILE, backend_client.RUNTIME_FILE)
backend_client.CONNECTION_FILE = self.root / "connection.json"
backend_client.RUNTIME_FILE = self.root / "runtime.json"
def tearDown(self):
self.bc.CONNECTION_FILE, self.bc.RUNTIME_FILE = self._saved
def _publish_local_backend(self, port):
import json
import os
self.bc.RUNTIME_FILE.write_text(json.dumps({
"pid": os.getpid(), "host": "127.0.0.1", "port": port,
"server_url": f"http://127.0.0.1:{port}",
"local_sync_token": "tok", "started_at": "2026-01-01T00:00:00+08:00",
}), encoding="utf-8")
def test_an_unpinned_local_url_still_follows_the_discovered_port(self):
import socket
listener = socket.socket()
listener.bind(("127.0.0.1", 0))
listener.listen(1)
self.addCleanup(listener.close)
port = listener.getsockname()[1]
self._publish_local_backend(port)
self.bc.save_settings({**self.bc.default_settings(),
"server_url": "http://127.0.0.1:1"})
self.assertEqual(
self.bc.load_settings()["server_url"], f"http://127.0.0.1:{port}")
def test_a_pinned_url_is_never_rewritten(self):
import socket
listener = socket.socket()
listener.bind(("127.0.0.1", 0))
listener.listen(1)
self.addCleanup(listener.close)
self._publish_local_backend(listener.getsockname()[1])
self.bc.save_settings({**self.bc.default_settings(),
"server_url": "http://127.0.0.1:54321",
"server_url_pinned": True})
self.assertEqual(
self.bc.load_settings()["server_url"], "http://127.0.0.1:54321")
def test_login_pins_the_url_it_was_given(self):
self.bc.save_settings({**self.bc.default_settings()})
self.assertFalse(self.bc.load_settings()["server_url_pinned"])
class SettingsPersistenceTest(TestCase):
"""save/load 都按 default_settings() 的键白名单过滤。
没在默认值里声明的键会被**静默丢掉**——写进去不报错,读出来是空的。
新增配置项时最容易漏这一步,而漏了之后代码看着完全正常,功能却是死的。
"""
def setUp(self):
import backend_client
self.bc = backend_client
self.root = Path(tempfile.mkdtemp())
self._saved = (backend_client.CONNECTION_FILE, backend_client.RUNTIME_FILE)
backend_client.CONNECTION_FILE = self.root / "conn.json"
backend_client.RUNTIME_FILE = self.root / "rt.json"
def tearDown(self):
self.bc.CONNECTION_FILE, self.bc.RUNTIME_FILE = self._saved
def test_the_gateway_config_survives_a_round_trip(self):
settings = self.bc.load_settings()
settings["gateway"] = {"enabled": True, "url": "http://127.0.0.1:8770/v1/answer"}
self.bc.save_settings(settings)
self.assertEqual(
self.bc.load_settings()["gateway"]["url"],
"http://127.0.0.1:8770/v1/answer",
)
def test_the_cached_model_plan_survives_a_round_trip(self):
settings = self.bc.load_settings()
settings["model_plan"] = {
"models": [{"id": "d1"}], "roles": {"answer_ids": "d1"}, "synced_at": "now",
}
self.bc.save_settings(settings)
self.assertEqual(self.bc.cached_model_plan()["roles"]["answer_ids"], "d1")
def test_the_device_id_is_generated_once_and_kept(self):
first = self.bc.device_id()
self.assertTrue(first.startswith("dev-"))
self.assertEqual(self.bc.device_id(), first)
def test_every_documented_key_is_declared_in_the_defaults(self):
"""新增配置项必须同时加进 default_settings(),否则读写都是空转。"""
defaults = self.bc.default_settings()
for key in ("gateway", "model_plan", "device_id", "server_url_pinned"):
self.assertIn(key, defaults, f"{key} 没在默认值里声明,会被静默丢掉")
class GatewayAsProviderTest(TestCase):
"""模型调用统一走网关,路由由后台的「角色编排」决定。
这是这次改版的核心:密钥加密存在后端、接口只返回遮罩值,桌面端从设计上就
拿不到明文,没法自己直连模型。把网关做成一种 `kind`,是为了让
`_chat_completion` 往下的所有调用点——MCP 多轮、视觉、受限分类器——一行都
不用改。这些测试盯的就是"一行都不用改"这句话是真的。
"""
def setUp(self) -> None:
self.provider = ai_chat.Provider(
kind=ai_chat.GATEWAY_KIND,
base_url="http://127.0.0.1:8770/v1/answer",
api_key="sync-key",
name="模型网关",
)
self.sent = []
def _urlopen(self, replies):
"""假冒网关。replies 是按顺序返回的响应体。"""
queue = list(replies)
class _Response:
def __init__(self, payload):
self._payload = json.dumps(payload, ensure_ascii=False).encode()
def read(self):
return self._payload
def __enter__(self):
return self
def __exit__(self, *exc):
return False
def fake(request, timeout=None):
self.sent.append(
{
"url": request.full_url,
"headers": dict(request.headers),
"body": json.loads(request.data.decode()),
"timeout": timeout,
}
)
return _Response(queue.pop(0) if queue else {"reply": ""})
return mock.patch("urllib.request.urlopen", fake)
def test_a_plain_reply_goes_through_the_gateway(self) -> None:
with self._urlopen([{"reply": "您好,两次一天。", "chosen": "主答题"}]):
message = ai_chat._chat_completion(
[{"role": "user", "content": "一天吃几次"}], provider=self.provider
)
self.assertEqual(message["content"], "您好,两次一天。")
self.assertEqual(self.sent[0]["headers"]["X-desktop-sync-key"], "sync-key")
def test_the_full_message_list_reaches_the_gateway(self) -> None:
"""系统提示词和历史必须一起送过去。
之前网关只收到一句裸的用户文本——人格、机构名、前几轮对话全丢了,回复
质量会静悄悄地垮掉,而且没有任何报错。
"""
messages = [
{"role": "system", "content": "你是甄养堂的贴心管家"},
{"role": "user", "content": "上次说的那个"},
{"role": "assistant", "content": "灵芝孢子粉"},
{"role": "user", "content": "一天吃几次"},
]
with self._urlopen([{"reply": "两次"}]):
ai_chat._chat_completion(messages, provider=self.provider)
body = self.sent[0]["body"]
self.assertEqual(body["messages"], messages, "整串消息必须原样送到")
self.assertEqual(
body["customer_text"], "一天吃几次", "裁判要拿最后一条用户消息打分"
)
def test_the_call_carries_a_device_id_and_a_fresh_task_id(self) -> None:
"""网关落库的那一行要能对上桌面端后来补报的审核原因——两边靠同一个 task_id 关联。
以前这里什么都不带,`model_calls.task_id` 只能存桌面端自己拿会话指纹拼的
假值,还会在同一个会话的下一轮撞车。真正的调用标识必须由这次请求自己生成。
"""
with self._urlopen([{"reply": "两次"}]):
message = ai_chat._chat_completion(
[{"role": "user", "content": "一天吃几次"}], provider=self.provider
)
sent_body = self.sent[0]["body"]
self.assertTrue(sent_body.get("task_id"), "必须真的带上 task_id")
self.assertIn("X-device-id", self.sent[0]["headers"])
self.assertEqual(message["_gateway"]["task_id"], sent_body["task_id"], "留痕里的 id 要和发出去的一致")
def test_two_separate_calls_never_reuse_the_same_task_id(self) -> None:
with self._urlopen([{"reply": "答一"}, {"reply": "答二"}]):
ai_chat._chat_completion(
[{"role": "user", "content": "第一句"}], provider=self.provider
)
ai_chat._chat_completion(
[{"role": "user", "content": "第二句"}], provider=self.provider
)
first_id = self.sent[0]["body"]["task_id"]
second_id = self.sent[1]["body"]["task_id"]
self.assertTrue(first_id and second_id)
self.assertNotEqual(first_id, second_id)
def test_a_normal_reply_is_tagged_as_a_customer_chat(self) -> None:
with self._urlopen([{"reply": "两次"}]):
ai_chat._chat_completion(
[{"role": "user", "content": "一天吃几次"}], provider=self.provider
)
self.assertEqual(self.sent[0]["body"]["purpose"], "chat")
def test_the_vision_classifier_is_tagged_as_an_internal_guard_call(self) -> None:
"""界面识别不是"发给客户的话",必须打上 guard。
这就是那个真实的 bug:守卫每轮轮询都要问一次模型,条数是真实对话的几十
倍。不打标记的话,调用日志里全是布局识别的 JSON、客户消息一律为空,本来
用来查"这句话怎么回出来的"那张表直接废掉;裁判分布量的也变成了"布局认得
准不准",而不是回复质量。
"""
with self._urlopen([{"reply": '{"state":"chat_ready"}'}]):
ai_chat._call_vision_classifier(
"判断界面", b"\x89PNG-fake", user="layout-guard", provider=self.provider
)
self.assertEqual(self.sent[0]["body"]["purpose"], "guard")
def test_tool_calls_come_back_intact_for_the_mcp_loop(self) -> None:
"""MCP 多轮靠 tool_calls 驱动。丢了它,模型说"我要查一下"就变成一句空回复。"""
calls = [
{
"id": "c1",
"type": "function",
"function": {"name": "查客户", "arguments": '{"phone":"138"}'},
}
]
with self._urlopen([{"reply": "", "tool_calls": calls}]):
message = ai_chat._chat_completion(
[{"role": "user", "content": "查一下我的订单"}],
tools=[{"type": "function", "function": {"name": "查客户"}}],
provider=self.provider,
)
self.assertEqual(message["tool_calls"], calls)
self.assertIn("tools", self.sent[0]["body"], "工具清单要送给网关")
def test_a_reply_without_tool_calls_omits_the_field(self) -> None:
"""没有工具调用时不要塞一个空的 tool_calls。
MCP 循环判的是 `msg.get("tool_calls")`,空列表是假值所以碰巧也能过;
但下一轮会把 `{"tool_calls": []}` 原样发回给模型,有些实现会因此报 400。
"""
with self._urlopen([{"reply": "好的"}]):
message = ai_chat._chat_completion(
[{"role": "user", "content": "嗯"}], provider=self.provider
)
self.assertNotIn("tool_calls", message)
def test_the_orchestration_trace_rides_along(self) -> None:
"""选了谁、裁判打了多少分——这些要能落进队列记录,出事时得解释得清。"""
with self._urlopen([
{
"reply": "两次",
"chosen": "主答题",
"judge": {"participated": True, "score": 0.82, "risk": "low"},
"judge_mode": "shadow",
"roles_version": 7,
"total_ms": 1830,
}
]):
message = ai_chat._chat_completion(
[{"role": "user", "content": "一天吃几次"}], provider=self.provider
)
trace = message["_gateway"]
self.assertEqual(trace["chosen"], "主答题")
self.assertEqual(trace["judge"]["score"], 0.82)
self.assertEqual(trace["roles_version"], 7)
def test_a_gateway_failure_is_raised_not_swallowed(self) -> None:
"""网关挂了必须报出来。
绝不能悄悄回落到本机出口:生产环境本机根本没有密钥,回落只会得到一个
更难懂的 401,把"网关挂了"这条真信息盖掉。
"""
import urllib.error
def boom(request, timeout=None):
raise urllib.error.URLError("connection refused")
with mock.patch("urllib.request.urlopen", boom):
with self.assertRaises(RuntimeError) as caught:
ai_chat._chat_completion(
[{"role": "user", "content": "在吗"}], provider=self.provider
)
self.assertIn("网关", str(caught.exception))
def test_an_http_error_surfaces_the_gateway_reason(self) -> None:
import io as _io
import urllib.error
def boom(request, timeout=None):
raise urllib.error.HTTPError(
request.full_url, 503, "Service Unavailable", {},
_io.BytesIO(json.dumps({"detail": "没有可用的答题出口"}).encode()),
)
with mock.patch("urllib.request.urlopen", boom):
with self.assertRaises(RuntimeError) as caught:
ai_chat._chat_completion(
[{"role": "user", "content": "在吗"}], provider=self.provider
)
message = str(caught.exception)
self.assertIn("503", message)
self.assertIn("没有可用的答题出口", message, "后端说的原因要带给人看")
def test_an_empty_reply_reports_the_upstream_error(self) -> None:
"""网关返回空回复时,把候选里的错误原因翻出来——"没有回复"没法排查。"""
with self._urlopen([
{"reply": "", "candidates": [{"provider": "主答题", "error": "熔断中"}]}
]):
with self.assertRaises(RuntimeError) as caught:
ai_chat._chat_completion(
[{"role": "user", "content": "在吗"}], provider=self.provider
)
self.assertIn("熔断中", str(caught.exception))
def test_a_gateway_provider_is_never_mistaken_for_an_openai_endpoint(self) -> None:
"""`detect_kind` 认不出 gateway,会按地址猜成 openai。
猜错的后果是拿网关地址去拼 /chat/completions,请求发到一个不存在的路径,
报出来的是 404——排查的人会去查模型服务商,而问题在本地。
"""
self.assertEqual(ai_chat._provider_type(self.provider), ai_chat.GATEWAY_KIND)
self.assertFalse(ai_chat._is_dify_endpoint(self.provider))
def test_vision_goes_through_the_gateway_too(self) -> None:
"""图片识别同样没有第二条路——密钥在后端。"""
with self._urlopen([{"reply": "客户发来一张检查单"}]):
with mock.patch.object(ai_chat, "_finalize_vision_reply", lambda text, *a: text):
answer = ai_chat.call_ai_vision(
b"\x89PNG-fake", chat_text="看看这个", provider=self.provider
)
self.assertEqual(answer, "客户发来一张检查单")
content = self.sent[0]["body"]["messages"][-1]["content"]
self.assertTrue(
any(part.get("type") == "image_url" for part in content),
"图片要在 messages 里带过去,由网关按各家协议自己转换",
)
class GatewayPreferenceTest(TestCase):
"""网关配了就用网关——这是「角色编排」生效的前提。"""
def _settings(self, gateway):
module = mock.MagicMock()
module.load_settings.return_value = {"gateway": gateway}
module.DESKTOP_SYNC_KEY = "fallback-key"
return mock.patch.dict("sys.modules", {"backend_client": module})
def test_a_configured_gateway_wins_over_the_local_single_model(self) -> None:
with self._settings({"enabled": True, "url": "http://gw/v1/answer"}):
provider = ai_chat.current_provider()
self.assertEqual(provider.kind, ai_chat.GATEWAY_KIND)
self.assertEqual(provider.base_url, "http://gw/v1/answer")
def test_the_sync_key_falls_back_to_the_client_constant(self) -> None:
with self._settings({"enabled": True, "url": "http://gw/v1/answer"}):
provider = ai_chat.current_provider()
self.assertEqual(provider.api_key, "fallback-key")
def test_a_disabled_gateway_is_ignored(self) -> None:
with self._settings({"enabled": False, "url": "http://gw/v1/answer"}):
self.assertIsNone(ai_chat.gateway_provider())
def test_a_gateway_without_a_url_is_ignored(self) -> None:
"""开关开着但地址没填——这种半配好的状态不能当成"已启用"。"""
with self._settings({"enabled": True, "url": " "}):
self.assertIsNone(ai_chat.gateway_provider())
def test_falling_back_to_the_local_model_says_so_out_loud(self) -> None:
"""静悄悄地用一份和后台对不上的配置,比直接报错更难查。"""
ai_chat._LOCAL_PROVIDER_WARNED = False
with self._settings(None):
with mock.patch("builtins.print") as printed:
provider = ai_chat.current_provider()
self.assertNotEqual(provider.kind, ai_chat.GATEWAY_KIND)
said = " ".join(str(call) for call in printed.call_args_list)
self.assertIn("角色编排", said)
def setUpModule():
# 别让测试读到开发机上的真实桌面端配置——配过模型网关的机器会让
# `ai_chat.current_provider()` 整体改走网关分支,一大片无关测试跟着变行为。
local_state_redirect.start()
def tearDownModule():
local_state_redirect.stop()
if __name__ == "__main__":
main()