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:
m4
2026-08-20 15:32:27 +08:00
parent 3683cbfc13
commit ce0bec1f91
2 changed files with 65 additions and 5 deletions
+10 -3
View File
@@ -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]
+55 -2
View File
@@ -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")])]