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

126 lines
4.4 KiB
Python

import json
import threading
import unittest
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import ClassVar
from javis.providers.base import ChatMessage, ResponseAbortedError
from javis.providers.ollama import OllamaProvider
class _OllamaHandler(BaseHTTPRequestHandler):
request_payload: ClassVar[dict[str, object]] = {}
abort_stream: ClassVar[bool] = False
def do_POST(self) -> None:
length = int(self.headers["Content-Length"])
type(self).request_payload = json.loads(self.rfile.read(length))
if type(self).request_payload["stream"]:
chunks = [
{
"message": {"role": "assistant", "content": "Lokale "},
"done": False,
}
]
if not type(self).abort_stream:
chunks.append(
{
"message": {"role": "assistant", "content": "Antwort"},
"done": True,
}
)
body = b"".join(json.dumps(chunk).encode() + b"\n" for chunk in chunks)
else:
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 do_GET(self) -> None:
body = json.dumps({"models": [{"name": "test-model"}]}).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 setUp(self) -> None:
_OllamaHandler.abort_stream = False
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")])
status = provider.probe()
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"])
self.assertTrue(status.reachable)
self.assertTrue(status.model_available)
def test_provider_streams_visible_text_without_thinking(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,
)
chunks = list(provider.stream_chat([ChatMessage("user", "Hallo")]))
finally:
server.shutdown()
server.server_close()
thread.join()
self.assertEqual(chunks, ["Lokale ", "Antwort"])
self.assertTrue(_OllamaHandler.request_payload["stream"])
self.assertFalse(_OllamaHandler.request_payload["think"])
def test_provider_rejects_incomplete_stream(self) -> None:
_OllamaHandler.abort_stream = True
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,
)
with self.assertRaises(ResponseAbortedError):
list(provider.stream_chat([ChatMessage("user", "Hallo")]))
finally:
server.shutdown()
server.server_close()
thread.join()
if __name__ == "__main__":
unittest.main()