feat: connect desktop sessions and streaming
This commit is contained in:
@@ -1,10 +1,16 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import tempfile
|
||||
import threading
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
from javis.core.chat_service import ChatService
|
||||
from javis.memory.sqlite_store import SQLiteSessionStore
|
||||
from javis.providers.base import ChatMessage
|
||||
from javis.ui.chat_controller import DesktopController, sanitized_error
|
||||
from javis.ui.models import GenerationState, StreamEventKind
|
||||
|
||||
|
||||
class _Service:
|
||||
@@ -44,6 +50,20 @@ class _Hybrid:
|
||||
self.mode = mode
|
||||
|
||||
|
||||
class _StreamingProvider:
|
||||
name = "ollama"
|
||||
model = "test"
|
||||
last_route = None
|
||||
mode = "local"
|
||||
|
||||
def stream_chat(self, _messages):
|
||||
yield "Teil "
|
||||
yield "Antwort"
|
||||
|
||||
def set_mode(self, mode: str) -> None:
|
||||
self.mode = mode
|
||||
|
||||
|
||||
class DesktopControllerTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.runtime = SimpleNamespace(
|
||||
@@ -71,3 +91,43 @@ class DesktopControllerTests(unittest.TestCase):
|
||||
|
||||
self.assertNotIn(secret, result)
|
||||
self.assertIn("MASKIERT", result)
|
||||
|
||||
def test_streaming_events_are_forwarded_and_completed(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
provider = _StreamingProvider()
|
||||
service = ChatService(
|
||||
SQLiteSessionStore(Path(directory) / "sessions.sqlite3"),
|
||||
provider,
|
||||
)
|
||||
controller = DesktopController(
|
||||
SimpleNamespace(service=service, hybrid_provider=provider)
|
||||
)
|
||||
|
||||
events = list(controller.stream_message("Hallo", threading.Event()))
|
||||
|
||||
self.assertEqual(
|
||||
[event.text for event in events if event.kind is StreamEventKind.CHUNK],
|
||||
["Teil ", "Antwort"],
|
||||
)
|
||||
self.assertEqual(events[-1].kind, StreamEventKind.FINISHED)
|
||||
session = service.list_sessions()[0]
|
||||
self.assertEqual(service.load_session(session.id).session.message_count, 2)
|
||||
|
||||
def test_cancel_drops_partial_exchange(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
provider = _StreamingProvider()
|
||||
service = ChatService(
|
||||
SQLiteSessionStore(Path(directory) / "sessions.sqlite3"),
|
||||
provider,
|
||||
)
|
||||
controller = DesktopController(
|
||||
SimpleNamespace(service=service, hybrid_provider=provider)
|
||||
)
|
||||
cancel = threading.Event()
|
||||
cancel.set()
|
||||
|
||||
events = list(controller.stream_message("Hallo", cancel))
|
||||
|
||||
self.assertEqual(events[-1].state, GenerationState.ABORTED)
|
||||
session = service.list_sessions()[0]
|
||||
self.assertEqual(service.load_session(session.id).messages, [])
|
||||
|
||||
Reference in New Issue
Block a user