132 lines
4.7 KiB
Python
132 lines
4.7 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_manual_title_is_normalized_and_limited(self) -> None:
|
|
session = self.service.new_session()
|
|
|
|
renamed = self.service.rename_session(session.id, " Mein Titel ")
|
|
shortened = self.service.rename_session(session.id, "x" * 80)
|
|
|
|
self.assertEqual(renamed.title, "Mein Titel")
|
|
self.assertEqual(len(shortened.title), 60)
|
|
self.assertTrue(shortened.title.endswith("…"))
|
|
|
|
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()
|