mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-09 00:40:45 +08:00
Dedupe WhatsApp access pending and cap execManagedNe per turn.
Reuse open pending requests without re-notifying admins, clarify already-waiting replies, and block further CLI after 4 calls or 2 failures in the same turn. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
d8bc08ee0d
commit
28a5f2586c
6 changed files with 332 additions and 27 deletions
|
|
@ -461,6 +461,20 @@ def handle_whatsapp_access(
|
|||
)
|
||||
return None
|
||||
|
||||
pending_id = ""
|
||||
already_pending = False
|
||||
finder = getattr(store, "find_open_whatsapp_access_pending", None)
|
||||
if callable(finder):
|
||||
existing = finder(
|
||||
tenant_id=tenant_id,
|
||||
account_id=account_id,
|
||||
external_user_id=raw_jid,
|
||||
phone=phone,
|
||||
)
|
||||
if isinstance(existing, dict) and str(existing.get("id") or "").strip():
|
||||
pending_id = str(existing.get("id") or "").strip()
|
||||
already_pending = True
|
||||
if not pending_id:
|
||||
pending_id = store.create_whatsapp_access_pending(
|
||||
tenant_id=tenant_id,
|
||||
account_id=account_id,
|
||||
|
|
@ -468,8 +482,8 @@ def handle_whatsapp_access(
|
|||
push_name=push_name,
|
||||
phone=phone,
|
||||
request_text=text,
|
||||
)
|
||||
if pending_id:
|
||||
) or ""
|
||||
if pending_id and not already_pending:
|
||||
_notify_admins(
|
||||
store,
|
||||
tenant_id=tenant_id,
|
||||
|
|
@ -500,12 +514,17 @@ def handle_whatsapp_access(
|
|||
{
|
||||
"channel": "whatsapp",
|
||||
"chat_id": inbound.external_chat_id,
|
||||
"text": denied_reply_text(lang=lang, pending_id=str(pending_id or "")),
|
||||
"text": denied_reply_text(
|
||||
lang=lang,
|
||||
pending_id=str(pending_id or ""),
|
||||
already_pending=already_pending,
|
||||
),
|
||||
"attachments": [],
|
||||
"metadata": reply_meta,
|
||||
}
|
||||
],
|
||||
"whatsapp_access": "denied",
|
||||
"whatsapp_access_already_pending": bool(already_pending),
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -416,6 +416,55 @@ def _load_turn_retry_forbidden_tools(store: Any, *, session_id: str, turn_uuid:
|
|||
return out
|
||||
|
||||
|
||||
def _load_turn_failed_name_counts(store: Any, *, session_id: str, turn_uuid: str) -> dict[str, int]:
|
||||
"""Count failed tool results by tool name earlier in this turn."""
|
||||
tu = str(turn_uuid or "").strip()
|
||||
sid = str(session_id or "").strip()
|
||||
out: dict[str, int] = {}
|
||||
if not tu or not sid:
|
||||
return out
|
||||
try:
|
||||
rows = store.get_messages(session_id=sid, limit=500)
|
||||
except Exception:
|
||||
return out
|
||||
for m in rows or []:
|
||||
if str(getattr(m, "role", "") or "").strip().lower() != "tool":
|
||||
continue
|
||||
if str(getattr(m, "turn_uuid", "") or "").strip() != tu:
|
||||
continue
|
||||
ep = _parse_event_payload(getattr(m, "event_payload", None))
|
||||
name = str(ep.get("tool_name") or "").strip()
|
||||
failed = ep.get("ok") is False
|
||||
if not failed:
|
||||
try:
|
||||
payload = json.loads(str(getattr(m, "content", "") or "") or "{}")
|
||||
except Exception:
|
||||
payload = {}
|
||||
failed = isinstance(payload, dict) and payload.get("ok") is False
|
||||
if not name and isinstance(payload, dict):
|
||||
raw_tc = getattr(m, "tool_calls", None)
|
||||
if isinstance(raw_tc, str):
|
||||
try:
|
||||
raw_tc = json.loads(raw_tc)
|
||||
except Exception:
|
||||
raw_tc = None
|
||||
if isinstance(raw_tc, dict):
|
||||
name = str(raw_tc.get("name") or "").strip()
|
||||
if failed and name:
|
||||
out[name] = int(out.get(name, 0)) + 1
|
||||
return out
|
||||
|
||||
|
||||
_HEAVY_CLI_NAME_SUFFIXES = ("execmanagedne",)
|
||||
_HEAVY_CLI_TURN_CALL_BUDGET = 4
|
||||
_HEAVY_CLI_TURN_FAIL_BUDGET = 2
|
||||
|
||||
|
||||
def _is_heavy_cli_tool(name: str) -> bool:
|
||||
low = str(name or "").strip().lower()
|
||||
return any(low.endswith(suf) for suf in _HEAVY_CLI_NAME_SUFFIXES)
|
||||
|
||||
|
||||
def normalize_tool_result(result: Any) -> dict[str, Any]:
|
||||
if isinstance(result, dict):
|
||||
out = dict(result)
|
||||
|
|
@ -1043,6 +1092,12 @@ class ToolExecutor:
|
|||
session_id=ctx.session_id,
|
||||
turn_uuid=str(ctx.turn_uuid or ""),
|
||||
)
|
||||
failed_name_counts = _load_turn_failed_name_counts(
|
||||
ctx.store,
|
||||
session_id=ctx.session_id,
|
||||
turn_uuid=str(ctx.turn_uuid or ""),
|
||||
)
|
||||
planned_name_counts: dict[str, int] = {}
|
||||
|
||||
results_by_id: dict[str, tuple[dict[str, Any], int]] = {}
|
||||
runnable_tool_uses: list[LLMToolCall] = []
|
||||
|
|
@ -1051,7 +1106,8 @@ class ToolExecutor:
|
|||
sig_seen: dict[str, int] = {}
|
||||
budget = max(1, min(int(signature_budget or 2), 8))
|
||||
for tc in tool_uses:
|
||||
if str(tc.name or "") in retry_forbidden_tools:
|
||||
tool_name = str(tc.name or "")
|
||||
if tool_name in retry_forbidden_tools:
|
||||
results_by_id[tc.id] = (
|
||||
{
|
||||
"ok": False,
|
||||
|
|
@ -1072,6 +1128,55 @@ class ToolExecutor:
|
|||
{"tool_name": tc.name},
|
||||
)
|
||||
continue
|
||||
if _is_heavy_cli_tool(tool_name):
|
||||
prior_calls = int(turn_tool_name_counts.get(tool_name, 0)) + int(
|
||||
planned_name_counts.get(tool_name, 0)
|
||||
)
|
||||
prior_fails = int(failed_name_counts.get(tool_name, 0))
|
||||
if prior_fails >= _HEAVY_CLI_TURN_FAIL_BUDGET:
|
||||
results_by_id[tc.id] = (
|
||||
{
|
||||
"ok": False,
|
||||
"error_code": "cli_fail_budget_exceeded",
|
||||
"failure_class": "retry_guard",
|
||||
"error": f"{tool_name} failed too many times this turn",
|
||||
"hint": (
|
||||
"execManagedNe already failed multiple times this turn. "
|
||||
"Stop CLI spam: report unreachable/timeout NEs, batch remaining "
|
||||
"show commands into one call with higher read_timeout_sec, or switch targets."
|
||||
),
|
||||
"fail_count": prior_fails,
|
||||
"fail_budget": _HEAVY_CLI_TURN_FAIL_BUDGET,
|
||||
},
|
||||
0,
|
||||
)
|
||||
_trace(
|
||||
"cli_fail_budget_exceeded",
|
||||
{"tool_name": tool_name, "fail_count": prior_fails},
|
||||
)
|
||||
continue
|
||||
if prior_calls >= _HEAVY_CLI_TURN_CALL_BUDGET:
|
||||
results_by_id[tc.id] = (
|
||||
{
|
||||
"ok": False,
|
||||
"error_code": "cli_call_budget_exceeded",
|
||||
"failure_class": "retry_guard",
|
||||
"error": f"{tool_name} call budget exceeded this turn",
|
||||
"hint": (
|
||||
"Too many execManagedNe calls this turn. "
|
||||
"Batch commands into fewer calls, reuse prior listCliTargets ids, "
|
||||
"and summarize what you already have."
|
||||
),
|
||||
"call_count": prior_calls,
|
||||
"call_budget": _HEAVY_CLI_TURN_CALL_BUDGET,
|
||||
},
|
||||
0,
|
||||
)
|
||||
_trace(
|
||||
"cli_call_budget_exceeded",
|
||||
{"tool_name": tool_name, "call_count": prior_calls},
|
||||
)
|
||||
continue
|
||||
if tc.name in _TABULAR_QUERY_TOOL_NAMES and not has_tabular_ref_in_session:
|
||||
results_by_id[tc.id] = (
|
||||
{
|
||||
|
|
@ -1226,6 +1331,7 @@ class ToolExecutor:
|
|||
continue
|
||||
first_tool_call_id_by_signature[sig] = str(tc.id or "")
|
||||
runnable_tool_uses.append(tc)
|
||||
planned_name_counts[tool_name] = int(planned_name_counts.get(tool_name, 0)) + 1
|
||||
|
||||
for batch in partition_tool_use_batches(runnable_tool_uses, ctx.tools):
|
||||
_check_stop()
|
||||
|
|
|
|||
|
|
@ -343,15 +343,25 @@ def coerce_whatsapp_access_target(value: str) -> str:
|
|||
return normalize_whatsapp_target(normalize_whatsapp_phone(value))
|
||||
|
||||
|
||||
def denied_reply_text(*, lang: str, pending_id: str = "") -> str:
|
||||
def denied_reply_text(*, lang: str, pending_id: str = "", already_pending: bool = False) -> str:
|
||||
pid = str(pending_id or "").strip()
|
||||
if str(lang or "").strip().lower().startswith("zh"):
|
||||
if already_pending and pid:
|
||||
return (
|
||||
f"访问申请仍在等待管理员处理(编号 {pid})。"
|
||||
"同意后即可使用;请稍候,无需重复发送。"
|
||||
)
|
||||
if pid:
|
||||
return (
|
||||
f"访问申请已提交(编号 {pid})。管理员同意后即可使用;"
|
||||
"请稍候,无需重复发送相同请求。"
|
||||
)
|
||||
return "无权限:您尚未获得使用此助手的授权。请联系管理员。"
|
||||
if already_pending and pid:
|
||||
return (
|
||||
f"Your access request is still pending (request {pid}). "
|
||||
"An administrator was already notified — please wait for YES; no need to resend."
|
||||
)
|
||||
if pid:
|
||||
return (
|
||||
f"Access pending (request {pid}): an administrator was notified. "
|
||||
|
|
|
|||
|
|
@ -6042,6 +6042,31 @@ class SqliteStore(ScheduledJobStoreMixin):
|
|||
if "last_attempt_at" not in cols:
|
||||
conn.execute("ALTER TABLE channel_outbound_message ADD COLUMN last_attempt_at TEXT")
|
||||
|
||||
def find_open_whatsapp_access_pending(
|
||||
self,
|
||||
*,
|
||||
tenant_id: str,
|
||||
account_id: str,
|
||||
external_user_id: str,
|
||||
phone: str = "",
|
||||
) -> dict[str, Any] | None:
|
||||
"""Return an existing pending access row for this sender, if any."""
|
||||
from runtime.extensions.whatsapp.access_control import contact_phone_key, whatsapp_users_match
|
||||
|
||||
pending = self.list_whatsapp_access_pending(
|
||||
tenant_id=tenant_id,
|
||||
account_id=account_id,
|
||||
status="pending",
|
||||
limit=200,
|
||||
)
|
||||
phone_val = str(phone or "").strip()
|
||||
for row in pending:
|
||||
if whatsapp_users_match(str(row.get("external_user_id") or ""), external_user_id):
|
||||
return dict(row)
|
||||
if phone_val and contact_phone_key(row) == phone_val:
|
||||
return dict(row)
|
||||
return None
|
||||
|
||||
def create_whatsapp_access_pending(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -6054,20 +6079,14 @@ class SqliteStore(ScheduledJobStoreMixin):
|
|||
) -> str | None:
|
||||
import uuid
|
||||
|
||||
from runtime.extensions.whatsapp.access_control import contact_phone_key, whatsapp_users_match
|
||||
|
||||
pending = self.list_whatsapp_access_pending(
|
||||
existing = self.find_open_whatsapp_access_pending(
|
||||
tenant_id=tenant_id,
|
||||
account_id=account_id,
|
||||
status="pending",
|
||||
limit=200,
|
||||
external_user_id=external_user_id,
|
||||
phone=phone,
|
||||
)
|
||||
phone_val = str(phone or "").strip()
|
||||
for row in pending:
|
||||
if whatsapp_users_match(str(row.get("external_user_id") or ""), external_user_id):
|
||||
return str(row.get("id") or "")
|
||||
if phone_val and contact_phone_key(row) == phone_val:
|
||||
return str(row.get("id") or "")
|
||||
if existing:
|
||||
return str(existing.get("id") or "") or None
|
||||
pending_id = uuid.uuid4().hex
|
||||
ts = utc_now_iso()
|
||||
with self._connect() as conn:
|
||||
|
|
|
|||
|
|
@ -128,3 +128,120 @@ def test_retry_forbidden_blocks_same_tool_different_args(tmp_path: Path) -> None
|
|||
assert calls["n"] == 1
|
||||
assert blocked.get("error_code") == "retry_forbidden_blocked"
|
||||
assert blocked.get("retry_forbidden") is True
|
||||
|
||||
|
||||
def test_exec_managed_ne_call_budget(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "cli.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
calls = {"n": 0}
|
||||
|
||||
def _handler(args):
|
||||
calls["n"] += 1
|
||||
return {"ok": True, "data": {"ne_id": args.get("ne_id"), "n": calls["n"]}}
|
||||
|
||||
reg = ToolRegistry(
|
||||
[
|
||||
ToolSpec(
|
||||
name="mcp__netx__execManagedNe",
|
||||
description="exec",
|
||||
parameters={"type": "object", "properties": {"ne_id": {"type": "string"}}},
|
||||
handler=_handler,
|
||||
read_only=False,
|
||||
)
|
||||
]
|
||||
)
|
||||
ctx = ToolExecutionContext(
|
||||
store=store,
|
||||
tools=reg,
|
||||
session_id=sess.id,
|
||||
turn_uuid="turn-cli-budget",
|
||||
lang="en",
|
||||
)
|
||||
# 4 distinct-arg calls allowed; 5th blocked by call budget.
|
||||
for i in range(4):
|
||||
ToolExecutor().execute_tool_uses(
|
||||
ctx=ctx,
|
||||
assistant_msg_id=i + 1,
|
||||
tool_uses=[
|
||||
LLMToolCall(
|
||||
id=f"c{i}",
|
||||
name="mcp__netx__execManagedNe",
|
||||
arguments={"ne_id": f"ne-{i}", "commands": [f"disp {i}"]},
|
||||
)
|
||||
],
|
||||
signature_budget=2,
|
||||
)
|
||||
assert calls["n"] == 4
|
||||
_, results = ToolExecutor().execute_tool_uses(
|
||||
ctx=ctx,
|
||||
assistant_msg_id=99,
|
||||
tool_uses=[
|
||||
LLMToolCall(
|
||||
id="cX",
|
||||
name="mcp__netx__execManagedNe",
|
||||
arguments={"ne_id": "ne-x", "commands": ["disp x"]},
|
||||
)
|
||||
],
|
||||
signature_budget=2,
|
||||
)
|
||||
blocked, _ = results["cX"]
|
||||
assert calls["n"] == 4
|
||||
assert blocked.get("error_code") == "cli_call_budget_exceeded"
|
||||
|
||||
|
||||
def test_exec_managed_ne_fail_budget(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "cli-fail.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
calls = {"n": 0}
|
||||
|
||||
def _handler(_args):
|
||||
calls["n"] += 1
|
||||
return {"ok": False, "error_code": "tool_timeout_or_failed", "error": "timeout"}
|
||||
|
||||
reg = ToolRegistry(
|
||||
[
|
||||
ToolSpec(
|
||||
name="mcp__netx__execManagedNe",
|
||||
description="exec",
|
||||
parameters={"type": "object", "properties": {"ne_id": {"type": "string"}}},
|
||||
handler=_handler,
|
||||
read_only=False,
|
||||
)
|
||||
]
|
||||
)
|
||||
ctx = ToolExecutionContext(
|
||||
store=store,
|
||||
tools=reg,
|
||||
session_id=sess.id,
|
||||
turn_uuid="turn-cli-fail",
|
||||
lang="en",
|
||||
)
|
||||
for i in range(2):
|
||||
ToolExecutor().execute_tool_uses(
|
||||
ctx=ctx,
|
||||
assistant_msg_id=i + 1,
|
||||
tool_uses=[
|
||||
LLMToolCall(
|
||||
id=f"f{i}",
|
||||
name="mcp__netx__execManagedNe",
|
||||
arguments={"ne_id": f"ne-{i}", "commands": ["disp"]},
|
||||
)
|
||||
],
|
||||
signature_budget=2,
|
||||
)
|
||||
assert calls["n"] == 2
|
||||
_, results = ToolExecutor().execute_tool_uses(
|
||||
ctx=ctx,
|
||||
assistant_msg_id=3,
|
||||
tool_uses=[
|
||||
LLMToolCall(
|
||||
id="f3",
|
||||
name="mcp__netx__execManagedNe",
|
||||
arguments={"ne_id": "ne-3", "commands": ["disp"]},
|
||||
)
|
||||
],
|
||||
signature_budget=2,
|
||||
)
|
||||
blocked, _ = results["f3"]
|
||||
assert calls["n"] == 2
|
||||
assert blocked.get("error_code") == "cli_fail_budget_exceeded"
|
||||
|
|
|
|||
|
|
@ -65,12 +65,19 @@ class _AccessStore:
|
|||
rows = [r for r in rows if str(r.get("list_type") or "") == list_type]
|
||||
return rows
|
||||
|
||||
def find_open_whatsapp_access_pending(self, **kwargs: Any) -> dict[str, Any] | None:
|
||||
jid = str(kwargs.get("external_user_id") or "")
|
||||
for row in self.pending:
|
||||
if row.get("status") == "pending" and row.get("external_user_id") == jid:
|
||||
return dict(row)
|
||||
return None
|
||||
|
||||
def create_whatsapp_access_pending(self, **kwargs: Any) -> str:
|
||||
existing = self.find_open_whatsapp_access_pending(**kwargs)
|
||||
if existing:
|
||||
return str(existing.get("id") or "")
|
||||
self._pending_seq += 1
|
||||
pending_id = f"pending{self._pending_seq}"
|
||||
for row in self.pending:
|
||||
if row.get("status") == "pending" and row.get("external_user_id") == kwargs.get("external_user_id"):
|
||||
return str(row.get("id") or pending_id)
|
||||
self.pending.append(
|
||||
{
|
||||
"id": pending_id,
|
||||
|
|
@ -347,7 +354,34 @@ def test_handle_whatsapp_access_denied_unknown_user(monkeypatch) -> None:
|
|||
assert store.pending[0].get("phone") == "8615601877957"
|
||||
assert store.outbound
|
||||
replies = out.get("replies") if isinstance(out.get("replies"), list) else []
|
||||
assert replies and "denied" in str(replies[0].get("text") or "").lower()
|
||||
text = str((replies[0] or {}).get("text") or "").lower()
|
||||
assert replies and ("pending" in text or "denied" in text)
|
||||
|
||||
|
||||
def test_handle_whatsapp_access_already_pending_skips_renotify(monkeypatch) -> None:
|
||||
import runtime.application.gateway.whatsapp_inbound_access as mod
|
||||
|
||||
monkeypatch.setattr(mod, "resolve_whatsapp_tenant_id", lambda store, account_id: "tenant1")
|
||||
store = _AccessStore()
|
||||
# blacklist mode: unknown senders are denied and create a pending request
|
||||
store.config["access_mode"] = "blacklist"
|
||||
store.contacts["111@s.whatsapp.net"] = {"external_user_id": "111@s.whatsapp.net", "list_type": "admin"}
|
||||
first = handle_whatsapp_access(store, inbound=_Inbound(), account_id="wa-default", text="hello")
|
||||
assert first is not None
|
||||
assert first.get("whatsapp_access") == "denied"
|
||||
assert first.get("whatsapp_access_already_pending") is not True
|
||||
assert len(store.outbound) == 1
|
||||
pending_id = str(store.pending[0].get("id") or "")
|
||||
|
||||
second = handle_whatsapp_access(store, inbound=_Inbound(), account_id="wa-default", text="hello again")
|
||||
assert second is not None
|
||||
assert second.get("whatsapp_access_already_pending") is True
|
||||
assert len(store.outbound) == 1 # no second admin notify
|
||||
assert len(store.pending) == 1
|
||||
assert store.pending[0].get("id") == pending_id
|
||||
replies = second.get("replies") if isinstance(second.get("replies"), list) else []
|
||||
text = str((replies[0] or {}).get("text") or "").lower()
|
||||
assert "still pending" in text or pending_id.lower() in text
|
||||
|
||||
|
||||
def test_handle_whatsapp_access_allows_whitelisted_user(monkeypatch) -> None:
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue