Files
EvoScientist-Multi/EvoScientist/middleware/recoverable_tools.py
T

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