Files
EvoScientist-Multi/EvoScientist/llm/gateway_proxy.py
T
m4 561e161123
Build / build (push) Has been cancelled
Docker / build (push) Has been cancelled
Lint / ruff (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.11) (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.12) (push) Has been cancelled
Test / pytest (windows-latest, 3.11) (push) Has been cancelled
Test / pytest (windows-latest, 3.12) (push) Has been cancelled
fix: fail-closed internal identity and preserve billing error semantics
- Gateway internal identity: when a service token is configured, reject
  wrong/missing tokens even from loopback (closes SSRF/local bypass).
- Terminal metering: classified AgentControlError propagates without
  retry; exhausted retries raise BILLING_UNAVAILABLE instead of a
  generic RuntimeError, keeping error attribution accurate.
2026-09-03 18:45:31 +08:00

194 lines
7.2 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
from EvoScientist.internal_service import internal_service_headers
from .contracts import EvoRuntimeError
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,
},
headers=internal_service_headers(),
)
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,
headers=internal_service_headers(),
) as response:
response.raise_for_status()
saw_done = False
async for line in response.aiter_lines():
if not line.startswith("data:"):
continue
data = line[len("data:") :].strip()
if data == "[DONE]":
saw_done = True
break
chunk = json.loads(data)
if chunk.get("type") == "error":
code = str(chunk.get("code") or "MODEL_PROVIDER_ERROR")
message = str(chunk.get("message") or code)
details = {
key: value
for key, value in {
"http_status": chunk.get("status"),
"retryable": chunk.get("retryable"),
}.items()
if isinstance(value, int | bool)
}
raise EvoRuntimeError(
code,
message,
details=(details,) if details else (),
)
message = _chunk_to_message(chunk)
yield ChatGenerationChunk(
message=message,
generation_info=chunk.get("generation_info"),
)
if not saw_done:
raise RuntimeError("AI4SCI_MODEL_STREAM_INCOMPLETE")
def _chunk_to_message(chunk: dict[str, Any]) -> BaseMessage:
delta = chunk.get("delta") or {}
message_dict = delta.get("message")
if message_dict is None:
raise RuntimeError("AI4SCI_MODEL_STREAM_DELTA_INVALID")
parsed = messages_from_dict([message_dict])
if len(parsed) != 1:
raise RuntimeError("AI4SCI_MODEL_STREAM_DELTA_INVALID")
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,
)