From 7373e516f09c4828afcb356517099a81508438e2 Mon Sep 17 00:00:00 2001 From: oliver Date: Mon, 10 Aug 2026 23:03:15 +0800 Subject: [PATCH] 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 --- .../application/gateway/ops_short_intent.py | 8 ++ runtime/chat/tool_runtime.py | 86 ++++++++++++++++++- runtime/tools/tool_error_hints.py | 42 +++++++++ svc/persistence/sqlite_store.py | 6 +- tests/test_failure_class_and_retry_guard.py | 72 ++++++++++++++++ 5 files changed, 212 insertions(+), 2 deletions(-) create mode 100644 tests/test_failure_class_and_retry_guard.py diff --git a/runtime/application/gateway/ops_short_intent.py b/runtime/application/gateway/ops_short_intent.py index 35129ff3..eda810bf 100644 --- a/runtime/application/gateway/ops_short_intent.py +++ b/runtime/application/gateway/ops_short_intent.py @@ -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 ( diff --git a/runtime/chat/tool_runtime.py b/runtime/chat/tool_runtime.py index f0b1d5fc..ed22e39e 100644 --- a/runtime/chat/tool_runtime.py +++ b/runtime/chat/tool_runtime.py @@ -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 {} diff --git a/runtime/tools/tool_error_hints.py b/runtime/tools/tool_error_hints.py index 4ed442e2..d8600742 100644 --- a/runtime/tools/tool_error_hints.py +++ b/runtime/tools/tool_error_hints.py @@ -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", ] diff --git a/svc/persistence/sqlite_store.py b/svc/persistence/sqlite_store.py index 67e466f3..65bf8baf 100644 --- a/svc/persistence/sqlite_store.py +++ b/svc/persistence/sqlite_store.py @@ -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), diff --git a/tests/test_failure_class_and_retry_guard.py b/tests/test_failure_class_and_retry_guard.py new file mode 100644 index 00000000..bf97cf48 --- /dev/null +++ b/tests/test_failure_class_and_retry_guard.py @@ -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"