fix: gateway proxy reads delta.message and detects truncated SSE
Standardize the stream delta on LangChain's own message dict (delta.message)
instead of a bespoke {text, tool_calls} shape, and raise when the stream ends
without the [DONE] sentinel so mid-stream truncation is no longer silent.
This commit is contained in:
@@ -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]
|
||||
|
||||
|
||||
|
||||
@@ -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")])]
|
||||
|
||||
Reference in New Issue
Block a user