from __future__ import annotations import tempfile import unittest from pathlib import Path from javis.core.provider_router import HybridProvider from javis.memory.usage_store import ProviderEvent, SQLiteUsageStore from javis.providers.base import ( ChatMessage, CloudModelUnavailableError, CloudNetworkError, InvalidApiKeyError, LocalModelProvider, MissingApiKeyError, ProviderRateLimitError, ProviderTimeoutError, ProviderUsage, ) from javis.security.privacy import CloudPolicy, PrivacyRouter class _RecordingProvider: def __init__( self, name: str, model: str, *, answer: str = "Antwort", error: Exception | None = None, ) -> None: self.name = name self.model = model self.answer = answer self.error = error self.calls: list[list[ChatMessage]] = [] self.last_usage = ProviderUsage(11, 5) def chat(self, messages: list[ChatMessage]) -> str: self.calls.append(messages) if self.error: raise self.error return self.answer class HybridProviderTests(unittest.TestCase): def setUp(self) -> None: self.temporary_directory = tempfile.TemporaryDirectory() self.usage_store = SQLiteUsageStore( Path(self.temporary_directory.name) / "provider-usage.sqlite3" ) self.local = _RecordingProvider("ollama", "qwen3:8b", answer="Lokal") self.cloud = _RecordingProvider( "gemini", "gemini-3.6-flash", answer="Cloud", ) self.notices: list[str] = [] self.approvals: list[object] = [] self.approval_result = False def tearDown(self) -> None: self.temporary_directory.cleanup() def _router(self, **overrides: object) -> HybridProvider: arguments: dict[str, object] = { "local_provider": self.local, "cloud_provider_factory": lambda: self.cloud, "cloud_model": "gemini-3.6-flash", "privacy_router": PrivacyRouter(), "usage_store": self.usage_store, "approval_callback": self._approve, "notice_callback": self.notices.append, "mode": "auto", "cloud_enabled": True, "billing_confirmed_disabled": True, "free_only": True, "max_cloud_requests_per_day": 25, "max_cloud_input_chars": 12_000, "max_cloud_context_messages": 6, } arguments.update(overrides) return HybridProvider(**arguments) def _approve(self, decision: object) -> bool: self.approvals.append(decision) return self.approval_result def test_contract_and_existing_local_session_are_supported(self) -> None: router = self._router() self.assertIsInstance(router, LocalModelProvider) self.assertTrue(router.supports_session("hybrid", "auto")) self.assertTrue(router.supports_session("ollama", "qwen3:8b")) self.assertFalse(router.supports_session("ollama", "other-model")) def test_allowed_technical_content_uses_cloud_with_minimal_context(self) -> None: router = self._router() answer = router.chat( [ ChatMessage("user", "Meine IBAN ist vertraulich"), ChatMessage("assistant", "Verstanden"), ChatMessage("user", "Wie funktioniert SQLite?"), ] ) self.assertEqual(answer, "Cloud") self.assertEqual(len(self.cloud.calls), 1) self.assertEqual( [message.content for message in self.cloud.calls[0]], ["Wie funktioniert SQLite?"], ) self.assertEqual(router.last_route.provider, "gemini") self.assertEqual(router.last_route.privacy_policy, CloudPolicy.ALLOWED) self.assertFalse(router.last_route.fallback) def test_never_policy_cannot_construct_cloud_provider(self) -> None: constructed = False def cloud_factory() -> _RecordingProvider: nonlocal constructed constructed = True return self.cloud router = self._router(cloud_provider_factory=cloud_factory, mode="gemini") answer = router.chat([ChatMessage("user", "Mein API-Key ist geheim")]) self.assertEqual(answer, "Lokal") self.assertFalse(constructed) self.assertTrue(router.last_route.fallback) self.assertIn("lokale Antwort", self.notices[-1]) def test_ask_defaults_to_local_when_approval_is_denied(self) -> None: router = self._router() answer = router.chat([ChatMessage("user", "Meine Familie plant Urlaub")]) self.assertEqual(answer, "Lokal") self.assertEqual(len(self.approvals), 1) self.assertFalse(self.cloud.calls) def test_ask_uses_cloud_after_explicit_approval(self) -> None: self.approval_result = True router = self._router() answer = router.chat([ChatMessage("user", "Meine Familie plant Urlaub")]) self.assertEqual(answer, "Cloud") self.assertEqual(len(self.cloud.calls), 1) def test_429_falls_back_without_losing_local_answer(self) -> None: self.cloud.error = ProviderRateLimitError("quota") router = self._router() answer = router.chat([ChatMessage("user", "Wie funktioniert Python?")]) self.assertEqual(answer, "Lokal") self.assertEqual(self.usage_store.cloud_requests_on(), 1) self.assertEqual(self.usage_store.local_fallbacks_on(), 1) self.assertIn("Gemini nicht verfügbar", self.notices[-1]) def test_network_failure_falls_back(self) -> None: self.cloud.error = CloudNetworkError("offline") router = self._router() answer = router.chat([ChatMessage("user", "Wie funktioniert Python?")]) self.assertEqual(answer, "Lokal") self.assertTrue(router.last_route.fallback) def test_timeout_and_missing_model_fall_back(self) -> None: for error in ( ProviderTimeoutError("timeout"), CloudModelUnavailableError("missing model"), ): with self.subTest(error=type(error).__name__): self.cloud.error = error router = self._router() answer = router.chat([ChatMessage("user", "Wie funktioniert Python?")]) self.assertEqual(answer, "Lokal") self.assertTrue(router.last_route.fallback) def test_missing_key_falls_back(self) -> None: def missing_key() -> _RecordingProvider: raise MissingApiKeyError("missing") router = self._router(cloud_provider_factory=missing_key) answer = router.chat([ChatMessage("user", "Wie funktioniert Python?")]) self.assertEqual(answer, "Lokal") self.assertTrue(router.last_route.fallback) def test_invalid_key_disables_cloud_for_remaining_process(self) -> None: self.cloud.error = InvalidApiKeyError("invalid") router = self._router() router.chat([ChatMessage("user", "Wie funktioniert Python?")]) self.cloud.error = None router.chat([ChatMessage("user", "Wie funktioniert SQLite?")]) self.assertEqual(len(self.cloud.calls), 1) self.assertEqual(len(self.local.calls), 2) self.assertIn("deaktiviert", self.notices[-1]) def test_local_daily_limit_stops_cloud_before_construction(self) -> None: self.usage_store.record( ProviderEvent( "gemini", "gemini-3.6-flash", True, None, False, "allowed", ) ) router = self._router(max_cloud_requests_per_day=1) answer = router.chat([ChatMessage("user", "Wie funktioniert Python?")]) self.assertEqual(answer, "Lokal") self.assertFalse(self.cloud.calls) self.assertIn("tägliches Cloudlimit", self.notices[-1]) def test_cloud_requires_local_activation_and_no_billing_confirmation(self) -> None: for override in ( {"cloud_enabled": False}, {"billing_confirmed_disabled": False}, ): with self.subTest(override=override): router = self._router(**override) answer = router.chat([ChatMessage("user", "Wie funktioniert Python?")]) self.assertEqual(answer, "Lokal") self.assertFalse(self.cloud.calls) def test_oversized_current_message_remains_local(self) -> None: router = self._router(max_cloud_input_chars=10) answer = router.chat([ChatMessage("user", "Wie funktioniert Python?")]) self.assertEqual(answer, "Lokal") self.assertFalse(self.cloud.calls) def test_local_mode_never_asks_or_calls_cloud(self) -> None: router = self._router(mode="local") answer = router.chat([ChatMessage("user", "Meine Familie plant Urlaub")]) self.assertEqual(answer, "Lokal") self.assertFalse(self.approvals) self.assertFalse(self.cloud.calls) if __name__ == "__main__": unittest.main()