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()