Files
Jarvis-Ai/tests/unit/test_provider_router.py
T

388 lines
14 KiB
Python

from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from javis.core.provider_router import ApprovalChoice, HybridProvider
from javis.memory.usage_store import ProviderEvent, SQLiteUsageStore
from javis.providers.base import (
ChatMessage,
CloudModelUnavailableError,
CloudNetworkError,
InvalidApiKeyError,
LocalModelProvider,
MissingApiKeyError,
ProviderRateLimitError,
ProviderTimeoutError,
ProviderUsage,
ResponseAbortedError,
)
from javis.security.privacy import CloudPolicy, PrivacyRouter
class _RecordingProvider:
def __init__(
self,
name: str,
model: str,
*,
answer: str = "Antwort",
error: Exception | None = None,
stream_chunks: tuple[str, ...] | None = None,
stream_error: Exception | None = None,
) -> None:
self.name = name
self.model = model
self.answer = answer
self.error = error
self.stream_chunks = stream_chunks
self.stream_error = stream_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
def stream_chat(self, messages: list[ChatMessage]):
self.calls.append(messages)
if self.stream_chunks is None:
if self.error:
raise self.error
yield self.answer
else:
yield from self.stream_chunks
if self.stream_error:
raise self.stream_error
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_queued_gui_approval_is_one_shot_and_uses_fake_cloud(self) -> None:
router = self._router()
question = [ChatMessage("user", "Meine Familie plant Urlaub")]
router.queue_approval(ApprovalChoice.ALLOW)
chunks = list(router.stream_chat(question))
second_answer = router.chat(question)
self.assertEqual(chunks, ["Cloud"])
self.assertEqual(second_answer, "Lokal")
self.assertEqual(len(self.approvals), 1)
def test_ask_supports_allow_local_and_cancel_choices(self) -> None:
question = [ChatMessage("user", "Meine Familie plant Urlaub")]
allowed = self._router(approval_callback=lambda _decision: ApprovalChoice.ALLOW)
self.assertEqual(allowed.chat(question), "Cloud")
local = self._router(approval_callback=lambda _decision: ApprovalChoice.LOCAL)
self.assertEqual(local.chat(question), "Lokal")
cancelled = self._router(approval_callback=lambda _decision: ApprovalChoice.CANCEL)
with self.assertRaises(ResponseAbortedError):
cancelled.chat(question)
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)
def test_allowed_cloud_response_streams_visible_chunks(self) -> None:
self.cloud.stream_chunks = ("Cloud ", "Stream")
router = self._router()
chunks = list(router.stream_chat([ChatMessage("user", "Wie funktioniert SQLite?")]))
self.assertEqual(chunks, ["Cloud ", "Stream"])
self.assertEqual(router.last_route.provider, "gemini")
self.assertFalse(router.last_route.fallback)
def test_cloud_stream_error_before_output_falls_back_locally(self) -> None:
self.cloud.stream_chunks = ()
self.cloud.stream_error = CloudNetworkError("offline")
self.local.stream_chunks = ("Lokaler ", "Ersatz")
router = self._router()
chunks = list(router.stream_chat([ChatMessage("user", "Wie funktioniert Python?")]))
self.assertEqual(chunks, ["Lokaler ", "Ersatz"])
self.assertTrue(router.last_route.fallback)
self.assertTrue(router.last_route.technical_fallback)
self.assertIn("lokale Antwort", self.notices[-1])
def test_cloud_stream_error_after_output_never_appends_local_answer(self) -> None:
self.cloud.stream_chunks = ("Teilantwort",)
self.cloud.stream_error = CloudNetworkError("offline")
self.local.stream_chunks = ("Lokaler Ersatz",)
router = self._router()
stream = router.stream_chat([ChatMessage("user", "Wie funktioniert Python?")])
self.assertEqual(next(stream), "Teilantwort")
with self.assertRaises(ResponseAbortedError):
next(stream)
self.assertFalse(self.local.calls)
self.assertFalse(router.last_route.fallback)
self.assertTrue(router.last_route.technical_fallback)
def test_medical_never_streams_locally_before_cloud_construction(self) -> None:
constructed = False
self.local.stream_chunks = ("Lokal",)
def cloud_factory() -> _RecordingProvider:
nonlocal constructed
constructed = True
return self.cloud
router = self._router(cloud_provider_factory=cloud_factory, mode="gemini")
chunks = list(
router.stream_chat([ChatMessage("user", "Welche Diagnose passt zu meinen Schmerzen?")])
)
self.assertEqual(chunks, ["Lokal"])
self.assertFalse(constructed)
self.assertFalse(self.approvals)
self.assertEqual(router.last_route.privacy_policy, CloudPolicy.NEVER)
if __name__ == "__main__":
unittest.main()