156 lines
5.3 KiB
Python
156 lines
5.3 KiB
Python
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, ResponseAbortedError
|
|
from javis.ui.chat_controller import DesktopController, sanitized_error
|
|
from javis.ui.models import GenerationState, StreamEventKind
|
|
|
|
|
|
class _Service:
|
|
def __init__(self) -> None:
|
|
self.sessions = [
|
|
SimpleNamespace(id="one", title="Eins"),
|
|
SimpleNamespace(id="two", title="Zwei"),
|
|
]
|
|
|
|
def new_session(self):
|
|
return self.sessions[0]
|
|
|
|
def list_sessions(self):
|
|
return self.sessions
|
|
|
|
def search_sessions(self, text: str):
|
|
return [session for session in self.sessions if text in session.title]
|
|
|
|
def load_session(self, session_id: str):
|
|
session = next(item for item in self.sessions if item.id == session_id)
|
|
return SimpleNamespace(
|
|
session=session,
|
|
messages=[ChatMessage("user", "Hallo")],
|
|
)
|
|
|
|
def rename_session(self, session_id: str, title: str):
|
|
session = next(item for item in self.sessions if item.id == session_id)
|
|
session.title = title
|
|
return session
|
|
|
|
|
|
class _Hybrid:
|
|
mode = "auto"
|
|
last_route = None
|
|
|
|
def set_mode(self, mode: str) -> None:
|
|
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 _AbortingProvider(_StreamingProvider):
|
|
def stream_chat(self, _messages):
|
|
raise ResponseAbortedError("Vom Dialog abgebrochen")
|
|
yield
|
|
|
|
|
|
class DesktopControllerTests(unittest.TestCase):
|
|
def setUp(self) -> None:
|
|
self.runtime = SimpleNamespace(
|
|
service=_Service(),
|
|
hybrid_provider=_Hybrid(),
|
|
)
|
|
self.controller = DesktopController(self.runtime)
|
|
|
|
def test_controller_manages_sessions_without_qt_window(self) -> None:
|
|
created = self.controller.new_session()
|
|
listed = self.controller.list_sessions()
|
|
searched = self.controller.list_sessions("Zwei")
|
|
loaded = self.controller.load_session("two")
|
|
renamed = self.controller.rename_active_session("Neu")
|
|
|
|
self.assertEqual(created.id, "one")
|
|
self.assertEqual(len(listed), 2)
|
|
self.assertEqual([session.id for session in searched], ["two"])
|
|
self.assertEqual(loaded.messages[0].content, "Hallo")
|
|
self.assertEqual(renamed.title, "Neu")
|
|
|
|
def test_error_sanitizer_masks_key_patterns(self) -> None:
|
|
secret = "AIza" + "x" * 25
|
|
result = sanitized_error(RuntimeError(f"Fehler {secret}"))
|
|
|
|
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, [])
|
|
|
|
def test_provider_abort_is_reported_as_aborted_not_error(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
provider = _AbortingProvider()
|
|
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(events[-1].state, GenerationState.ABORTED)
|
|
self.assertNotIn(StreamEventKind.ERROR, [event.kind for event in events])
|