119 lines
3.9 KiB
Python
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()
|