# -*- 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()