feat: add streaming _astream to gateway proxy chat model
This commit is contained in:
@@ -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 = ""
|
||||
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user