c683f6e739
Docker / build (push) Has been cancelled
Build / 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
Add bounded document ingestion, controlled web search, recoverable session support, subagent timeouts, and the native sandbox runtime contract. Unify package versioning and add release-focused regression coverage.
217 lines
7.5 KiB
Python
217 lines
7.5 KiB
Python
"""At-least-once tool-effect recovery for Ai4Sci Graph runs."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import json
|
|
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
|
|
|
|
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 = {
|
|
"web_search",
|
|
"glob",
|
|
"grep",
|
|
"ls",
|
|
"view_image",
|
|
"fetch_url",
|
|
}
|
|
_IDEMPOTENT_NAMES = {
|
|
"write_file",
|
|
"create_directory",
|
|
"mkdir",
|
|
"update_file",
|
|
}
|
|
|
|
|
|
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"],
|
|
},
|
|
)
|
|
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)
|
|
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(tool_name),
|
|
"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:
|
|
await _post(
|
|
proxy,
|
|
"terminal",
|
|
{
|
|
"effect_id": effect_id,
|
|
"fencing_token": fencing_token,
|
|
"outcome": "failed",
|
|
"result": {},
|
|
},
|
|
)
|
|
raise
|
|
successful = isinstance(result, ToolMessage) and result.status != "error"
|
|
payload = (
|
|
{"message": message_to_dict(result)}
|
|
if isinstance(result, ToolMessage)
|
|
else {"command_result": True}
|
|
)
|
|
await _post(
|
|
proxy,
|
|
"terminal",
|
|
{
|
|
"effect_id": effect_id,
|
|
"fencing_token": fencing_token,
|
|
"outcome": "succeeded" if successful else "failed",
|
|
"result": payload,
|
|
},
|
|
)
|
|
return result
|