1081 lines
48 KiB
Python
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()
|