mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-08 23:33:16 +08:00
Relax absolute tool-count caps; keep only anti-loop guards.
Field turns may legitimately use many distinct CLI/shell calls (e.g. 200 rounds). Drop per-turn call/fail budgets and default short-intent round caps; identical-arg and retry_forbidden blocks remain. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
9d447b1c51
commit
021ed0cc4f
4 changed files with 25 additions and 266 deletions
|
|
@ -416,67 +416,6 @@ 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
|
||||
_SHELL_TURN_CALL_BUDGET = 5
|
||||
_SHELL_TURN_FAIL_BUDGET = 3
|
||||
|
||||
|
||||
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 _turn_name_budgets(name: str) -> tuple[int, int, str] | None:
|
||||
"""Return (call_budget, fail_budget, kind) for tools that spam WhatsApp turns."""
|
||||
low = str(name or "").strip().lower()
|
||||
if any(low.endswith(suf) for suf in _HEAVY_CLI_NAME_SUFFIXES):
|
||||
return (_HEAVY_CLI_TURN_CALL_BUDGET, _HEAVY_CLI_TURN_FAIL_BUDGET, "cli")
|
||||
if low == "run_command" or low.endswith("__run_command") or low.endswith(".run_command"):
|
||||
return (_SHELL_TURN_CALL_BUDGET, _SHELL_TURN_FAIL_BUDGET, "shell")
|
||||
return None
|
||||
|
||||
|
||||
def normalize_tool_result(result: Any) -> dict[str, Any]:
|
||||
if isinstance(result, dict):
|
||||
out = dict(result)
|
||||
|
|
@ -1104,12 +1043,6 @@ 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] = []
|
||||
|
|
@ -1140,70 +1073,6 @@ class ToolExecutor:
|
|||
{"tool_name": tc.name},
|
||||
)
|
||||
continue
|
||||
budgets = _turn_name_budgets(tool_name)
|
||||
if budgets is not None:
|
||||
call_budget, fail_budget, kind = budgets
|
||||
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 >= fail_budget:
|
||||
if kind == "cli":
|
||||
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."
|
||||
)
|
||||
code = "cli_fail_budget_exceeded"
|
||||
else:
|
||||
hint = (
|
||||
"run_command already failed multiple times this turn. "
|
||||
"Stop shell retries; use write_xlsx / ume_alarm_xlsx_report / dedicated tools instead."
|
||||
)
|
||||
code = "shell_fail_budget_exceeded"
|
||||
results_by_id[tc.id] = (
|
||||
{
|
||||
"ok": False,
|
||||
"error_code": code,
|
||||
"failure_class": "retry_guard",
|
||||
"error": f"{tool_name} failed too many times this turn",
|
||||
"hint": hint,
|
||||
"fail_count": prior_fails,
|
||||
"fail_budget": fail_budget,
|
||||
},
|
||||
0,
|
||||
)
|
||||
_trace(code, {"tool_name": tool_name, "fail_count": prior_fails})
|
||||
continue
|
||||
if prior_calls >= call_budget:
|
||||
if kind == "cli":
|
||||
hint = (
|
||||
"Too many execManagedNe calls this turn. "
|
||||
"Batch commands into fewer calls, reuse prior listCliTargets ids, "
|
||||
"and summarize what you already have."
|
||||
)
|
||||
code = "cli_call_budget_exceeded"
|
||||
else:
|
||||
hint = (
|
||||
"Too many run_command calls this turn. "
|
||||
"Prefer dedicated tools (write_xlsx, ume_alarm_xlsx_report, MCP) "
|
||||
"and summarize without more shell."
|
||||
)
|
||||
code = "shell_call_budget_exceeded"
|
||||
results_by_id[tc.id] = (
|
||||
{
|
||||
"ok": False,
|
||||
"error_code": code,
|
||||
"failure_class": "retry_guard",
|
||||
"error": f"{tool_name} call budget exceeded this turn",
|
||||
"hint": hint,
|
||||
"call_count": prior_calls,
|
||||
"call_budget": call_budget,
|
||||
},
|
||||
0,
|
||||
)
|
||||
_trace(code, {"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] = (
|
||||
{
|
||||
|
|
@ -1358,7 +1227,6 @@ 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()
|
||||
|
|
|
|||
|
|
@ -485,22 +485,28 @@ class OclawGateway:
|
|||
return str(maybe_ops_short_intent_system_hint(text=str(msg.text or ""), lang=lang) or "").strip()
|
||||
|
||||
def _resolve_max_tool_rounds(self, msg: StandardMessage, *, base: int) -> int:
|
||||
"""Cap tool rounds for WhatsApp/WeChat ops short intents (cut 12+ tool loops)."""
|
||||
"""Optionally cap short-intent rounds when explicitly configured.
|
||||
|
||||
Field ops often raise AIA_TURN_MAX_TOOL_ROUNDS (e.g. 200). Absolute caps hurt
|
||||
legitimate multi-NE work; identical-arg / retry_forbidden guards handle loops.
|
||||
Set AIA_OPS_SHORT_INTENT_MAX_TOOL_ROUNDS only if you want an explicit short-intent ceiling.
|
||||
"""
|
||||
rounds = max(1, min(int(base), 300))
|
||||
if not self._is_channel_delivery_channel(msg):
|
||||
return rounds
|
||||
try:
|
||||
raw = str(self.store.get_setting("AIA_OPS_SHORT_INTENT_MAX_TOOL_ROUNDS") or "").strip()
|
||||
except Exception:
|
||||
raw = ""
|
||||
if not raw.isdigit():
|
||||
return rounds
|
||||
md = msg.metadata if isinstance(msg.metadata, dict) else {}
|
||||
from runtime.application.gateway.ops_short_intent import detect_ops_short_intent
|
||||
|
||||
intent = detect_ops_short_intent(str(msg.text or md.get("raw_inbound_text") or ""))
|
||||
if not intent:
|
||||
return rounds
|
||||
try:
|
||||
raw = str(self.store.get_setting("AIA_OPS_SHORT_INTENT_MAX_TOOL_ROUNDS") or "").strip()
|
||||
cap = int(raw) if raw.isdigit() else 8
|
||||
except Exception:
|
||||
cap = 8
|
||||
cap = max(3, min(int(cap), 30))
|
||||
cap = max(3, min(int(raw), 300))
|
||||
return min(rounds, cap)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -130,7 +130,8 @@ def test_retry_forbidden_blocks_same_tool_different_args(tmp_path: Path) -> None
|
|||
assert blocked.get("retry_forbidden") is True
|
||||
|
||||
|
||||
def test_exec_managed_ne_call_budget(tmp_path: Path) -> None:
|
||||
def test_distinct_exec_managed_ne_calls_not_capped_by_count(tmp_path: Path) -> None:
|
||||
"""Field ops may legitimately CLI many NEs; only identical-arg loops are blocked."""
|
||||
store = SqliteStore(str(tmp_path / "cli.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
calls = {"n": 0}
|
||||
|
|
@ -154,11 +155,10 @@ def test_exec_managed_ne_call_budget(tmp_path: Path) -> None:
|
|||
store=store,
|
||||
tools=reg,
|
||||
session_id=sess.id,
|
||||
turn_uuid="turn-cli-budget",
|
||||
turn_uuid="turn-cli-many",
|
||||
lang="en",
|
||||
)
|
||||
# 4 distinct-arg calls allowed; 5th blocked by call budget.
|
||||
for i in range(4):
|
||||
for i in range(8):
|
||||
ToolExecutor().execute_tool_uses(
|
||||
ctx=ctx,
|
||||
assistant_msg_id=i + 1,
|
||||
|
|
@ -171,123 +171,4 @@ def test_exec_managed_ne_call_budget(tmp_path: Path) -> None:
|
|||
],
|
||||
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"
|
||||
|
||||
|
||||
def test_run_command_call_budget(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "shell.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
calls = {"n": 0}
|
||||
|
||||
def _handler(_args):
|
||||
calls["n"] += 1
|
||||
return {"ok": True, "stdout": "ok", "exit_code": 0}
|
||||
|
||||
reg = ToolRegistry(
|
||||
[
|
||||
ToolSpec(
|
||||
name="run_command",
|
||||
description="shell",
|
||||
parameters={"type": "object", "properties": {"command": {"type": "string"}}},
|
||||
handler=_handler,
|
||||
read_only=False,
|
||||
)
|
||||
]
|
||||
)
|
||||
ctx = ToolExecutionContext(
|
||||
store=store,
|
||||
tools=reg,
|
||||
session_id=sess.id,
|
||||
turn_uuid="turn-shell-budget",
|
||||
lang="en",
|
||||
)
|
||||
for i in range(5):
|
||||
ToolExecutor().execute_tool_uses(
|
||||
ctx=ctx,
|
||||
assistant_msg_id=i + 1,
|
||||
tool_uses=[LLMToolCall(id=f"s{i}", name="run_command", arguments={"command": f"echo {i}"})],
|
||||
signature_budget=2,
|
||||
)
|
||||
assert calls["n"] == 5
|
||||
_, results = ToolExecutor().execute_tool_uses(
|
||||
ctx=ctx,
|
||||
assistant_msg_id=99,
|
||||
tool_uses=[LLMToolCall(id="sX", name="run_command", arguments={"command": "echo x"})],
|
||||
signature_budget=2,
|
||||
)
|
||||
blocked, _ = results["sX"]
|
||||
assert calls["n"] == 5
|
||||
assert blocked.get("error_code") == "shell_call_budget_exceeded"
|
||||
assert calls["n"] == 8
|
||||
|
|
|
|||
|
|
@ -1031,6 +1031,10 @@ def test_ops_short_intent_caps_tool_rounds(tmp_path) -> None:
|
|||
attachments=[],
|
||||
metadata={},
|
||||
)
|
||||
assert gw._resolve_max_tool_rounds(short, base=100) == 8
|
||||
assert gw._resolve_max_tool_rounds(long, base=100) == 100
|
||||
assert gw._resolve_max_tool_rounds(admin, base=100) == 100
|
||||
assert gw._resolve_max_tool_rounds(short, base=200) == 200
|
||||
assert gw._resolve_max_tool_rounds(long, base=200) == 200
|
||||
assert gw._resolve_max_tool_rounds(admin, base=200) == 200
|
||||
|
||||
store.set_setting("AIA_OPS_SHORT_INTENT_MAX_TOOL_ROUNDS", "12")
|
||||
assert gw._resolve_max_tool_rounds(short, base=200) == 12
|
||||
assert gw._resolve_max_tool_rounds(long, base=200) == 200
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue