From 386c8130ea3aea5a73ad4c2f7233a0d2614b5e16 Mon Sep 17 00:00:00 2001 From: m4 Date: Thu, 20 Aug 2026 09:11:10 +0800 Subject: [PATCH] feat: add streaming _astream to gateway proxy chat model --- EvoScientist/llm/gateway_proxy.py | 49 +++++++++++++++++++++++- tests/test_gateway_proxy.py | 63 +++++++++++++++++++++++++++++++ 2 files changed, 111 insertions(+), 1 deletion(-) create mode 100644 tests/test_gateway_proxy.py diff --git a/EvoScientist/llm/gateway_proxy.py b/EvoScientist/llm/gateway_proxy.py index c0cf3ae..75531d9 100644 --- a/EvoScientist/llm/gateway_proxy.py +++ b/EvoScientist/llm/gateway_proxy.py @@ -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 = "" diff --git a/tests/test_gateway_proxy.py b/tests/test_gateway_proxy.py new file mode 100644 index 0000000..6107156 --- /dev/null +++ b/tests/test_gateway_proxy.py @@ -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"