Files
Jarvis-Ai/tests/unit/test_gemini_provider.py

223 lines
7.6 KiB
Python

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
self.stream_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
def generate_content_stream(self, **kwargs: object):
self.stream_call = kwargs
if isinstance(self.result, BaseException):
raise self.result
if isinstance(self.result, list):
return iter(self.result)
return iter((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,
),
)
def _stream_response(
visible_text: str,
*,
thought_text: str | None = None,
with_usage: bool = False,
) -> SimpleNamespace:
parts = []
if thought_text:
parts.append(SimpleNamespace(text=thought_text, thought=True))
parts.append(SimpleNamespace(text=visible_text, thought=False))
usage = SimpleNamespace(prompt_token_count=12, candidates_token_count=7) if with_usage else None
return SimpleNamespace(
candidates=[SimpleNamespace(content=SimpleNamespace(parts=parts))],
usage_metadata=usage,
)
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)
self.assertFalse(call["config"].thinking_config.include_thoughts)
def test_streaming_yields_only_visible_text_and_tracks_usage(self) -> None:
provider, factory = self._provider(
[
_stream_response("Cloud ", thought_text="verstecktes Denken"),
_stream_response("Antwort", with_usage=True),
]
)
chunks = list(provider.stream_chat([ChatMessage("user", "Was ist SQLite?")]))
self.assertEqual(chunks, ["Cloud ", "Antwort"])
self.assertNotIn("verstecktes Denken", "".join(chunks))
self.assertEqual(provider.last_usage.input_tokens, 12)
self.assertEqual(provider.last_usage.output_tokens, 7)
self.assertEqual(
factory.client.models.stream_call["model"],
"gemini-3.6-flash",
)
def test_streaming_error_is_sanitized(self) -> None:
provider, _ = self._provider(
errors.ClientError(
429,
{"message": "quota test-key-not-a-real-secret"},
)
)
with self.assertRaises(ProviderRateLimitError) as raised:
list(provider.stream_chat([ChatMessage("user", "Was ist Python?")]))
self.assertNotIn("test-key-not-a-real-secret", str(raised.exception))
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()