feat: stream provider responses
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user