mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-09 00:40:45 +08:00
Cache inventory MCP lists and accept YES as WhatsApp confirm.
Extend short TTL reuse to listManagedNe/queryUmeNeInventory to cut list self-loops, and treat short YES/continue as confirmation when a token is pending. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
f1d025ed60
commit
d3f621b58b
6 changed files with 130 additions and 42 deletions
|
|
@ -1204,7 +1204,7 @@ def process_inbound_payload(payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
if channel_is_wa:
|
if channel_is_wa:
|
||||||
reply = (
|
reply = (
|
||||||
f"This action needs confirmation. "
|
f"This action needs confirmation. "
|
||||||
f"Reply `confirm {token}` or include `[confirm:{token}]`."
|
f"Reply `YES` / `confirm` or `confirm {token}` / `[confirm:{token}]`."
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
reply = f"该动作需要确认。请回复 `confirm {token}` 或包含 `[confirm:{token}]`。"
|
reply = f"该动作需要确认。请回复 `confirm {token}` 或包含 `[confirm:{token}]`。"
|
||||||
|
|
|
||||||
|
|
@ -967,12 +967,26 @@ class ToolExecutor:
|
||||||
continue
|
continue
|
||||||
sig = f"{tc.name}:{self._json_dumps_safe(dict(tc.arguments or {}))}"
|
sig = f"{tc.name}:{self._json_dumps_safe(dict(tc.arguments or {}))}"
|
||||||
count = int(sig_seen.get(sig, 0))
|
count = int(sig_seen.get(sig, 0))
|
||||||
|
name_low = str(tc.name or "").strip().lower()
|
||||||
|
listish = name_low.endswith(
|
||||||
|
(
|
||||||
|
"listclitargets",
|
||||||
|
"listmanagedne",
|
||||||
|
"queryumeneinventory",
|
||||||
|
)
|
||||||
|
)
|
||||||
if count >= budget:
|
if count >= budget:
|
||||||
results_by_id[tc.id] = (
|
results_by_id[tc.id] = (
|
||||||
{
|
{
|
||||||
"ok": False,
|
"ok": False,
|
||||||
"error_code": "tool_loop_guard",
|
"error_code": "tool_loop_guard",
|
||||||
"error": f"tool loop guard triggered for signature: {tc.name}",
|
"error": f"tool loop guard triggered for signature: {tc.name}",
|
||||||
|
"hint": (
|
||||||
|
"Identical list/inventory call already ran this turn; reuse prior ids/rows "
|
||||||
|
"instead of listing again."
|
||||||
|
if listish
|
||||||
|
else "Identical tool call already ran this turn; change arguments or continue without retry."
|
||||||
|
),
|
||||||
},
|
},
|
||||||
0,
|
0,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -26,6 +26,26 @@ def evaluate_risk(task: AgentTask) -> GuardrailResult:
|
||||||
return GuardrailResult(allowed=True, needs_confirmation=False, reason="Low-risk request")
|
return GuardrailResult(allowed=True, needs_confirmation=False, reason="Low-risk request")
|
||||||
|
|
||||||
|
|
||||||
|
_AFFIRMATIVE_SHORT = frozenset(
|
||||||
|
{
|
||||||
|
"yes",
|
||||||
|
"y",
|
||||||
|
"ok",
|
||||||
|
"okay",
|
||||||
|
"confirm",
|
||||||
|
"continue",
|
||||||
|
"please continue",
|
||||||
|
"go ahead",
|
||||||
|
"可以",
|
||||||
|
"确认",
|
||||||
|
"继续",
|
||||||
|
"好的",
|
||||||
|
"行",
|
||||||
|
"是",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def has_explicit_confirmation(user_text: str) -> bool:
|
def has_explicit_confirmation(user_text: str) -> bool:
|
||||||
return has_explicit_confirmation_token(user_text, token=None)
|
return has_explicit_confirmation_token(user_text, token=None)
|
||||||
|
|
||||||
|
|
@ -48,6 +68,13 @@ def has_explicit_confirmation_token(user_text: str, token: str | None) -> bool:
|
||||||
return False
|
return False
|
||||||
if f"[confirm:{t}]".lower() in low or f"confirm:{t}".lower() in low:
|
if f"[confirm:{t}]".lower() in low or f"confirm:{t}".lower() in low:
|
||||||
return True
|
return True
|
||||||
|
# WhatsApp field: short YES/continue while a token is outstanding.
|
||||||
|
import re
|
||||||
|
|
||||||
|
compact = re.sub(r"@\S+", " ", low)
|
||||||
|
compact = " ".join(compact.split())
|
||||||
|
if compact in _AFFIRMATIVE_SHORT:
|
||||||
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -30,9 +30,15 @@ _MCP_TOOL_TIMEOUT_OVERRIDES_S: dict[str, float] = {
|
||||||
"aggregateUmeAlarms": 60.0,
|
"aggregateUmeAlarms": 60.0,
|
||||||
}
|
}
|
||||||
|
|
||||||
_LIST_CLI_CACHE_TTL_S = 120.0
|
# Read-mostly inventory/list tools that agents re-call in tight self-loops on WhatsApp.
|
||||||
_LIST_CLI_CACHE_LOCK = threading.Lock()
|
_MCP_LIST_CACHE_TTL_S: dict[str, float] = {
|
||||||
_LIST_CLI_CACHE: dict[str, tuple[float, dict[str, Any]]] = {}
|
"listCliTargets": 120.0,
|
||||||
|
"listManagedNe": 120.0,
|
||||||
|
"queryUmeNeInventory": 90.0,
|
||||||
|
}
|
||||||
|
|
||||||
|
_MCP_LIST_CACHE_LOCK = threading.Lock()
|
||||||
|
_MCP_LIST_CACHE: dict[str, tuple[float, float, dict[str, Any]]] = {}
|
||||||
|
|
||||||
|
|
||||||
def mcp_timeout_for_tool(tool_name: str, row_timeout_s: float | None = None) -> float:
|
def mcp_timeout_for_tool(tool_name: str, row_timeout_s: float | None = None) -> float:
|
||||||
|
|
@ -45,43 +51,51 @@ def mcp_timeout_for_tool(tool_name: str, row_timeout_s: float | None = None) ->
|
||||||
return max(5.0, base)
|
return max(5.0, base)
|
||||||
|
|
||||||
|
|
||||||
def _list_cli_cache_key(server_id: str, args: dict[str, Any]) -> str:
|
def _mcp_list_cache_ttl(tool_name: str) -> float | None:
|
||||||
payload = {
|
return _MCP_LIST_CACHE_TTL_S.get(str(tool_name or "").strip())
|
||||||
"server_id": server_id,
|
|
||||||
"source": str(args.get("source") or "all"),
|
|
||||||
"keyword": str(args.get("keyword") or ""),
|
def _mcp_list_cache_key(server_id: str, tool_name: str, args: dict[str, Any]) -> str:
|
||||||
"page": int(args.get("page") or 1),
|
# Keep key stable; drop obviously volatile noise keys if present.
|
||||||
"page_size": int(args.get("page_size") or 50),
|
cleaned = {
|
||||||
|
str(k): args.get(k)
|
||||||
|
for k in sorted(str(x) for x in (args or {}).keys())
|
||||||
|
if str(k) not in {"trace_id", "request_id", "run_id"}
|
||||||
}
|
}
|
||||||
return json.dumps(payload, sort_keys=True, ensure_ascii=False)
|
payload = {"server_id": server_id, "tool_name": tool_name, "args": cleaned}
|
||||||
|
return json.dumps(payload, sort_keys=True, ensure_ascii=False, default=str)
|
||||||
|
|
||||||
|
|
||||||
def _get_list_cli_cache(key: str) -> dict[str, Any] | None:
|
def _get_mcp_list_cache(key: str) -> dict[str, Any] | None:
|
||||||
now = time.monotonic()
|
now = time.monotonic()
|
||||||
with _LIST_CLI_CACHE_LOCK:
|
with _MCP_LIST_CACHE_LOCK:
|
||||||
hit = _LIST_CLI_CACHE.get(key)
|
hit = _MCP_LIST_CACHE.get(key)
|
||||||
if not hit:
|
if not hit:
|
||||||
return None
|
return None
|
||||||
ts, payload = hit
|
ts, ttl, payload = hit
|
||||||
if now - ts > _LIST_CLI_CACHE_TTL_S:
|
if now - ts > float(ttl):
|
||||||
_LIST_CLI_CACHE.pop(key, None)
|
_MCP_LIST_CACHE.pop(key, None)
|
||||||
return None
|
return None
|
||||||
return dict(payload)
|
return dict(payload)
|
||||||
|
|
||||||
|
|
||||||
def _set_list_cli_cache(key: str, payload: dict[str, Any]) -> None:
|
def _set_mcp_list_cache(key: str, payload: dict[str, Any], *, ttl_s: float) -> None:
|
||||||
with _LIST_CLI_CACHE_LOCK:
|
with _MCP_LIST_CACHE_LOCK:
|
||||||
# Bound memory: drop oldest when large.
|
if len(_MCP_LIST_CACHE) >= 96:
|
||||||
if len(_LIST_CLI_CACHE) >= 64:
|
oldest = sorted(_MCP_LIST_CACHE.items(), key=lambda kv: kv[1][0])[:24]
|
||||||
oldest = sorted(_LIST_CLI_CACHE.items(), key=lambda kv: kv[1][0])[:16]
|
|
||||||
for k, _ in oldest:
|
for k, _ in oldest:
|
||||||
_LIST_CLI_CACHE.pop(k, None)
|
_MCP_LIST_CACHE.pop(k, None)
|
||||||
_LIST_CLI_CACHE[key] = (time.monotonic(), dict(payload))
|
_MCP_LIST_CACHE[key] = (time.monotonic(), float(ttl_s), dict(payload))
|
||||||
|
|
||||||
|
|
||||||
def clear_list_cli_targets_cache() -> None:
|
def clear_list_cli_targets_cache() -> None:
|
||||||
with _LIST_CLI_CACHE_LOCK:
|
"""Clear inventory/list MCP caches (name kept for test compatibility)."""
|
||||||
_LIST_CLI_CACHE.clear()
|
with _MCP_LIST_CACHE_LOCK:
|
||||||
|
_MCP_LIST_CACHE.clear()
|
||||||
|
|
||||||
|
|
||||||
|
def clear_mcp_list_cache() -> None:
|
||||||
|
clear_list_cli_targets_cache()
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|
@ -108,17 +122,18 @@ class _McpBoundTool:
|
||||||
|
|
||||||
def _handler(args: dict[str, Any]) -> dict[str, Any]:
|
def _handler(args: dict[str, Any]) -> dict[str, Any]:
|
||||||
call_args = dict(args or {})
|
call_args = dict(args or {})
|
||||||
|
cache_ttl = _mcp_list_cache_ttl(tool_name)
|
||||||
cache_key = ""
|
cache_key = ""
|
||||||
if tool_name == "listCliTargets":
|
if cache_ttl is not None:
|
||||||
cache_key = _list_cli_cache_key(server_id, call_args)
|
cache_key = _mcp_list_cache_key(server_id, tool_name, call_args)
|
||||||
cached = _get_list_cli_cache(cache_key)
|
cached = _get_mcp_list_cache(cache_key)
|
||||||
if cached is not None:
|
if cached is not None:
|
||||||
out = dict(cached)
|
out = dict(cached)
|
||||||
out["cache_hit"] = True
|
out["cache_hit"] = True
|
||||||
out["cache_ttl_s"] = _LIST_CLI_CACHE_TTL_S
|
out["cache_ttl_s"] = float(cache_ttl)
|
||||||
out["hint"] = (
|
out["hint"] = (
|
||||||
out.get("hint")
|
out.get("hint")
|
||||||
or "Reused listCliTargets result from short TTL cache; do not re-list before every execManagedNe."
|
or f"Reused {tool_name} result from short TTL cache; do not re-list before every follow-up tool."
|
||||||
)
|
)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
@ -130,13 +145,12 @@ class _McpBoundTool:
|
||||||
from runtime.tools.tool_error_hints import enrich_mcp_scope_error
|
from runtime.tools.tool_error_hints import enrich_mcp_scope_error
|
||||||
|
|
||||||
res = enrich_mcp_scope_error(res)
|
res = enrich_mcp_scope_error(res)
|
||||||
if tool_name == "listCliTargets" and res.get("ok") is not False and cache_key:
|
if cache_ttl is not None and cache_key and res.get("ok") is not False:
|
||||||
_set_list_cli_cache(cache_key, res)
|
_set_mcp_list_cache(cache_key, res, ttl_s=float(cache_ttl))
|
||||||
res = dict(res)
|
res = dict(res)
|
||||||
res["cache_hit"] = False
|
res["cache_hit"] = False
|
||||||
res["hint"] = (
|
res["hint"] = (
|
||||||
"Cache listCliTargets ids for this session; call execManagedNe with ne_id/ume_ne_id "
|
f"Cache {tool_name} results briefly; reuse ids/rows instead of listing again in the same turn."
|
||||||
"instead of listing again."
|
|
||||||
)
|
)
|
||||||
if tool_name == "execManagedNe" and res.get("ok") is False:
|
if tool_name == "execManagedNe" and res.get("ok") is False:
|
||||||
from runtime.tools.tool_error_hints import enrich_exec_managed_ne_error
|
from runtime.tools.tool_error_hints import enrich_exec_managed_ne_error
|
||||||
|
|
@ -296,6 +310,7 @@ def materialize_mcp_skills_for_specialist(
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"clear_list_cli_targets_cache",
|
"clear_list_cli_targets_cache",
|
||||||
|
"clear_mcp_list_cache",
|
||||||
"materialize_mcp_tools",
|
"materialize_mcp_tools",
|
||||||
"materialize_mcp_tools_for_specialist",
|
"materialize_mcp_tools_for_specialist",
|
||||||
"materialize_mcp_skills_for_specialist",
|
"materialize_mcp_skills_for_specialist",
|
||||||
|
|
|
||||||
13
tests/test_confirm_affirmative.py
Normal file
13
tests/test_confirm_affirmative.py
Normal file
|
|
@ -0,0 +1,13 @@
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from runtime.orchestration.security import has_explicit_confirmation_token
|
||||||
|
|
||||||
|
|
||||||
|
def test_whatsapp_yes_confirms_outstanding_token() -> None:
|
||||||
|
assert has_explicit_confirmation_token("@bot YES", "abc123")
|
||||||
|
assert has_explicit_confirmation_token("YES", "abc123")
|
||||||
|
assert has_explicit_confirmation_token("please continue", "abc123")
|
||||||
|
assert has_explicit_confirmation_token("确认", "abc123")
|
||||||
|
assert not has_explicit_confirmation_token("YES", None)
|
||||||
|
assert not has_explicit_confirmation_token("maybe later", "abc123")
|
||||||
|
assert has_explicit_confirmation_token("confirm abc123", "abc123")
|
||||||
|
|
@ -70,22 +70,41 @@ class McpTimeoutAndCacheTests(unittest.TestCase):
|
||||||
"tool_name": "listCliTargets",
|
"tool_name": "listCliTargets",
|
||||||
"description": "list",
|
"description": "list",
|
||||||
"parameters": {"type": "object", "properties": {}},
|
"parameters": {"type": "object", "properties": {}},
|
||||||
}
|
},
|
||||||
|
{
|
||||||
|
"tool_name": "queryUmeNeInventory",
|
||||||
|
"description": "inventory",
|
||||||
|
"parameters": {"type": "object", "properties": {}},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"tool_name": "listManagedNe",
|
||||||
|
"description": "managed",
|
||||||
|
"parameters": {"type": "object", "properties": {}},
|
||||||
|
},
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
spec = next(s for s in materialize_mcp_tools(store) if s.name.endswith("listCliTargets"))
|
specs = {s.name: s for s in materialize_mcp_tools(store)}
|
||||||
calls = {"n": 0}
|
calls = {"n": 0}
|
||||||
|
|
||||||
def fake_call_tool(self, tool_name, arguments=None): # type: ignore[no-untyped-def]
|
def fake_call_tool(self, tool_name, arguments=None): # type: ignore[no-untyped-def]
|
||||||
calls["n"] += 1
|
calls["n"] += 1
|
||||||
return {"ok": True, "data": {"items": [{"ne_id": "1"}]}}
|
return {"ok": True, "data": {"items": [{"ne_id": "1"}], "tool": tool_name}}
|
||||||
|
|
||||||
with patch("runtime.tools.mcp.adapter.McpProcessRuntime.call_tool", fake_call_tool):
|
with patch("runtime.tools.mcp.adapter.McpProcessRuntime.call_tool", fake_call_tool):
|
||||||
first = spec.handler({"keyword": "PE", "source": "ume"})
|
cli = specs["mcp__netx__listCliTargets"]
|
||||||
second = spec.handler({"keyword": "PE", "source": "ume"})
|
inv = specs["mcp__netx__queryUmeNeInventory"]
|
||||||
self.assertEqual(calls["n"], 1)
|
managed = specs["mcp__netx__listManagedNe"]
|
||||||
|
first = cli.handler({"keyword": "PE", "source": "ume"})
|
||||||
|
second = cli.handler({"keyword": "PE", "source": "ume"})
|
||||||
|
inv1 = inv.handler({"keyword": "core"})
|
||||||
|
inv2 = inv.handler({"keyword": "core"})
|
||||||
|
m1 = managed.handler({"keyword": "x", "vendor": "huawei", "connect_status": "online"})
|
||||||
|
m2 = managed.handler({"keyword": "x", "vendor": "huawei", "connect_status": "online"})
|
||||||
|
self.assertEqual(calls["n"], 3)
|
||||||
self.assertFalse(first.get("cache_hit"))
|
self.assertFalse(first.get("cache_hit"))
|
||||||
self.assertTrue(second.get("cache_hit"))
|
self.assertTrue(second.get("cache_hit"))
|
||||||
|
self.assertTrue(inv2.get("cache_hit"))
|
||||||
|
self.assertTrue(m2.get("cache_hit"))
|
||||||
self.assertEqual(second.get("data", {}).get("items", [])[0]["ne_id"], "1")
|
self.assertEqual(second.get("data", {}).get("items", [])[0]["ne_id"], "1")
|
||||||
clear_list_cli_targets_cache()
|
clear_list_cli_targets_cache()
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue