Files
EvoScientist-Multi/EvoScientist/llm/gateway_proxy.py
T
m4 3683cbfc13 fix: propagate streamed-model attempt id via shared configurable
The legacy BaseChatModel.astream path does not forward run_manager to
_astream, so the per-call tracing run_id was unreachable and the stream
fell back to the conversation run_id, which the gateway's
verify_model_attempt rejected (401 RUN_ATTEMPT_NOT_ACCEPTED). Publish the
tracing run_id into the shared configurable dict from
on_chat_model_start and read it back in _attempt_id.
2026-08-20 15:32:18 +08:00

161 lines
5.9 KiB
Python

"""Secretless ChatModel proxy for Ai4Sci Graph-native runs."""
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, ChatGenerationChunk, ChatResult
from langchain_core.tools import BaseTool
from langchain_core.utils.function_calling import convert_to_openai_tool
from pydantic import Field
class GatewayProxyChatModel(BaseChatModel):
gateway_url: str
run_id: str
envelope_signature: str
provider_id: str = ""
model_id: str = ""
bound_tools: list[dict[str, Any]] = Field(default_factory=list)
bound_tool_choice: Any = None
@property
def _llm_type(self) -> str:
return "ai4sci-gateway-proxy"
@property
def _identifying_params(self) -> dict[str, Any]:
return {"provider_id": self.provider_id, "model_id": self.model_id}
def bind_tools(
self,
tools: Sequence[dict[str, Any] | type | BaseTool],
*,
tool_choice: str | dict[str, Any] | bool | None = None,
**kwargs: Any,
):
del kwargs
serialized = [convert_to_openai_tool(tool) for tool in tools]
return self.model_copy(
update={"bound_tools": serialized, "bound_tool_choice": tool_choice}
)
def _attempt_id(self, run_manager: Any) -> str:
# The metering callback publishes the per-call tracing run_id into the
# shared configurable dict (see RecoverableMeteringCallback). The legacy
# BaseChatModel.astream path does NOT forward run_manager to _astream, so
# run_manager.run_id is unavailable here; the configurable channel is the
# only reliable source of the id the gateway registered under.
try:
from langgraph.config import get_config
config = get_config()
except Exception:
config = None
if isinstance(config, dict):
configurable = config.get("configurable")
if isinstance(configurable, dict):
current = configurable.get("ai4sci_attempt_id")
if current:
return str(current)
return str(getattr(run_manager, "run_id", None) or self.run_id)
def _generate(self, *args: Any, **kwargs: Any) -> ChatResult:
del args, kwargs
raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_MODEL_PATH")
async def _agenerate(
self,
messages: list[BaseMessage],
stop: list[str] | None = None,
run_manager: Any = None,
**kwargs: Any,
) -> ChatResult:
del stop, kwargs
attempt_id = self._attempt_id(run_manager)
async with httpx.AsyncClient(timeout=httpx.Timeout(660.0, connect=5.0)) as client:
response = await client.post(
f"{self.gateway_url.rstrip('/')}/api/internal/recoverable-runs/model/invoke",
json={
"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,
},
)
response.raise_for_status()
value = response.json()
parsed = messages_from_dict([value["message"]])
if len(parsed) != 1:
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 = self._attempt_id(run_manager)
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 = ""
) -> GatewayProxyChatModel:
required = {
name: str(value.get(name) or "")
for name in ("gateway_url", "run_id", "envelope_signature")
}
if not all(required.values()):
raise RuntimeError("AI4SCI_MODEL_PROXY_CONFIG_INVALID")
return GatewayProxyChatModel(
**required,
provider_id=provider_id,
model_id=model_id,
)