156 lines
5.1 KiB
Python
156 lines
5.1 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,
|
|
)
|
|
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_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()
|