58 lines
1.8 KiB
Python
58 lines
1.8 KiB
Python
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from javis.core.chat_service import ChatService
|
|
from javis.memory.sqlite_store import SQLiteSessionStore
|
|
from javis.providers.base import ChatMessage, LocalModelProvider
|
|
|
|
|
|
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 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",
|
|
],
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|