feat: stream provider responses
This commit is contained in:
@@ -16,6 +16,7 @@ from javis.providers.base import (
|
||||
ProviderRateLimitError,
|
||||
ProviderTimeoutError,
|
||||
ProviderUsage,
|
||||
ResponseAbortedError,
|
||||
)
|
||||
from javis.security.privacy import CloudPolicy, PrivacyRouter
|
||||
|
||||
@@ -28,11 +29,15 @@ class _RecordingProvider:
|
||||
*,
|
||||
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)
|
||||
|
||||
@@ -42,6 +47,17 @@ class _RecordingProvider:
|
||||
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:
|
||||
@@ -283,6 +299,64 @@ class HybridProviderTests(unittest.TestCase):
|
||||
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()
|
||||
|
||||
Reference in New Issue
Block a user