test: cover chat core and session storage
This commit is contained in:
@@ -0,0 +1,57 @@
|
||||
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()
|
||||
@@ -0,0 +1,40 @@
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from javis.core.chat_service import ChatService
|
||||
from javis.interface.cli import run_chat
|
||||
from javis.memory.sqlite_store import SQLiteSessionStore
|
||||
from javis.providers.base import ChatMessage
|
||||
|
||||
|
||||
class _FakeProvider:
|
||||
name = "ollama"
|
||||
model = "test-model"
|
||||
|
||||
def chat(self, messages: list[ChatMessage]) -> str:
|
||||
return f"Echo: {messages[-1].content}"
|
||||
|
||||
|
||||
class CliTests(unittest.TestCase):
|
||||
def test_basic_chat_commands(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
store = SQLiteSessionStore(Path(directory) / "sessions.sqlite3")
|
||||
service = ChatService(store, _FakeProvider())
|
||||
inputs = iter(["Hallo", "/sessions", "/new", "/clear", "/exit"])
|
||||
output: list[str] = []
|
||||
|
||||
result = run_chat(
|
||||
service,
|
||||
input_fn=lambda _prompt: next(inputs),
|
||||
output=output.append,
|
||||
)
|
||||
|
||||
self.assertEqual(result, 0)
|
||||
self.assertTrue(any(line == "Javis: Echo: Hallo" for line in output))
|
||||
self.assertTrue(any("Nachrichten" in line for line in output))
|
||||
self.assertEqual(output[-1], "Chat beendet.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,56 @@
|
||||
import json
|
||||
import threading
|
||||
import unittest
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from typing import ClassVar
|
||||
|
||||
from javis.providers.base import ChatMessage
|
||||
from javis.providers.ollama import OllamaProvider
|
||||
|
||||
|
||||
class _OllamaHandler(BaseHTTPRequestHandler):
|
||||
request_payload: ClassVar[dict[str, object]] = {}
|
||||
|
||||
def do_POST(self) -> None:
|
||||
length = int(self.headers["Content-Length"])
|
||||
type(self).request_payload = json.loads(self.rfile.read(length))
|
||||
body = json.dumps(
|
||||
{
|
||||
"message": {"role": "assistant", "content": "Lokale Antwort"},
|
||||
"done": True,
|
||||
}
|
||||
).encode()
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
|
||||
def log_message(self, format: str, *args: object) -> None:
|
||||
return
|
||||
|
||||
|
||||
class OllamaProviderTests(unittest.TestCase):
|
||||
def test_provider_uses_local_chat_endpoint(self) -> None:
|
||||
server = ThreadingHTTPServer(("127.0.0.1", 0), _OllamaHandler)
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
provider = OllamaProvider(
|
||||
"test-model",
|
||||
f"http://127.0.0.1:{server.server_port}",
|
||||
2,
|
||||
)
|
||||
response = provider.chat([ChatMessage("user", "Hallo")])
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join()
|
||||
|
||||
self.assertEqual(response, "Lokale Antwort")
|
||||
self.assertEqual(_OllamaHandler.request_payload["model"], "test-model")
|
||||
self.assertFalse(_OllamaHandler.request_payload["stream"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,39 @@
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from javis.config.settings import ConfigurationError, Settings, default_data_dir
|
||||
|
||||
|
||||
class SettingsTests(unittest.TestCase):
|
||||
def test_windows_default_uses_local_app_data(self) -> None:
|
||||
result = default_data_dir(
|
||||
{"LOCALAPPDATA": "C:\\Local"},
|
||||
platform_name="nt",
|
||||
home=Path("C:\\Users\\test"),
|
||||
)
|
||||
self.assertEqual(result, Path("C:\\Local") / "Javis")
|
||||
|
||||
def test_posix_default_uses_xdg_data_home(self) -> None:
|
||||
result = default_data_dir(
|
||||
{"XDG_DATA_HOME": "/tmp/data"},
|
||||
platform_name="posix",
|
||||
home=Path("/home/test"),
|
||||
)
|
||||
self.assertEqual(result, Path("/tmp/data/javis"))
|
||||
|
||||
def test_explicit_data_dir_must_be_absolute(self) -> None:
|
||||
with self.assertRaises(ConfigurationError):
|
||||
Settings.from_env({"JAVIS_DATA_DIR": "relative"})
|
||||
|
||||
def test_provider_url_must_remain_local(self) -> None:
|
||||
with self.assertRaises(ConfigurationError):
|
||||
Settings.from_env(
|
||||
{
|
||||
"JAVIS_DATA_DIR": str(Path.cwd()),
|
||||
"JAVIS_OLLAMA_URL": "https://example.com",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,58 @@
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from javis.memory.sqlite_store import SessionNotFoundError, SQLiteSessionStore
|
||||
|
||||
|
||||
class SQLiteSessionStoreTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.temporary_directory = tempfile.TemporaryDirectory()
|
||||
database = Path(self.temporary_directory.name) / "sessions.sqlite3"
|
||||
self.store = SQLiteSessionStore(database)
|
||||
|
||||
def tearDown(self) -> None:
|
||||
self.temporary_directory.cleanup()
|
||||
|
||||
def test_create_list_and_load_session(self) -> None:
|
||||
created = self.store.create_session("ollama", "test-model")
|
||||
loaded = self.store.get_session(created.id)
|
||||
|
||||
self.assertEqual(loaded.id, created.id)
|
||||
self.assertEqual(loaded.provider, "ollama")
|
||||
self.assertEqual(loaded.model, "test-model")
|
||||
self.assertEqual(self.store.list_sessions()[0].id, created.id)
|
||||
|
||||
def test_messages_keep_exchange_order(self) -> None:
|
||||
session = self.store.create_session("ollama", "test-model")
|
||||
self.store.append_exchange(session.id, "eins", "zwei")
|
||||
self.store.append_exchange(session.id, "drei", "vier")
|
||||
|
||||
messages = self.store.get_messages(session.id)
|
||||
|
||||
self.assertEqual(
|
||||
[(message.role, message.content) for message in messages],
|
||||
[
|
||||
("user", "eins"),
|
||||
("assistant", "zwei"),
|
||||
("user", "drei"),
|
||||
("assistant", "vier"),
|
||||
],
|
||||
)
|
||||
|
||||
def test_unknown_session_raises_clear_error(self) -> None:
|
||||
with self.assertRaises(SessionNotFoundError):
|
||||
self.store.get_session("missing")
|
||||
|
||||
def test_clear_keeps_session_and_removes_messages(self) -> None:
|
||||
session = self.store.create_session("ollama", "test-model")
|
||||
self.store.append_exchange(session.id, "eins", "zwei")
|
||||
|
||||
self.store.clear_messages(session.id)
|
||||
|
||||
self.assertEqual(self.store.get_messages(session.id), [])
|
||||
self.assertEqual(self.store.get_session(session.id).message_count, 0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user