from __future__ import annotations import unittest from unittest.mock import patch from backend.llm import LLMGateway, LLMGatewayError from backend.llm.service import LLMServiceMixin 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,)) def test_saved_system_model_can_reach_the_connection_probe(self) -> None: service = LLMServiceMixin() service._system_credentials = { "llm_models": [ { "id": "primary-model", "name": "主模型", "api_key": "secret", "base_url": "https://model.example/v1", "model": "model-name", } ] } service.llm_gateway = self.gateway() expected = {"ok": True, "reply": "OK"} with patch( "backend.llm.service.test_llm_connection", return_value=expected ) as connection_probe: result = service.test_system_llm_profile("primary-model", {}) self.assertEqual(result, expected) connection_probe.assert_called_once_with( "secret", "https://model.example/v1", "model-name" ) if __name__ == "__main__": unittest.main()