64 lines
1.5 KiB
Python
64 lines
1.5 KiB
Python
import httpx
|
|
import pytest
|
|
from langchain_core.messages import HumanMessage
|
|
|
|
from EvoScientist.llm.gateway_proxy import GatewayProxyChatModel
|
|
|
|
|
|
class _FakeStream:
|
|
def __init__(self, lines):
|
|
self._lines = lines
|
|
|
|
async def __aenter__(self):
|
|
return self
|
|
|
|
async def __aexit__(self, *exc):
|
|
return False
|
|
|
|
def raise_for_status(self):
|
|
pass
|
|
|
|
async def aiter_lines(self):
|
|
for line in self._lines:
|
|
yield line
|
|
|
|
|
|
class _FakeClient:
|
|
def __init__(self, lines):
|
|
self._lines = lines
|
|
self.sent = None
|
|
|
|
async def __aenter__(self):
|
|
return self
|
|
|
|
async def __aexit__(self, *exc):
|
|
return False
|
|
|
|
def stream(self, method, url, json=None):
|
|
assert method == "POST"
|
|
assert url.endswith("/api/internal/recoverable-runs/model/stream")
|
|
self.sent = json
|
|
return _FakeStream(self._lines)
|
|
|
|
|
|
@pytest.mark.anyio
|
|
async def test_astream_yields_chunks_from_sse(monkeypatch):
|
|
model = GatewayProxyChatModel(
|
|
gateway_url="http://gw",
|
|
run_id="run-1",
|
|
envelope_signature="sig",
|
|
)
|
|
lines = [
|
|
'data: {"delta": {"text": "hello"}}\n',
|
|
'data: {"delta": {"text": " world"}}\n',
|
|
"data: [DONE]\n",
|
|
]
|
|
fake = _FakeClient(lines)
|
|
monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: fake)
|
|
|
|
chunks = [c async for c in model._astream([HumanMessage(content="hi")])]
|
|
|
|
assert fake.sent["stream"] is True
|
|
assert len(chunks) == 2
|
|
assert chunks[0].message.content == "hello"
|