270 lines
9.5 KiB
Python
270 lines
9.5 KiB
Python
"""At-least-once tool-effect recovery for Ai4Sci Graph runs."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import hashlib
|
|
import json
|
|
import logging
|
|
from collections.abc import Awaitable, Callable, Mapping
|
|
from typing import TYPE_CHECKING, Any
|
|
|
|
import httpx
|
|
from langchain.agents.middleware.types import AgentMiddleware
|
|
from langchain_core.messages import ToolMessage, message_to_dict, messages_from_dict
|
|
from langgraph.types import Command
|
|
|
|
try:
|
|
from langgraph.errors import GraphInterrupt as _GraphInterrupt
|
|
except ImportError: # pragma: no cover - compatibility with older LangGraph
|
|
_GraphInterrupt = None # type: ignore[assignment,misc]
|
|
|
|
from EvoScientist.internal_service import internal_service_headers
|
|
from EvoScientist.llm.contracts import EvoRuntimeError
|
|
|
|
if TYPE_CHECKING:
|
|
from langchain.agents.middleware.types import ToolCallRequest
|
|
|
|
_READ_ONLY_PREFIXES = ("read_", "get_", "list_", "search_", "find_", "check_")
|
|
_READ_ONLY_NAMES = {
|
|
"tavily_search",
|
|
"web_search",
|
|
"glob",
|
|
"grep",
|
|
"ls",
|
|
"view_image",
|
|
"fetch_url",
|
|
}
|
|
_IDEMPOTENT_NAMES = {
|
|
"write_file",
|
|
"create_directory",
|
|
"mkdir",
|
|
"update_file",
|
|
}
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def _callback_unavailable(exc: BaseException) -> bool:
|
|
if isinstance(exc, httpx.TransportError):
|
|
return True
|
|
return (
|
|
isinstance(exc, httpx.HTTPStatusError)
|
|
and exc.response.status_code >= 500
|
|
)
|
|
|
|
|
|
def _canonical(value: Any) -> str:
|
|
return json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"), default=str)
|
|
|
|
|
|
def _hash(value: Any) -> str:
|
|
return hashlib.sha256(_canonical(value).encode()).hexdigest()
|
|
|
|
|
|
def _context() -> tuple[dict[str, str] | None, dict[str, Any]]:
|
|
try:
|
|
from langgraph.config import get_config
|
|
|
|
config = get_config()
|
|
except Exception:
|
|
return None, {}
|
|
configurable = config.get("configurable") if isinstance(config, dict) else None
|
|
if not isinstance(configurable, Mapping):
|
|
return None, {}
|
|
proxy = configurable.get("ai4sci_tool_effect")
|
|
metadata = dict(config.get("metadata") or {})
|
|
if not isinstance(proxy, Mapping):
|
|
# Compatibility for Runs dispatched before the tool-effect grant was
|
|
# split from the model proxy. Detached EvoMemory graphs must never use
|
|
# the parent conversation's tool-effect authority.
|
|
run_kind = str(metadata.get("run_kind") or "")
|
|
if not run_kind.startswith("evomemory_"):
|
|
proxy = configurable.get("ai4sci_model_proxy")
|
|
if not isinstance(proxy, Mapping):
|
|
return None, metadata
|
|
normalized = {
|
|
name: str(proxy.get(name) or "")
|
|
for name in ("gateway_url", "run_id", "envelope_signature")
|
|
}
|
|
return (
|
|
normalized if all(normalized.values()) else None,
|
|
metadata,
|
|
)
|
|
|
|
|
|
def _effect_class(tool_name: str) -> str:
|
|
lowered = tool_name.lower()
|
|
if lowered in _READ_ONLY_NAMES or lowered.startswith(_READ_ONLY_PREFIXES):
|
|
return "read_only"
|
|
if lowered in _IDEMPOTENT_NAMES:
|
|
return "idempotent"
|
|
return "non_idempotent"
|
|
|
|
|
|
async def _post(proxy: Mapping[str, str], phase: str, payload: dict[str, Any]) -> dict[str, Any]:
|
|
async with httpx.AsyncClient(timeout=httpx.Timeout(30.0, connect=3.0)) as client:
|
|
response = await client.post(
|
|
f"{proxy['gateway_url'].rstrip('/')}/api/internal/recoverable-runs/tool-effect/{phase}",
|
|
json={
|
|
**payload,
|
|
"run_id": proxy["run_id"],
|
|
"attempt_id": proxy["run_id"],
|
|
"envelope_signature": proxy["envelope_signature"],
|
|
},
|
|
headers=internal_service_headers(),
|
|
)
|
|
if response.is_error:
|
|
try:
|
|
error_body = response.json()
|
|
except ValueError:
|
|
error_body = None
|
|
detail = error_body.get("detail") if isinstance(error_body, Mapping) else None
|
|
code = detail.get("code") if isinstance(detail, Mapping) else None
|
|
if isinstance(code, str) and code.isascii() and code.replace("_", "").isalnum():
|
|
raise EvoRuntimeError(code)
|
|
response.raise_for_status()
|
|
return dict(response.json())
|
|
|
|
|
|
class RecoverableToolEffectMiddleware(AgentMiddleware):
|
|
name = "recoverable_tool_effect"
|
|
|
|
def wrap_tool_call(
|
|
self,
|
|
request: ToolCallRequest,
|
|
handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]],
|
|
) -> ToolMessage | Command[Any]:
|
|
proxy, _ = _context()
|
|
if proxy is None:
|
|
return handler(request)
|
|
raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_TOOL_PATH")
|
|
|
|
async def awrap_tool_call(
|
|
self,
|
|
request: ToolCallRequest,
|
|
handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]],
|
|
) -> ToolMessage | Command[Any]:
|
|
proxy, metadata = _context()
|
|
if proxy is None:
|
|
return await handler(request)
|
|
tool_call = dict(request.tool_call)
|
|
tool_name = str(tool_call.get("name") or "unknown_tool")
|
|
tool_call_id = str(tool_call.get("id") or "")
|
|
arguments = tool_call.get("args") or {}
|
|
request_hash = _hash(arguments)
|
|
effect_class = _effect_class(tool_name)
|
|
checkpoint_ns = str(metadata.get("checkpoint_ns") or metadata.get("langgraph_checkpoint_ns") or "")
|
|
task_path = ":".join(
|
|
str(metadata.get(name) or "")
|
|
for name in ("langgraph_step", "langgraph_node", "langgraph_task_idx")
|
|
)
|
|
effect_id = _hash(
|
|
{
|
|
"run_id": proxy["run_id"],
|
|
"checkpoint_ns": checkpoint_ns,
|
|
"task_path": task_path,
|
|
"tool_call_id": tool_call_id,
|
|
"tool_name": tool_name,
|
|
"request_hash": request_hash,
|
|
}
|
|
)
|
|
prepared = await _post(
|
|
proxy,
|
|
"prepare",
|
|
{
|
|
"effect_id": effect_id,
|
|
"checkpoint_ns": checkpoint_ns,
|
|
"task_path": task_path,
|
|
"tool_call_id": tool_call_id,
|
|
"tool_name": tool_name,
|
|
"effect_class": effect_class,
|
|
"request_hash": request_hash,
|
|
},
|
|
)
|
|
if prepared.get("action") == "manual_reconcile":
|
|
return ToolMessage(
|
|
content=(
|
|
f"Tool '{tool_name}' may already have produced an external side effect. "
|
|
"Automatic retry is blocked; user confirmation is required."
|
|
),
|
|
tool_call_id=tool_call_id,
|
|
name=tool_name,
|
|
status="error",
|
|
)
|
|
if prepared.get("action") == "cached":
|
|
result = prepared.get("result")
|
|
if isinstance(result, dict) and isinstance(result.get("message"), dict):
|
|
messages = messages_from_dict([result["message"]])
|
|
if len(messages) == 1 and isinstance(messages[0], ToolMessage):
|
|
return messages[0]
|
|
return ToolMessage(
|
|
content="Cached tool result is unavailable; user confirmation is required.",
|
|
tool_call_id=tool_call_id,
|
|
name=tool_name,
|
|
status="error",
|
|
)
|
|
fencing_token = int(prepared["fencing_token"])
|
|
try:
|
|
result = await handler(request)
|
|
except BaseException as exc:
|
|
if isinstance(exc, asyncio.CancelledError) or (
|
|
_GraphInterrupt is not None and isinstance(exc, _GraphInterrupt)
|
|
):
|
|
raise
|
|
try:
|
|
await _post(
|
|
proxy,
|
|
"terminal",
|
|
{
|
|
"effect_id": effect_id,
|
|
"fencing_token": fencing_token,
|
|
"outcome": "failed",
|
|
"result": {},
|
|
},
|
|
)
|
|
except BaseException as callback_exc:
|
|
if not _callback_unavailable(callback_exc):
|
|
raise
|
|
logger.exception(
|
|
"tool-effect failure callback transport failed run_id=%s tool=%s effect_id=%s",
|
|
proxy["run_id"],
|
|
tool_name,
|
|
effect_id,
|
|
)
|
|
raise
|
|
successful = isinstance(result, ToolMessage) and result.status != "error"
|
|
payload = (
|
|
{"message": message_to_dict(result)}
|
|
if isinstance(result, ToolMessage)
|
|
else {"command_result": True}
|
|
)
|
|
try:
|
|
await _post(
|
|
proxy,
|
|
"terminal",
|
|
{
|
|
"effect_id": effect_id,
|
|
"fencing_token": fencing_token,
|
|
"outcome": "succeeded" if successful else "failed",
|
|
"result": payload,
|
|
},
|
|
)
|
|
except BaseException as callback_exc:
|
|
if not _callback_unavailable(callback_exc):
|
|
raise
|
|
logger.exception(
|
|
"tool-effect terminal callback transport failed run_id=%s tool=%s effect_id=%s",
|
|
proxy["run_id"],
|
|
tool_name,
|
|
effect_id,
|
|
)
|
|
if effect_class != "read_only":
|
|
raise EvoRuntimeError("TOOL_EFFECT_TERMINAL_UNAVAILABLE") from None
|
|
return ToolMessage(
|
|
content="TOOL_EFFECT_TERMINAL_UNAVAILABLE",
|
|
tool_call_id=tool_call_id,
|
|
name=tool_name,
|
|
status="error",
|
|
)
|
|
return result
|