Files
xiaobaifupan/tests/test_llm_gateway.py
T

119 lines
3.9 KiB
Python

from __future__ import annotations
import unittest
from backend.llm import LLMGateway, LLMGatewayError
class ProviderFailure(RuntimeError):
pass
class FakeDatabase:
def __init__(self, used: int = 0) -> None:
self.used = used
self.audit: list[dict[str, object]] = []
def count_llm_usage_since(self, user_id: int, source: str, since: str) -> int:
return self.used
def record_llm_usage(
self,
user_id: int,
feature: str,
source: str,
model: str,
status: str,
latency_ms: int,
**metadata,
) -> None:
self.audit.append(
{
"user_id": user_id,
"feature": feature,
"source": source,
"model": model,
"status": status,
**metadata,
}
)
def profile() -> dict[str, object]:
return {
"source": "platform",
"primary": {"api_key": "p", "base_url": "https://p", "model": "primary"},
"fallback": {"api_key": "f", "base_url": "https://f", "model": "fallback"},
}
class LLMGatewayTests(unittest.TestCase):
def gateway(self, database: FakeDatabase | None = None) -> LLMGateway:
database = database or FakeDatabase()
return LLMGateway(
database=database,
user_id_supplier=lambda: 7,
membership_supplier=lambda: {"active": True},
settings_supplier=lambda: {"member_daily_limit": 50},
profile_supplier=profile,
)
def test_non_streaming_call_falls_back_and_audits_once(self) -> None:
database = FakeDatabase()
gateway = self.gateway(database)
def invoke(model):
if model.role == "primary":
raise ProviderFailure("primary failed")
return {"answer": "ok"}
result = gateway.call("heaven_trend", "heaven-trend-v1", invoke, (ProviderFailure,))
self.assertEqual(result.role, "fallback")
self.assertEqual(result.value, {"answer": "ok"})
self.assertEqual(len(database.audit), 1)
self.assertEqual(database.audit[0]["model"], "fallback")
self.assertEqual(database.audit[0]["prompt_version"], "heaven-trend-v1")
def test_stream_falls_back_before_first_delta(self) -> None:
database = FakeDatabase()
gateway = self.gateway(database)
def invoke(model):
if model.role == "primary":
raise ProviderFailure("primary failed")
yield "a"
yield "b"
events = list(gateway.stream("mentor", "mentor-v1", invoke, (ProviderFailure,)))
self.assertEqual([event.value for event in events[:-1]], ["a", "b"])
self.assertEqual(events[-1].kind, "complete")
self.assertEqual(events[-1].role, "fallback")
self.assertEqual(database.audit[0]["status"], "success")
def test_stream_does_not_switch_model_after_output_started(self) -> None:
database = FakeDatabase()
gateway = self.gateway(database)
def invoke(model):
yield "first"
raise ProviderFailure(f"{model.role} interrupted")
iterator = gateway.stream("assistant", "assistant-v1", invoke, (ProviderFailure,))
self.assertEqual(next(iterator).value, "first")
with self.assertRaisesRegex(LLMGatewayError, "连接中断"):
list(iterator)
self.assertEqual(len(database.audit), 1)
self.assertEqual(database.audit[0]["model"], "primary")
self.assertEqual(database.audit[0]["status"], "failed")
def test_daily_quota_is_enforced_before_provider_call(self) -> None:
gateway = self.gateway(FakeDatabase(used=50))
with self.assertRaisesRegex(LLMGatewayError, "额度已用完"):
gateway.call("mentor", "mentor-v1", lambda model: "unused", (ProviderFailure,))
if __name__ == "__main__":
unittest.main()