feat: add streaming _astream to gateway proxy chat model

This commit is contained in:
m4
2026-08-20 09:11:10 +08:00
parent 8d7d95a20d
commit 386c8130ea
2 changed files with 111 additions and 1 deletions
+48 -1
View File
@@ -2,13 +2,14 @@
from __future__ import annotations
import json
from collections.abc import Mapping, Sequence
from typing import Any
import httpx
from langchain_core.language_models.chat_models import BaseChatModel
from langchain_core.messages import BaseMessage, messages_from_dict, messages_to_dict
from langchain_core.outputs import ChatGeneration, ChatResult
from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult
from langchain_core.tools import BaseTool
from langchain_core.utils.function_calling import convert_to_openai_tool
from pydantic import Field
@@ -76,6 +77,52 @@ class GatewayProxyChatModel(BaseChatModel):
raise RuntimeError("AI4SCI_MODEL_PROXY_RESPONSE_INVALID")
return ChatResult(generations=[ChatGeneration(message=parsed[0])])
async def _astream(
self,
messages: list[BaseMessage],
stop: list[str] | None = None,
run_manager: Any = None,
**kwargs: Any,
):
del stop, kwargs
attempt_id = str(getattr(run_manager, "run_id", None) or self.run_id)
payload = {
"run_id": self.run_id,
"attempt_id": attempt_id,
"envelope_signature": self.envelope_signature,
"messages": messages_to_dict(messages),
"tools": self.bound_tools,
"tool_choice": self.bound_tool_choice,
"stream": True,
}
async with httpx.AsyncClient(timeout=httpx.Timeout(660.0, connect=5.0)) as client:
async with client.stream(
"POST",
f"{self.gateway_url.rstrip('/')}/api/internal/recoverable-runs/model/stream",
json=payload,
) as response:
response.raise_for_status()
async for line in response.aiter_lines():
if not line.startswith("data:"):
continue
data = line[len("data:"):].strip()
if data == "[DONE]":
break
chunk = json.loads(data)
message = _chunk_to_message(chunk)
yield ChatGenerationChunk(
message=message,
generation_info=chunk.get("generation_info"),
)
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", "")}}]
)
return parsed[0]
def proxy_from_config(
value: Mapping[str, Any], *, provider_id: str = "", model_id: str = ""
+63
View File
@@ -0,0 +1,63 @@
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"