mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-09 01:50:44 +08:00
Classify tool failures and block identical blind retries in-turn.
Stamp failure_class for schema/timeout/runtime analytics, refuse same tool+args after a failure in the turn, and add a license short-intent recipe for WhatsApp ops. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
4646170743
commit
7373e516f0
5 changed files with 212 additions and 2 deletions
|
|
@ -28,6 +28,12 @@ _HINTS: dict[str, tuple[str, str]] = {
|
|||
"do not build xlsx via run_command.]",
|
||||
"[短指令:导出 Excel。优先 ume_alarm_xlsx_report 或 write_xlsx(deliverable=true);禁止 run_command 造表。]",
|
||||
),
|
||||
"license": (
|
||||
"[Ops short-intent: license/capacity. Prefer aggregateUmeAlarms / queryUmeAlarmsRaw with license keywords, "
|
||||
"or ume_alarm_xlsx_report(mode=list, keyword=license); ≤3 tool calls — no CLI spam.]",
|
||||
"[短指令:License/容量。优先 aggregateUmeAlarms / queryUmeAlarmsRaw(license 关键字)"
|
||||
"或 ume_alarm_xlsx_report(mode=list, keyword=license);≤3 次工具,勿刷 CLI。]",
|
||||
),
|
||||
"continue": (
|
||||
"[Ops short-intent: continue/confirm. Resume the unfinished prior task immediately; "
|
||||
"do not re-ask confirmation or restart the query from scratch.]",
|
||||
|
|
@ -58,6 +64,8 @@ def detect_ops_short_intent(text: str) -> str | None:
|
|||
return "offline"
|
||||
if any(k in t for k in ("excel", "xlsx", "spreadsheet", "export", "send me the table", "表格", "导出")):
|
||||
return "excel_export"
|
||||
if any(k in t for k in ("license", "licence", "capacity", "带宽", "拥塞", "license到期")):
|
||||
return "license"
|
||||
if any(
|
||||
k in t
|
||||
for k in (
|
||||
|
|
|
|||
|
|
@ -318,6 +318,51 @@ def _session_has_video_ref(store: Any, session_id: str, *, limit: int = 300) ->
|
|||
return False
|
||||
|
||||
|
||||
def _parse_event_payload(raw: Any) -> dict[str, Any]:
|
||||
if isinstance(raw, dict):
|
||||
return dict(raw)
|
||||
if isinstance(raw, str) and raw.strip():
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
return data if isinstance(data, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
return {}
|
||||
|
||||
|
||||
def _load_turn_failed_signatures(store: Any, *, session_id: str, turn_uuid: str) -> set[str]:
|
||||
"""Signatures that already failed earlier in this turn (blind-retry guard)."""
|
||||
tu = str(turn_uuid or "").strip()
|
||||
sid = str(session_id or "").strip()
|
||||
out: set[str] = set()
|
||||
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))
|
||||
sig = str(ep.get("tool_signature") or "").strip()
|
||||
if not sig:
|
||||
continue
|
||||
if ep.get("ok") is False:
|
||||
out.add(sig)
|
||||
continue
|
||||
# Fallback: inspect tool body when older rows lack ok in event_payload.
|
||||
try:
|
||||
payload = json.loads(str(getattr(m, "content", "") or "") or "{}")
|
||||
except Exception:
|
||||
payload = {}
|
||||
if isinstance(payload, dict) and payload.get("ok") is False:
|
||||
out.add(sig)
|
||||
return out
|
||||
|
||||
|
||||
def normalize_tool_result(result: Any) -> dict[str, Any]:
|
||||
if isinstance(result, dict):
|
||||
out = dict(result)
|
||||
|
|
@ -335,6 +380,9 @@ def normalize_tool_result(result: Any) -> dict[str, Any]:
|
|||
raw_err = str(out.get("error") or "").strip()
|
||||
if not raw_ec:
|
||||
out["error_code"] = _TOOL_ERROR_MAP.get(raw_err, "tool_failed")
|
||||
from runtime.tools.tool_error_hints import stamp_tool_failure_class
|
||||
|
||||
out = stamp_tool_failure_class(out)
|
||||
return out
|
||||
|
||||
|
||||
|
|
@ -881,6 +929,11 @@ class ToolExecutor:
|
|||
has_text_ref_in_session = _session_has_text_ref(ctx.store, ctx.session_id)
|
||||
has_image_ref_in_session = _session_has_image_ref(ctx.store, ctx.session_id)
|
||||
has_video_ref_in_session = _session_has_video_ref(ctx.store, ctx.session_id)
|
||||
failed_signatures = _load_turn_failed_signatures(
|
||||
ctx.store,
|
||||
session_id=ctx.session_id,
|
||||
turn_uuid=str(ctx.turn_uuid or ""),
|
||||
)
|
||||
|
||||
results_by_id: dict[str, tuple[dict[str, Any], int]] = {}
|
||||
runnable_tool_uses: list[LLMToolCall] = []
|
||||
|
|
@ -966,6 +1019,28 @@ class ToolExecutor:
|
|||
)
|
||||
continue
|
||||
sig = f"{tc.name}:{self._json_dumps_safe(dict(tc.arguments or {}))}"
|
||||
if sig in failed_signatures:
|
||||
results_by_id[tc.id] = (
|
||||
{
|
||||
"ok": False,
|
||||
"error_code": "identical_retry_blocked",
|
||||
"failure_class": "retry_guard",
|
||||
"error": f"identical failed call already ran this turn: {tc.name}",
|
||||
"hint": (
|
||||
"Same tool+arguments already failed earlier in this turn. "
|
||||
"Change arguments, raise timeouts, or switch tools — do not blind-retry."
|
||||
),
|
||||
},
|
||||
0,
|
||||
)
|
||||
_trace(
|
||||
"identical_retry_blocked",
|
||||
{
|
||||
"tool_name": tc.name,
|
||||
"signature": sig[:300],
|
||||
},
|
||||
)
|
||||
continue
|
||||
count = int(sig_seen.get(sig, 0))
|
||||
name_low = str(tc.name or "").strip().lower()
|
||||
listish = name_low.endswith(
|
||||
|
|
@ -1059,6 +1134,8 @@ class ToolExecutor:
|
|||
)
|
||||
result, duration_ms = results_by_id[tc.id]
|
||||
result = normalize_tool_result(result)
|
||||
if isinstance(result, dict) and result.get("ok") is False:
|
||||
failed_signatures.add(f"{tc.name}:{self._json_dumps_safe(dict(tc.arguments or {}))}")
|
||||
persisted_result, ingested_refs = ingest_embedded_image_blobs_as_refs(
|
||||
result,
|
||||
filename_prefix=f"{str(tc.name or 'tool')}-{str(tc.id or '')}",
|
||||
|
|
@ -1100,6 +1177,7 @@ class ToolExecutor:
|
|||
trunc_ms = int((time.perf_counter() - t_trunc) * 1000)
|
||||
tool_content = self._json_dumps_safe(result_for_llm)
|
||||
t_db2 = time.perf_counter()
|
||||
tool_sig = f"{tc.name}:{self._json_dumps_safe(dict(tc.arguments or {}))}"
|
||||
msg_row = ctx.store.add_message(
|
||||
session_id=ctx.session_id,
|
||||
role="tool",
|
||||
|
|
@ -1108,7 +1186,13 @@ class ToolExecutor:
|
|||
attachments=(_merge_attachments(_attachments_from_tool_result(persisted_result), ingested_refs) or None),
|
||||
turn_uuid=ctx.turn_uuid,
|
||||
event_type="tool_result",
|
||||
event_payload={"tool_name": tc.name, "observed_rows": int(observed_rows_this_call)},
|
||||
event_payload={
|
||||
"tool_name": tc.name,
|
||||
"observed_rows": int(observed_rows_this_call),
|
||||
"ok": bool(result.get("ok")) if isinstance(result, dict) else False,
|
||||
"tool_signature": tool_sig[:800],
|
||||
"failure_class": str((result or {}).get("failure_class") or "") if isinstance(result, dict) else "",
|
||||
},
|
||||
)
|
||||
try:
|
||||
owner = ctx.store.get_ui_session_owner(session_id=ctx.session_id) or {}
|
||||
|
|
|
|||
|
|
@ -199,9 +199,51 @@ def enrich_exec_managed_ne_error(result: dict[str, Any]) -> dict[str, Any]:
|
|||
return out
|
||||
|
||||
|
||||
def classify_tool_failure(result: dict[str, Any]) -> str:
|
||||
"""Coarse failure class for analytics / retry guards (English field ops)."""
|
||||
if not isinstance(result, dict) or result.get("ok") is not False:
|
||||
return ""
|
||||
existing = str(result.get("failure_class") or result.get("error_class") or "").strip().lower()
|
||||
if existing:
|
||||
return existing
|
||||
code = str(result.get("error_code") or "").strip().lower()
|
||||
err = str(result.get("error") or "").strip().lower()
|
||||
hint = str(result.get("hint") or "").strip().lower()
|
||||
blob = f"{code} {err} {hint}"
|
||||
if code in {"tool_invalid_arguments", "invalid_arguments"} or "invalid arguments" in blob or "参数不合法" in err:
|
||||
return "schema_validation"
|
||||
if "insufficient_scope" in blob or code == "insufficient_scope":
|
||||
return "scope"
|
||||
if "timeout" in blob or code == "tool_timeout_or_failed":
|
||||
return "timeout"
|
||||
if any(x in blob for x in ("unreachable", "connect_failed", "connection refused", "no route")):
|
||||
return "unreachable"
|
||||
if any(x in blob for x in ("auth", "permission denied", "authentication", "login failed")):
|
||||
return "auth"
|
||||
if code in {"tool_loop_guard", "identical_retry_blocked"}:
|
||||
return "retry_guard"
|
||||
if code in {"tool_not_registered"}:
|
||||
return "not_registered"
|
||||
return "runtime"
|
||||
|
||||
|
||||
def stamp_tool_failure_class(result: dict[str, Any]) -> dict[str, Any]:
|
||||
if not isinstance(result, dict) or result.get("ok") is not False:
|
||||
return result if isinstance(result, dict) else {"ok": False, "error": "tool_result_not_dict"}
|
||||
out = dict(result)
|
||||
klass = classify_tool_failure(out)
|
||||
if klass:
|
||||
out["failure_class"] = klass
|
||||
if not str(out.get("error_class") or "").strip():
|
||||
out["error_class"] = klass
|
||||
return out
|
||||
|
||||
|
||||
__all__ = [
|
||||
"classify_tool_failure",
|
||||
"enrich_exec_managed_ne_error",
|
||||
"enrich_mcp_scope_error",
|
||||
"format_unregistered_tool_error",
|
||||
"stamp_tool_failure_class",
|
||||
"suggest_tool_names",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -2300,7 +2300,11 @@ class SqliteStore(ScheduledJobStoreMixin):
|
|||
if raw_cap.isdigit():
|
||||
cap = max(20_000, min(int(raw_cap), 2_000_000))
|
||||
args_capped = self._cap_json_for_log(args, max_chars=cap, keep_keys=())
|
||||
result_capped = self._cap_json_for_log(result, max_chars=cap, keep_keys=("ok", "error_code", "error"))
|
||||
result_capped = self._cap_json_for_log(
|
||||
result,
|
||||
max_chars=cap,
|
||||
keep_keys=("ok", "error_code", "error", "failure_class", "error_class", "hint"),
|
||||
)
|
||||
self._tool_log_queries_repo().insert_tool_log(
|
||||
session_id=str(session_id),
|
||||
tool_name=str(tool_name),
|
||||
|
|
|
|||
72
tests/test_failure_class_and_retry_guard.py
Normal file
72
tests/test_failure_class_and_retry_guard.py
Normal file
|
|
@ -0,0 +1,72 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from runtime.application.gateway.ops_short_intent import detect_ops_short_intent
|
||||
from runtime.chat.tool_runtime import ToolExecutionContext, ToolExecutor, normalize_tool_result
|
||||
from runtime.tools.base import ToolRegistry, ToolSpec
|
||||
from runtime.tools.tool_error_hints import classify_tool_failure
|
||||
from svc.llm.chat_models import LLMToolCall
|
||||
from svc.persistence.sqlite_store import SqliteStore
|
||||
|
||||
|
||||
def test_classify_schema_and_timeout() -> None:
|
||||
assert (
|
||||
classify_tool_failure({"ok": False, "error_code": "tool_invalid_arguments", "error": "bad"})
|
||||
== "schema_validation"
|
||||
)
|
||||
assert (
|
||||
classify_tool_failure({"ok": False, "error_code": "tool_timeout_or_failed", "error": "timeout"})
|
||||
== "timeout"
|
||||
)
|
||||
|
||||
|
||||
def test_normalize_stamps_failure_class() -> None:
|
||||
out = normalize_tool_result({"ok": False, "error_code": "tool_invalid_arguments", "error": "x"})
|
||||
assert out["failure_class"] == "schema_validation"
|
||||
|
||||
|
||||
def test_license_short_intent() -> None:
|
||||
assert detect_ops_short_intent("license expiry report") == "license"
|
||||
assert detect_ops_short_intent("@bot licence check") == "license"
|
||||
|
||||
|
||||
def test_identical_failed_retry_blocked_across_rounds(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "retry.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-retry-1",
|
||||
lang="en",
|
||||
)
|
||||
uses = [LLMToolCall(id="c1", name="mcp__netx__execManagedNe", arguments={"ne_id": "ne-1", "commands": ["disp"]})]
|
||||
ToolExecutor().execute_tool_uses(ctx=ctx, assistant_msg_id=1, tool_uses=uses, signature_budget=2)
|
||||
assert calls["n"] == 1
|
||||
|
||||
uses2 = [LLMToolCall(id="c2", name="mcp__netx__execManagedNe", arguments={"ne_id": "ne-1", "commands": ["disp"]})]
|
||||
_, results = ToolExecutor().execute_tool_uses(
|
||||
ctx=ctx, assistant_msg_id=2, tool_uses=uses2, signature_budget=2
|
||||
)
|
||||
blocked, _ = results["c2"]
|
||||
assert calls["n"] == 1
|
||||
assert blocked.get("error_code") == "identical_retry_blocked"
|
||||
assert blocked.get("failure_class") == "retry_guard"
|
||||
Loading…
Add table
Add a link
Reference in a new issue