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