This commit is contained in:
Your Name
2026-08-27 14:04:28 +08:00
parent f7720831be
commit 334890171e
3016 changed files with 263403 additions and 27971 deletions
+584
View File
@@ -0,0 +1,584 @@
# -*- coding: utf-8 -*-
"""模型网关:并发、背压、熔断、幂等、兜底。
这些不是"锦上添花"的特性,是网关能不能上线的底线。每一条对应一种线上事故:
并发不真 → 两个 7 秒的调用变成 14 秒,客户等不了
没有背压 → 请求堆到超时,用户等 60 秒拿到失败
没有熔断 → 每个请求都去撞一个已经挂了的上游
没有幂等 → 桌面端超时重发,双倍扣费
兜底不全 → 上游抖一下,客户就收不到回复
"""
import asyncio
import json
import sqlite3
import tempfile
import time
from pathlib import Path
from unittest import TestCase, main, mock
import httpx
import model_gateway as gw
import model_protocol as mp
def _outlet(name="A", kind="openai", inflight=32, **cfg):
config = {
"id": name.lower(), "name": name, "kind": kind,
"base_url": "https://up/v1", "api_key": "k", "model": "m",
"max_tokens": 500, "temperature": 0.35, "timeout_ms": 30000,
"max_inflight": inflight,
}
config.update(cfg)
return gw.Outlet(config=config, gate=asyncio.Semaphore(inflight))
def _openai_body(text):
return {"choices": [{"message": {"content": text}}]}
def _client(handler):
return httpx.AsyncClient(transport=httpx.MockTransport(handler))
class ProtocolShapeTest(TestCase):
"""协议层是同步和异步两个传输方共用的,形状错了两边一起错。"""
def test_each_kind_lands_on_its_own_endpoint(self):
self.assertEqual(
mp.endpoint_url("claude", "https://api.anthropic.com"),
"https://api.anthropic.com/v1/messages")
self.assertEqual(
mp.endpoint_url("dify", "http://ai/v1"), "http://ai/v1/chat-messages")
self.assertEqual(
mp.endpoint_url("openai", "https://api.x.com"),
"https://api.x.com/chat/completions")
def test_a_fully_written_url_is_not_appended_to(self):
for kind, url in (
("claude", "https://api.anthropic.com/v1/messages"),
("dify", "http://ai/v1/chat-messages"),
("openai", "https://api.x.com/chat/completions"),
):
self.assertEqual(mp.endpoint_url(kind, url), url)
def test_claude_payload_lifts_system_and_forces_max_tokens(self):
payload = mp.chat_payload(
"claude", model="c", max_tokens=0, temperature=0.3,
messages=[{"role": "system", "content": "家规"},
{"role": "user", "content": "在吗"}])
self.assertEqual(payload["system"], "家规")
self.assertGreaterEqual(payload["max_tokens"], 1)
self.assertNotIn("system", [m["role"] for m in payload["messages"]])
def test_dify_payload_carries_query_and_user_not_messages(self):
payload = mp.chat_payload(
"dify", model="", max_tokens=500, temperature=0.3,
messages=[{"role": "system", "content": "忽略"},
{"role": "user", "content": "一天吃几次"}])
self.assertEqual(payload["query"], "一天吃几次")
self.assertIn("user", payload)
self.assertNotIn("messages", payload)
def test_parsers_raise_instead_of_returning_an_empty_string(self):
for kind, body in (("openai", {}), ("claude", {"content": []}), ("dify", {})):
with self.assertRaises(ValueError):
mp.parse_chat(kind, body)
def test_comfyui_is_not_a_chat_kind(self):
self.assertIn("comfyui", mp.ALL_KINDS)
self.assertNotIn("comfyui", mp.CHAT_KINDS)
class OutletCallTest(TestCase):
def test_a_normal_call_returns_text_and_latency(self):
async def run():
def handler(request):
return httpx.Response(200, json=_openai_body("好的"))
async with _client(handler) as client:
return await gw.call_outlet(client, _outlet(), [{"role": "user", "content": ""}])
out = asyncio.run(run())
self.assertEqual(out["text"], "好的")
self.assertEqual(out["error"], "")
self.assertGreaterEqual(out["latency_ms"], 0)
def test_a_retriable_status_is_retried_once(self):
calls = []
async def run():
def handler(request):
calls.append(1)
if len(calls) == 1:
return httpx.Response(503, text="busy")
return httpx.Response(200, json=_openai_body("好的"))
async with _client(handler) as client:
return await gw.call_outlet(client, _outlet(), [{"role": "user", "content": ""}])
out = asyncio.run(run())
self.assertEqual(out["text"], "好的")
self.assertEqual(len(calls), 2)
def test_a_4xx_is_never_retried(self):
calls = []
async def run():
def handler(request):
calls.append(1)
return httpx.Response(401, text="bad key")
async with _client(handler) as client:
return await gw.call_outlet(client, _outlet(), [{"role": "user", "content": ""}])
out = asyncio.run(run())
self.assertEqual(len(calls), 1, "确定性故障重试只会错三遍还多花钱")
self.assertIn("401", out["error"])
def test_a_failing_outlet_eventually_trips_the_breaker(self):
async def run():
def handler(request):
return httpx.Response(401, text="bad key")
outlet = _outlet()
async with _client(handler) as client:
for _ in range(gw.BREAKER_THRESHOLD):
await gw.call_outlet(client, outlet, [{"role": "user", "content": ""}])
tripped = await gw.call_outlet(client, outlet, [{"role": "user", "content": ""}])
return outlet, tripped
outlet, tripped = asyncio.run(run())
self.assertEqual(outlet.breaker.state, "open")
self.assertIn("熔断", tripped["error"])
def test_the_breaker_half_opens_after_cooling_down(self):
outlet = _outlet()
for _ in range(gw.BREAKER_THRESHOLD):
outlet.breaker.record(False)
self.assertFalse(outlet.breaker.allow())
outlet.breaker.opened_at -= gw.BREAKER_COOLDOWN + 1
self.assertTrue(outlet.breaker.allow())
self.assertEqual(outlet.breaker.state, "half_open")
def test_the_inflight_gate_caps_concurrency_per_outlet(self):
peak = []
live = []
async def run():
async def handler(request):
live.append(1)
peak.append(len(live))
await asyncio.sleep(0.05)
live.pop()
return httpx.Response(200, json=_openai_body(""))
outlet = _outlet(inflight=2)
async with _client(handler) as client:
await asyncio.gather(*[
gw.call_outlet(client, outlet, [{"role": "user", "content": ""}])
for _ in range(6)
])
asyncio.run(run())
self.assertLessEqual(max(peak), 2, "一家变慢不能把所有 worker 吃光")
def test_a_hung_upstream_is_cut_at_the_deadline(self):
async def run():
async def handler(request):
await asyncio.sleep(5)
return httpx.Response(200, json=_openai_body("太慢了"))
async with _client(handler) as client:
return await gw.call_outlet(
client, _outlet(), [{"role": "user", "content": ""}], deadline=0.3)
started = time.monotonic()
out = asyncio.run(run())
self.assertLess(time.monotonic() - started, 3.0)
self.assertTrue(out["error"])
self.assertEqual(out["text"], "")
class JudgeTest(TestCase):
def _judge(self, body, second="乙说"):
async def run():
def handler(request):
return httpx.Response(200, json=_openai_body(body))
async with _client(handler) as client:
return await gw.run_judge(client, _outlet("裁判"), "一天吃几次", "甲说", second)
return asyncio.run(run())
def test_reads_the_verdict(self):
v = self._judge('{"winner":"B","score":0.88,"risk":"low","reason":"更短"}')
self.assertEqual(v["winner"], "B")
self.assertAlmostEqual(v["score"], 0.88)
self.assertTrue(v["participated"])
def test_unparseable_output_does_not_participate(self):
self.assertFalse(self._judge("我觉得 B 好")["participated"])
def test_a_single_candidate_always_wins(self):
v = self._judge('{"winner":"B","score":0.7,"risk":"low"}', second="")
self.assertEqual(v["winner"], "A")
def test_no_judge_configured_is_a_cheap_no_op(self):
async def run():
async with _client(lambda r: httpx.Response(500)) as client:
return await gw.run_judge(client, None, "x", "甲说")
self.assertFalse(asyncio.run(run())["participated"])
class GatewayEndpointTest(TestCase):
"""端到端:起真的 app,用 MockTransport 假冒上游。"""
def setUp(self):
import admin_backend
self.root = Path(tempfile.mkdtemp())
self.db_path = self.root / "t.db"
db = admin_backend.Database(self.db_path)
db.initialize("InitialAdmin123")
uid = db.authenticate("admin", "InitialAdmin123")["id"]
for pid, name in (("a", "模型A"), ("b", "模型B"), ("j", "裁判")):
db.save_model_provider(
{"id": pid, "name": name, "kind": "openai",
"base_url": "https://up/v1", "api_key": "sk-x", "model": "m"},
uid, "1.1.1.1")
db.save_model_roles(
{"answer_ids": "a,b", "judge_id": "j", "judge_mode": "arbitrate"},
uid, "1.1.1.1")
self.app = gw.create_app(self.db_path, "sync-key")
def _post(self, handler, *, headers=None, body=None):
async def run():
from httpx import ASGITransport
async with httpx.AsyncClient(
transport=ASGITransport(app=self.app), base_url="http://gw"
) as caller:
# app 启动时建的真实 client 换成假冒上游
await self.app.gateway_startup()
await self.app.state.client.aclose()
self.app.state.client = _client(handler)
try:
started = time.monotonic()
response = await caller.post(
"/v1/answer",
headers={"X-Desktop-Sync-Key": "sync-key", **(headers or {})},
json=body or {"customer_text": "一天吃几次", "task_id": "t1"},
)
# 只计请求本身。gateway_startup/shutdown 每次要 0.3~0.4 秒
# (建目录缓存、开关 httpx 连接池),把它们算进去的话,任何
# 耗时断言测的其实是"今天这台机器忙不忙"。
self.request_seconds = time.monotonic() - started
return response
finally:
await self.app.gateway_shutdown()
await self.app.state.client.aclose()
return asyncio.run(run())
def _upstream(self, judge_json='{"winner":"B","score":0.9,"risk":"low"}'):
def handler(request):
body = json.loads(request.content.decode())
text = str(body.get("messages", [{}])[-1].get("content") or "")
if "评审" in text:
return httpx.Response(200, json=_openai_body(judge_json))
return httpx.Response(200, json=_openai_body("上游回复"))
return handler
def test_a_missing_sync_key_is_rejected(self):
response = self._post(self._upstream(), headers={"X-Desktop-Sync-Key": "wrong"})
self.assertEqual(response.status_code, 401)
def test_a_normal_request_returns_reply_candidates_and_verdict(self):
response = self._post(self._upstream())
self.assertEqual(response.status_code, 200)
data = response.json()
self.assertEqual(data["reply"], "上游回复")
self.assertEqual(len(data["candidates"]), 2)
self.assertTrue(data["judge"]["participated"])
self.assertEqual(data["judge_mode"], "arbitrate")
self.assertGreater(data["roles_version"], 0)
def _post_then_drain(self, handler, *, body):
"""发一次请求,等异步落库真的写完,再把那一行读回来。
专门补这条是因为用户问"为什么调用记录里客户消息和模型回复全是空"
在这之前,测试只直接调 `_write_log`,从没验证过 `/v1/answer` 这条真实
路径到底往队列里放了什么——写日志的那段代码等于没有测试覆盖。
"""
async def run():
from httpx import ASGITransport
async with httpx.AsyncClient(
transport=ASGITransport(app=self.app), base_url="http://gw"
) as caller:
await self.app.gateway_startup()
await self.app.state.client.aclose()
self.app.state.client = _client(handler)
try:
resp = await caller.post(
"/v1/answer",
headers={"X-Desktop-Sync-Key": "sync-key", "X-Device-Id": "desk-1"},
json=body,
)
await self.app.gateway_log_queue.join()
return resp
finally:
await self.app.gateway_shutdown()
await self.app.state.client.aclose()
response = asyncio.run(run())
con = sqlite3.connect(self.db_path)
con.row_factory = sqlite3.Row
try:
row = con.execute(
"SELECT * FROM model_calls WHERE task_id = ?",
(str(body.get("task_id") or ""),),
).fetchone()
finally:
con.close()
return response, row
def test_a_real_request_stores_what_the_customer_said_and_what_the_model_answered(self):
"""这就是"记录模型返回的话"这件事的端到端证明。
接口返回 200、界面上却是两个「(空)」,问题就出在这一段:请求路径拿到的
文本有没有真的进到那一行里。断言直接对着数据库读,不看接口返回。
"""
response, row = self._post_then_drain(
self._upstream(),
body={"customer_text": "我这几天总是失眠", "task_id": "e2e-1"},
)
self.assertEqual(response.status_code, 200)
self.assertIsNotNone(row, "这一次调用必须留下一行记录")
self.assertEqual(row["customer_text"], "我这几天总是失眠")
self.assertEqual(row["reply_text"], "上游回复")
self.assertEqual(row["purpose"], "chat", "没标 purpose 的默认是客服对话")
self.assertEqual(row["device_id"], "desk-1")
def test_a_guard_request_is_stored_as_an_internal_call(self):
_resp, row = self._post_then_drain(
self._upstream(),
body={"customer_text": "判断界面", "task_id": "e2e-2", "purpose": "guard"},
)
self.assertEqual(row["purpose"], "guard", "界面识别不能混进客服对话日志")
def test_all_upstreams_down_returns_502_not_a_fake_reply(self):
response = self._post(lambda r: httpx.Response(500, text="down"))
self.assertEqual(response.status_code, 502)
self.assertEqual(response.json()["reply"], "")
def test_the_same_idempotency_key_is_not_charged_twice(self):
hits = []
def handler(request):
hits.append(1)
body = json.loads(request.content.decode())
text = str(body.get("messages", [{}])[-1].get("content") or "")
if "评审" in text:
return httpx.Response(200, json=_openai_body('{"winner":"A","score":0.8,"risk":"low"}'))
return httpx.Response(200, json=_openai_body("上游回复"))
async def run():
from httpx import ASGITransport
async with httpx.AsyncClient(
transport=ASGITransport(app=self.app), base_url="http://gw"
) as caller:
await self.app.gateway_startup()
await self.app.state.client.aclose()
self.app.state.client = _client(handler)
try:
headers = {"X-Desktop-Sync-Key": "sync-key", "X-Idempotency-Key": "k1"}
first = await caller.post("/v1/answer", headers=headers,
json={"customer_text": "x", "task_id": "t"})
second = await caller.post("/v1/answer", headers=headers,
json={"customer_text": "x", "task_id": "t"})
return first.json(), second.json(), len(hits)
finally:
await self.app.gateway_shutdown()
await self.app.state.client.aclose()
first, second, upstream_calls = asyncio.run(run())
self.assertEqual(first, second)
self.assertLessEqual(upstream_calls, 3, "重发不该再打一遍上游")
def test_health_reports_breakers_and_the_active_plan(self):
async def run():
from httpx import ASGITransport
async with httpx.AsyncClient(
transport=ASGITransport(app=self.app), base_url="http://gw"
) as caller:
await self.app.gateway_startup()
try:
return (await caller.get("/health")).json()
finally:
await self.app.gateway_shutdown()
await self.app.state.client.aclose()
data = asyncio.run(run())
self.assertEqual(data["status"], "ok")
self.assertEqual(sorted(data["roles"]["answer"]), ["模型A", "模型B"])
self.assertEqual(data["roles"]["judge"], "裁判")
self.assertTrue(all(o["breaker"] == "closed" for o in data["outlets"]))
def test_candidates_are_asked_in_parallel(self):
# 上游延迟。断言用它算出串行基线,不写死秒数——改这一个常量,
# 下面的判据自动跟着走。
delay = 0.08
live, peak = [], []
async def handler(request):
body = json.loads(request.content.decode())
text = str(body.get("messages", [{}])[-1].get("content") or "")
if "评审" in text:
return httpx.Response(200, json=_openai_body('{"winner":"A","score":0.8,"risk":"low"}'))
live.append(1)
peak.append(len(live))
await asyncio.sleep(delay)
live.pop()
return httpx.Response(200, json=_openai_body("上游回复"))
response = self._post(handler)
self.assertEqual(response.status_code, 200)
# 这一条是并发的确证:同一时刻有两路在飞。它是计数,不是计时,
# 机器再忙也不会动摇。
self.assertEqual(max(peak), 2, "两路必须同时在飞")
# 再补一条耗时判据,挡住"确实同时发出去了、但下游还是被某处串起来等"
# 这种情况。基线由 delay 现算:串行至少 2×delay,并发应该接近 1×。
self.assertLess(
self.request_seconds,
delay * 2,
f"两路答题若是串行,光这一段就要 {delay * 2:.2f}s"
f"实测 {self.request_seconds:.3f}s 说明没并起来",
)
class CatalogTest(TestCase):
def setUp(self):
import admin_backend
self.root = Path(tempfile.mkdtemp())
self.db_path = self.root / "t.db"
self.db = admin_backend.Database(self.db_path)
self.db.initialize("InitialAdmin123")
self.uid = self.db.authenticate("admin", "InitialAdmin123")["id"]
def test_keys_are_decrypted_for_the_gateway_only(self):
self.db.save_model_provider(
{"id": "a", "name": "A", "kind": "openai", "base_url": "https://u/v1",
"api_key": "sk-REAL"}, self.uid, "1.1.1.1")
catalog = gw.Catalog(self.db_path)
catalog.refresh()
self.assertEqual(catalog.outlets["a"].config["api_key"], "sk-REAL")
def test_comfyui_never_enters_the_answer_pool(self):
self.db.save_model_provider(
{"id": "img", "name": "文生图", "kind": "comfyui",
"base_url": "http://127.0.0.1:8188"}, self.uid, "1.1.1.1")
self.db.save_model_provider(
{"id": "a", "name": "A", "kind": "openai", "base_url": "https://u/v1",
"api_key": "k"}, self.uid, "1.1.1.1")
self.db.save_model_roles(
{"answer_ids": "a,img", "judge_id": "img"}, self.uid, "1.1.1.1")
catalog = gw.Catalog(self.db_path)
catalog.refresh()
answers, judge, _mode, _v = catalog.plan()
self.assertEqual([o.name for o in answers], ["A"])
self.assertIsNone(judge)
def test_a_disabled_outlet_disappears_from_the_plan(self):
self.db.save_model_provider(
{"id": "a", "name": "A", "kind": "openai", "base_url": "https://u/v1",
"api_key": "k", "enabled": False}, self.uid, "1.1.1.1")
catalog = gw.Catalog(self.db_path)
catalog.refresh()
self.assertNotIn("a", catalog.outlets)
def test_a_broken_database_keeps_the_previous_catalog(self):
self.db.save_model_provider(
{"id": "a", "name": "A", "kind": "openai", "base_url": "https://u/v1",
"api_key": "k"}, self.uid, "1.1.1.1")
catalog = gw.Catalog(self.db_path)
catalog.refresh()
catalog.db_path = self.root / "没有这个文件.db"
catalog.refresh()
self.assertIn("a", catalog.outlets, "读库失败要降级,不能失效")
self.assertTrue(catalog.last_error)
def test_the_call_log_lands_in_the_shared_table(self):
gw._write_log(self.db_path, {
"device_id": "d1", "task_id": "t1", "roles_version": 1,
"judge_mode": "arbitrate", "chosen": "模型B",
"judge": {"winner": "B", "score": 0.9, "risk": "low"},
"candidates": [{"provider": "A"}], "total_ms": 5200})
con = sqlite3.connect(self.db_path)
try:
row = con.execute(
"SELECT device_id,chosen,judge_winner,judge_score FROM model_calls"
).fetchone()
finally:
con.close()
self.assertEqual(row, ("d1", "模型B", "B", 0.9))
def test_the_call_log_also_keeps_customer_and_reply_text(self):
"""跟踪问题要能看到"客户说了什么、模型答了什么",不能只有分数。"""
gw._write_log(self.db_path, {
"device_id": "d1", "task_id": "t1", "chosen": "模型B",
"judge": {}, "candidates": [], "total_ms": 100,
"customer_text": "我这几天总是失眠", "reply_text": "建议您注意休息",
})
con = sqlite3.connect(self.db_path)
try:
row = con.execute(
"SELECT customer_text,reply_text FROM model_calls"
).fetchone()
finally:
con.close()
self.assertEqual(row, ("我这几天总是失眠", "建议您注意休息"))
def test_a_second_write_with_the_same_task_id_refreshes_data_but_keeps_review_reason(self):
"""桌面端可能先落盘(它知道审核规则命中原因,网关不知道)——网关补写权威数据时不能把这个原因擦掉。"""
gw._write_log(self.db_path, {
"device_id": "d1", "task_id": "t1", "chosen": "旧值",
"judge": {}, "candidates": [], "total_ms": 1,
"customer_text": "旧客户消息", "reply_text": "旧回复",
})
con = sqlite3.connect(self.db_path)
try:
con.execute(
"UPDATE model_calls SET review_reason=? WHERE task_id='t1'",
("命中审核规则「诊断」",),
)
con.commit()
finally:
con.close()
gw._write_log(self.db_path, {
"device_id": "d1", "task_id": "t1", "chosen": "新值(服务端权威)",
"judge": {"winner": "A", "score": 0.5}, "candidates": [{"provider": "A"}],
"total_ms": 800, "customer_text": "新客户消息", "reply_text": "新回复",
})
con = sqlite3.connect(self.db_path)
try:
rows = con.execute("SELECT * FROM model_calls WHERE task_id='t1'").fetchall()
row = con.execute(
"SELECT chosen,reply_text,review_reason FROM model_calls WHERE task_id='t1'"
).fetchone()
finally:
con.close()
self.assertEqual(len(rows), 1, "同一次调用只能落一行")
self.assertEqual(row, ("新值(服务端权威)", "新回复", "命中审核规则「诊断」"))
if __name__ == "__main__":
main()