57 lines
1.8 KiB
Python
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()
|