from __future__ import annotations import unittest from types import SimpleNamespace from typing import Any import httpx from google.genai import errors from javis.providers.base import ( ChatMessage, CloudNetworkError, InvalidApiKeyError, InvalidProviderResponseError, MissingApiKeyError, ProviderRateLimitError, ProviderTimeoutError, ) from javis.providers.gemini import GeminiProvider class _FakeModels: def __init__(self, result: object) -> None: self.result = result self.call: dict[str, Any] | None = None def generate_content(self, **kwargs: object) -> object: self.call = kwargs if isinstance(self.result, BaseException): raise self.result return self.result class _FakeClient: def __init__(self, result: object) -> None: self.models = _FakeModels(result) class _ClientFactory: def __init__(self, result: object) -> None: self.client = _FakeClient(result) self.arguments: dict[str, Any] | None = None def __call__(self, **kwargs: object) -> _FakeClient: self.arguments = kwargs return self.client def _response(text: str = "Cloud-Antwort") -> SimpleNamespace: return SimpleNamespace( text=text, usage_metadata=SimpleNamespace( prompt_token_count=12, candidates_token_count=7, ), ) class GeminiProviderTests(unittest.TestCase): def _provider( self, result: object, **overrides: object ) -> tuple[GeminiProvider, _ClientFactory]: factory = _ClientFactory(result) arguments: dict[str, object] = { "model": "gemini-3.6-flash", "api_key": "test-key-not-a-real-secret", "timeout_seconds": 30, "max_output_tokens": 1024, "max_retries": 1, "client_factory": factory, } arguments.update(overrides) return GeminiProvider(**arguments), factory def test_success_uses_bounded_sdk_configuration_and_tracks_usage(self) -> None: provider, factory = self._provider(_response()) answer = provider.chat( [ ChatMessage("user", "Hallo"), ChatMessage("assistant", "Guten Tag"), ChatMessage("user", "Was ist SQLite?"), ] ) self.assertEqual(answer, "Cloud-Antwort") self.assertEqual(provider.last_usage.input_tokens, 12) self.assertEqual(provider.last_usage.output_tokens, 7) self.assertNotIn("test-key-not-a-real-secret", repr(factory.client.models.call)) http_options = factory.arguments["http_options"] self.assertEqual(http_options.api_version, "v1") self.assertEqual(http_options.timeout, 30_000) self.assertEqual(http_options.retry_options.attempts, 2) self.assertNotIn(429, http_options.retry_options.http_status_codes) call = factory.client.models.call self.assertEqual(call["model"], "gemini-3.6-flash") self.assertEqual([content.role for content in call["contents"]], ["user", "model", "user"]) self.assertEqual(call["config"].max_output_tokens, 1024) self.assertIsNone(call["config"].tools) def test_missing_key_stops_before_client_creation(self) -> None: factory = _ClientFactory(_response()) with self.assertRaises(MissingApiKeyError): GeminiProvider( model="gemini-3.6-flash", api_key=None, timeout_seconds=30, max_output_tokens=1024, max_retries=1, client_factory=factory, ) self.assertIsNone(factory.arguments) def test_invalid_key_is_sanitized(self) -> None: provider, _ = self._provider( errors.ClientError( 401, {"message": "rejected test-key-not-a-real-secret"}, ) ) with self.assertRaisesRegex(InvalidApiKeyError, "wurde abgelehnt") as raised: provider.chat([ChatMessage("user", "Hallo")]) self.assertNotIn("test-key-not-a-real-secret", str(raised.exception)) def test_rate_limit_is_a_distinct_fallback_signal(self) -> None: provider, _ = self._provider(errors.ClientError(429, {"message": "quota"})) with self.assertRaises(ProviderRateLimitError): provider.chat([ChatMessage("user", "Was ist Python?")]) def test_network_failure_is_a_distinct_fallback_signal(self) -> None: request = httpx.Request("POST", "https://example.invalid") provider, _ = self._provider(httpx.ConnectError("offline", request=request)) with self.assertRaises(CloudNetworkError): provider.chat([ChatMessage("user", "Was ist Python?")]) def test_timeout_is_a_distinct_fallback_signal(self) -> None: request = httpx.Request("POST", "https://example.invalid") provider, _ = self._provider(httpx.ReadTimeout("timeout", request=request)) with self.assertRaises(ProviderTimeoutError): provider.chat([ChatMessage("user", "Was ist Python?")]) def test_empty_response_is_rejected(self) -> None: provider, _ = self._provider(_response(" ")) with self.assertRaises(InvalidProviderResponseError): provider.chat([ChatMessage("user", "Was ist Python?")]) def test_more_than_one_retry_is_rejected(self) -> None: with self.assertRaises(ValueError): self._provider(_response(), max_retries=2) if __name__ == "__main__": unittest.main()