Files
Jarvis-Ai/tests/unit/test_chat_service.py
T

122 lines
4.3 KiB
Python

import tempfile
import unittest
from pathlib import Path
from javis.core.chat_service import ChatService, derive_session_title
from javis.memory.sqlite_store import SQLiteSessionStore
from javis.providers.base import ChatMessage, LocalModelProvider, ProviderUnavailableError
class RecordingProvider:
name = "ollama"
model = "test-model"
def __init__(self) -> None:
self.calls: list[list[ChatMessage]] = []
def chat(self, messages: list[ChatMessage]) -> str:
self.calls.append(messages)
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()
database = Path(self.temporary_directory.name) / "sessions.sqlite3"
self.provider = RecordingProvider()
self.service = ChatService(SQLiteSessionStore(database), self.provider)
def tearDown(self) -> None:
self.temporary_directory.cleanup()
def test_test_double_implements_provider_contract(self) -> None:
self.assertIsInstance(self.provider, LocalModelProvider)
def test_send_passes_and_persists_conversation_history(self) -> None:
session = self.service.new_session()
first = self.service.send(session.id, "Hallo")
second = self.service.send(session.id, "Und jetzt?")
self.assertEqual(first, "Antwort 1")
self.assertEqual(second, "Antwort 2")
self.assertEqual(len(self.provider.calls[1]), 3)
loaded = self.service.load_session(session.id)
self.assertEqual(
[message.content for message in loaded.messages],
[
"Hallo",
"Antwort 1",
"Und jetzt?",
"Antwort 2",
],
)
self.assertEqual(loaded.session.title, "Hallo")
def test_title_is_local_short_and_redacts_obvious_sensitive_content(self) -> None:
self.assertEqual(
derive_session_title("Erkläre SQLite-Transaktionen für Anfänger"),
"Erkläre SQLite-Transaktionen für Anfänger",
)
self.assertLessEqual(derive_session_title("Wort " * 30).__len__(), 60)
self.assertEqual(
derive_session_title("Mein API-Key ist ABC123 und funktioniert nicht"),
"Sensible Anfrage",
)
self.assertEqual(
derive_session_title("Wenn ich blute, sollte ich zum Arzt?"),
"Gesundheitsfrage",
)
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()