289 lines
10 KiB
Python
289 lines
10 KiB
Python
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_medical_never_routes_locally_before_cloud_construction(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", "Wenn ich blute, sollte ich dann zum Arzt?")])
|
|
|
|
self.assertEqual(answer, "Lokal")
|
|
self.assertFalse(constructed)
|
|
self.assertFalse(self.approvals)
|
|
self.assertEqual(router.last_route.privacy_policy, CloudPolicy.NEVER)
|
|
|
|
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)
|
|
self.assertFalse(router.last_route.fallback)
|
|
self.assertTrue(router.last_route.cloud_suppressed_by_local_mode)
|
|
self.assertFalse(router.last_route.technical_fallback)
|
|
|
|
def test_network_failure_is_reported_as_technical_fallback(self) -> None:
|
|
self.cloud.error = CloudNetworkError("offline")
|
|
router = self._router()
|
|
|
|
router.chat([ChatMessage("user", "Wie funktioniert Python?")])
|
|
|
|
self.assertFalse(router.last_route.cloud_suppressed_by_local_mode)
|
|
self.assertTrue(router.last_route.technical_fallback)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|