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, [])
|
||||
|
||||
@@ -4,7 +4,7 @@ import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from javis.core.provider_router import HybridProvider
|
||||
from javis.core.provider_router import ApprovalChoice, HybridProvider
|
||||
from javis.memory.usage_store import ProviderEvent, SQLiteUsageStore
|
||||
from javis.providers.base import (
|
||||
ChatMessage,
|
||||
@@ -182,6 +182,19 @@ class HybridProviderTests(unittest.TestCase):
|
||||
self.assertEqual(answer, "Cloud")
|
||||
self.assertEqual(len(self.cloud.calls), 1)
|
||||
|
||||
def test_ask_supports_allow_local_and_cancel_choices(self) -> None:
|
||||
question = [ChatMessage("user", "Meine Familie plant Urlaub")]
|
||||
|
||||
allowed = self._router(approval_callback=lambda _decision: ApprovalChoice.ALLOW)
|
||||
self.assertEqual(allowed.chat(question), "Cloud")
|
||||
|
||||
local = self._router(approval_callback=lambda _decision: ApprovalChoice.LOCAL)
|
||||
self.assertEqual(local.chat(question), "Lokal")
|
||||
|
||||
cancelled = self._router(approval_callback=lambda _decision: ApprovalChoice.CANCEL)
|
||||
with self.assertRaises(ResponseAbortedError):
|
||||
cancelled.chat(question)
|
||||
|
||||
def test_429_falls_back_without_losing_local_answer(self) -> None:
|
||||
self.cloud.error = ProviderRateLimitError("quota")
|
||||
router = self._router()
|
||||
|
||||
Reference in New Issue
Block a user