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

57 lines
1.8 KiB
Python

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