feat: stream provider responses
This commit is contained in:
@@ -4,22 +4,39 @@ import unittest
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from typing import ClassVar
|
||||
|
||||
from javis.providers.base import ChatMessage
|
||||
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))
|
||||
body = json.dumps(
|
||||
{
|
||||
"message": {"role": "assistant", "content": "Lokale Antwort"},
|
||||
"done": True,
|
||||
}
|
||||
).encode()
|
||||
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)))
|
||||
@@ -39,6 +56,9 @@ class _OllamaHandler(BaseHTTPRequestHandler):
|
||||
|
||||
|
||||
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)
|
||||
@@ -62,6 +82,44 @@ class OllamaProviderTests(unittest.TestCase):
|
||||
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()
|
||||
|
||||
Reference in New Issue
Block a user