diff --git a/EvoScientist/llm/gateway_proxy.py b/EvoScientist/llm/gateway_proxy.py index 6fb9a2e..94a9b35 100644 --- a/EvoScientist/llm/gateway_proxy.py +++ b/EvoScientist/llm/gateway_proxy.py @@ -122,11 +122,13 @@ class GatewayProxyChatModel(BaseChatModel): json=payload, ) as response: response.raise_for_status() + saw_done = False async for line in response.aiter_lines(): if not line.startswith("data:"): continue data = line[len("data:"):].strip() if data == "[DONE]": + saw_done = True break chunk = json.loads(data) message = _chunk_to_message(chunk) @@ -134,13 +136,18 @@ class GatewayProxyChatModel(BaseChatModel): message=message, generation_info=chunk.get("generation_info"), ) + if not saw_done: + raise RuntimeError("AI4SCI_MODEL_STREAM_INCOMPLETE") def _chunk_to_message(chunk: dict[str, Any]) -> BaseMessage: delta = chunk.get("delta") or {} - parsed = messages_from_dict( - [delta.get("message") or {"type": "AIMessageChunk", "data": {"content": delta.get("text", "")}}] - ) + message_dict = delta.get("message") + if message_dict is None: + raise RuntimeError("AI4SCI_MODEL_STREAM_DELTA_INVALID") + parsed = messages_from_dict([message_dict]) + if len(parsed) != 1: + raise RuntimeError("AI4SCI_MODEL_STREAM_DELTA_INVALID") return parsed[0] diff --git a/tests/test_gateway_proxy.py b/tests/test_gateway_proxy.py index 6107156..cccf7ad 100644 --- a/tests/test_gateway_proxy.py +++ b/tests/test_gateway_proxy.py @@ -1,3 +1,5 @@ +import json + import httpx import pytest from langchain_core.messages import HumanMessage @@ -48,9 +50,10 @@ async def test_astream_yields_chunks_from_sse(monkeypatch): run_id="run-1", envelope_signature="sig", ) + msg = {"type": "AIMessageChunk", "data": {"content": "hello"}} lines = [ - 'data: {"delta": {"text": "hello"}}\n', - 'data: {"delta": {"text": " world"}}\n', + f'data: {json.dumps({"delta": {"message": msg}})}\n', + 'data: {"delta": {"message": {"type": "AIMessageChunk", "data": {"content": " world"}}}}\n', "data: [DONE]\n", ] fake = _FakeClient(lines) @@ -61,3 +64,53 @@ async def test_astream_yields_chunks_from_sse(monkeypatch): assert fake.sent["stream"] is True assert len(chunks) == 2 assert chunks[0].message.content == "hello" + + +@pytest.mark.anyio +async def test_astream_roundtrips_streaming_tool_call_chunks(monkeypatch): + model = GatewayProxyChatModel( + gateway_url="http://gw", + run_id="run-1", + envelope_signature="sig", + ) + msg = { + "type": "AIMessageChunk", + "data": { + "content": "", + "tool_call_chunks": [ + { + "name": "read_file", + "args": '{"path":', + "id": "call-1", + "index": 0, + "type": "tool_call_chunk", + } + ], + }, + } + lines = [f'data: {json.dumps({"delta": {"message": msg}})}\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 len(chunks) == 1 + assert chunks[0].message.tool_call_chunks[0]["name"] == "read_file" + assert chunks[0].message.tool_call_chunks[0]["id"] == "call-1" + + +@pytest.mark.anyio +async def test_astream_raises_on_missing_done(monkeypatch): + model = GatewayProxyChatModel( + gateway_url="http://gw", + run_id="run-1", + envelope_signature="sig", + ) + lines = [ + 'data: {"delta": {"message": {"type": "AIMessageChunk", "data": {"content": "hi"}}}}\n' + ] + fake = _FakeClient(lines) + monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: fake) + + with pytest.raises(RuntimeError, match="AI4SCI_MODEL_STREAM_INCOMPLETE"): + _ = [c async for c in model._astream([HumanMessage(content="hi")])]