3683cbfc13
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.
161 lines
5.9 KiB
Python
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,
|
|
)
|