feat: stream provider responses

This commit is contained in:
2026-07-30 19:38:53 +02:00
parent 9148193806
commit 00a5a6ce8c
12 changed files with 850 additions and 41 deletions
+49 -1
View File
@@ -4,7 +4,7 @@ from pathlib import Path
from javis.core.chat_service import ChatService
from javis.memory.sqlite_store import SQLiteSessionStore
from javis.providers.base import ChatMessage, LocalModelProvider
from javis.providers.base import ChatMessage, LocalModelProvider, ProviderUnavailableError
class RecordingProvider:
@@ -19,6 +19,19 @@ class RecordingProvider:
return f"Antwort {len(self.calls)}"
class StreamingProvider(RecordingProvider):
def __init__(self, *, fail_after_first: bool = False) -> None:
super().__init__()
self.fail_after_first = fail_after_first
def stream_chat(self, messages: list[ChatMessage]):
self.calls.append(messages)
yield "Teil "
if self.fail_after_first:
raise ProviderUnavailableError("Streamfehler")
yield "Antwort"
class ChatServiceTests(unittest.TestCase):
def setUp(self) -> None:
self.temporary_directory = tempfile.TemporaryDirectory()
@@ -52,6 +65,41 @@ class ChatServiceTests(unittest.TestCase):
],
)
def test_streaming_response_is_persisted_exactly_once_after_completion(self) -> None:
provider = StreamingProvider()
service = ChatService(self.service.store, provider)
session = service.new_session()
chunks = list(service.stream_send(session.id, "Hallo"))
self.assertEqual(chunks, ["Teil ", "Antwort"])
self.assertEqual(
[message.content for message in service.load_session(session.id).messages],
["Hallo", "Teil Antwort"],
)
self.assertEqual(service.load_session(session.id).session.message_count, 2)
def test_closed_stream_does_not_persist_partial_response(self) -> None:
provider = StreamingProvider()
service = ChatService(self.service.store, provider)
session = service.new_session()
stream = service.stream_send(session.id, "Hallo")
self.assertEqual(next(stream), "Teil ")
stream.close()
self.assertEqual(service.load_session(session.id).messages, [])
def test_provider_error_during_stream_does_not_persist_partial_response(self) -> None:
provider = StreamingProvider(fail_after_first=True)
service = ChatService(self.service.store, provider)
session = service.new_session()
with self.assertRaises(ProviderUnavailableError):
list(service.stream_send(session.id, "Hallo"))
self.assertEqual(service.load_session(session.id).messages, [])
if __name__ == "__main__":
unittest.main()