feat: stream provider responses
This commit is contained in:
@@ -23,6 +23,7 @@ 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
|
||||
@@ -30,6 +31,14 @@ class _FakeModels:
|
||||
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:
|
||||
@@ -56,6 +65,23 @@ def _response(text: str = "Cloud-Antwort") -> SimpleNamespace:
|
||||
)
|
||||
|
||||
|
||||
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
|
||||
@@ -99,6 +125,39 @@ class GeminiProviderTests(unittest.TestCase):
|
||||
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())
|
||||
|
||||
Reference in New Issue
Block a user