From 076bf3570dd08da2e98b3f9afa735cadd006c4b1 Mon Sep 17 00:00:00 2001 From: mmprotest Date: Thu, 9 Apr 2026 18:45:11 +1000 Subject: [PATCH] Add lightweight runner thrash guards for repeated failures --- tests/test_loop.py | 125 +++++++++++++++++++++++ villani_code/state.py | 182 ++++++++++++++++++++++++++++++++++ villani_code/state_tooling.py | 11 ++ 3 files changed, 318 insertions(+) diff --git a/tests/test_loop.py b/tests/test_loop.py index 6ad25ecf..c8fd582e 100644 --- a/tests/test_loop.py +++ b/tests/test_loop.py @@ -681,3 +681,128 @@ def create_message(self, payload, stream): assert out['execution']['terminated_reason'] == 'benchmark_repeated_mutation_denials' assert any(e.get('type') == 'benchmark_repeated_mutation_denials' for e in events) + + +def test_failure_loop_guard_blocks_mutation_until_grounding(tmp_path: Path) -> None: + class Client: + def __init__(self) -> None: + self.calls = 0 + self.payloads: list[dict] = [] + + def create_message(self, payload, stream): + self.calls += 1 + self.payloads.append(payload) + if self.calls <= 3: + return { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": f"patch-{self.calls}", + "name": "Patch", + "input": {"file_path": "note.txt", "unified_diff": "not-a-diff"}, + } + ], + } + if self.calls == 4: + return { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "write-blocked", + "name": "Write", + "input": {"file_path": "note.txt", "content": "blocked"}, + } + ], + } + if self.calls == 5: + return { + "role": "assistant", + "content": [{"type": "tool_use", "id": "read-ground", "name": "Read", "input": {"file_path": "note.txt"}}], + } + if self.calls == 6: + return { + "role": "assistant", + "content": [{"type": "tool_use", "id": "write-allowed", "name": "Write", "input": {"file_path": "note.txt", "content": "ok"}}], + } + return {"role": "assistant", "content": [{"type": "text", "text": "done"}]} + + (tmp_path / "note.txt").write_text("start\n", encoding="utf-8") + client = Client() + runner = Runner(client=client, repo=tmp_path, model="m", stream=False, plan_mode="off") + out = runner.run("apply minimal fix") + + tool_results = out["transcript"]["tool_results"] + assert any("Recent tool attempts keep failing." in str(item.get("content", "")) for item in tool_results) + assert "Recent tool attempts keep failing." in str(tool_results[3]["content"]) + assert "read/inspect step" in str(tool_results[3]["content"]) + assert tool_results[4]["is_error"] is False + assert tool_results[5]["is_error"] is False + + +def test_brittle_near_duplicate_shell_retry_guard_and_clear(tmp_path: Path) -> None: + class Client: + def __init__(self) -> None: + self.calls = 0 + + def create_message(self, payload, stream): + self.calls += 1 + if self.calls == 1: + return { + "role": "assistant", + "content": [{"type": "tool_use", "id": "bash-1", "name": "Bash", "input": {"command": "python -c \"import sys; sys.exit(1)\""}}], + } + if self.calls == 2: + return { + "role": "assistant", + "content": [{"type": "tool_use", "id": "bash-2", "name": "Bash", "input": {"command": "python -c 'import sys; sys.exit(2)'"}}], + } + if self.calls == 3: + return { + "role": "assistant", + "content": [{"type": "tool_use", "id": "bash-3", "name": "Bash", "input": {"command": "pwd"}}], + } + if self.calls == 4: + return { + "role": "assistant", + "content": [{"type": "tool_use", "id": "bash-4", "name": "Bash", "input": {"command": "python -c \"print('ok')\""}}], + } + return {"role": "assistant", "content": [{"type": "text", "text": "done"}]} + + runner = Runner(client=Client(), repo=tmp_path, model="m", stream=False) + out = runner.run("inspect and run shell checks") + results = out["transcript"]["tool_results"] + + assert "exit_code" in str(results[0]["content"]) + assert "brittle after recent failure" in str(results[1]["content"]) + assert results[1]["is_error"] is True + assert results[2]["is_error"] is False + assert results[3]["is_error"] is False + + +def test_single_failures_do_not_trigger_thrash_guards(tmp_path: Path) -> None: + class Client: + def __init__(self) -> None: + self.calls = 0 + + def create_message(self, payload, stream): + self.calls += 1 + if self.calls == 1: + return { + "role": "assistant", + "content": [{"type": "tool_use", "id": "read-1", "name": "Read", "input": {"file_path": "missing.txt"}}], + } + if self.calls == 2: + return { + "role": "assistant", + "content": [{"type": "tool_use", "id": "ls-1", "name": "Ls", "input": {"path": "."}}], + } + return {"role": "assistant", "content": [{"type": "text", "text": "done"}]} + + runner = Runner(client=Client(), repo=tmp_path, model="m", stream=False) + out = runner.run("inspect files") + tool_results = out["transcript"]["tool_results"] + assert tool_results[0]["is_error"] is True + assert "Recent tool attempts keep failing." not in str(tool_results[0]["content"]) + assert all("brittle shell retry" not in str(item.get("content", "")) for item in tool_results) diff --git a/villani_code/state.py b/villani_code/state.py index 1a2c80d4..dc6345c2 100644 --- a/villani_code/state.py +++ b/villani_code/state.py @@ -4,6 +4,7 @@ import json import re import time +from collections import deque from pathlib import Path from typing import Any, Callable @@ -498,6 +499,12 @@ def __init__( self._no_progress_cycles = 0 self._recovery_count = 0 self._last_failed_tool_sig = "" + self._recent_failed_actions: deque[dict[str, str]] = deque(maxlen=5) + self._surface_grounding_block: str | None = None + self._surface_grounding_block_message = "" + self._pending_recovery_message = "" + self._recent_failed_shell_forms: deque[str] = deque(maxlen=3) + self._shell_retry_block_form: str | None = None self._repo_map = "" self._retriever: Retriever | None = None self._context_budget = ( @@ -1445,6 +1452,7 @@ def _budget_reason( "occurrence": failure.occurrence_count, } ) + self._record_failure_and_update_guards(tool_name, tool_input, result) self.hooks.run_event( "PostToolUse", @@ -1528,6 +1536,14 @@ def _budget_reason( else self._pending_verification ) self._pending_verification = "" + if self._pending_recovery_message and next_user_content: + existing = str(next_user_content[-1].get("content", "")) + next_user_content[-1]["content"] = ( + f"{existing}\n\n{self._pending_recovery_message}" + if existing + else self._pending_recovery_message + ) + self._pending_recovery_message = "" messages.append({"role": "user", "content": next_user_content}) reason = _budget_reason() @@ -1555,6 +1571,172 @@ def _is_mutating_tool_call( return not command.startswith(readonly_prefixes) return False + def _normalize_failure_reason(self, content: str) -> str: + text = str(content or "").strip().lower() + if not text: + return "unknown_failure" + if "invalid input" in text or "missing" in text or "must " in text or "tool_result" in text: + return "invalid_input" + if "denied by permission policy" in text or "user denied tool execution" in text: + return "permission_denied" + if "benchmark policy blocked this mutation" in text: + return "policy_blocked" + if "no such file" in text or "filenotfounderror" in text: + return "path_not_found" + if "timeout" in text: + return "timeout" + compact = re.sub(r"'[^']*'|\"[^\"]*\"", "\"x\"", text) + compact = re.sub(r"\b\d+\b", "#", compact) + compact = re.sub(r"\s+", " ", compact).strip() + return compact[:96] or "runtime_error" + + def _failure_kind(self, content: str) -> str: + lowered = str(content or "").lower() + if any( + token in lowered + for token in ("invalid input", "missing required", "unknown tool", "tool_result", "must ") + ): + return "protocol" + return "runtime" + + def _shell_failure_from_result(self, result: dict[str, Any]) -> tuple[bool, str, str]: + if bool(result.get("is_error", False)): + reason = self._normalize_failure_reason(str(result.get("content", ""))) + return True, reason, "" + try: + payload = json.loads(str(result.get("content", ""))) + except Exception: + return False, "", "" + if not isinstance(payload, dict): + return False, "", "" + exit_code = payload.get("exit_code") + command = str(payload.get("command", "")) + if isinstance(exit_code, int) and exit_code != 0: + return True, f"exit_{exit_code}", self._normalize_shell_command(command) + return False, "", "" + + def _normalize_shell_command(self, command: str) -> str: + normalized = str(command or "").strip() + normalized = normalized.replace('"', "").replace("'", "") + normalized = re.sub(r"\b\d+\b", "#", normalized) + normalized = re.sub(r"\s+", " ", normalized.lower()).strip() + return normalized + + def _is_brittle_shell_command(self, command: str) -> bool: + cmd = str(command or "") + lowered = cmd.lower() + return ( + len(cmd) > 220 + or "<<" in cmd + or "\\n" in cmd + or lowered.count("'") + lowered.count('"') > 8 + or any(marker in lowered for marker in ("python -c", "node -e", "perl -e", "ruby -e")) + or any(marker in lowered for marker in ("&&", "||", ";")) + or bool(re.search(r"(>\s*[^ ]+\s*2>&1|2>\s*[^ ]+)", lowered)) + ) + + def _is_grounding_action(self, tool_name: str, tool_input: dict[str, Any]) -> bool: + if tool_name in {"Read", "Ls", "Grep", "Search", "Glob", "GitStatus", "GitDiff", "GitLog", "GitBranch"}: + return True + if tool_name != "Bash": + return False + command = str(tool_input.get("command", "")).strip().lower() + return bool(command) and not self._is_mutating_tool_call("Bash", tool_input) + + def _record_failure_and_update_guards(self, tool_name: str, tool_input: dict[str, Any], result: dict[str, Any]) -> None: + surface = "shell" if tool_name == "Bash" else "tool" + action_name = "shell" if surface == "shell" else tool_name + content = str(result.get("content", "")) + normalized_reason = self._normalize_failure_reason(content) + if surface == "shell": + failed, shell_reason, shell_form = self._shell_failure_from_result(result) + if not failed: + return + normalized_reason = shell_reason or normalized_reason + if shell_form: + self._recent_failed_shell_forms.append(shell_form) + elif not bool(result.get("is_error", False)): + return + + self._recent_failed_actions.append( + { + "surface": surface, + "action_name": action_name, + "normalized_reason": normalized_reason, + "failure_kind": self._failure_kind(content), + } + ) + + last3 = list(self._recent_failed_actions)[-3:] + last5 = list(self._recent_failed_actions)[-5:] + trigger = False + if surface == "tool": + same_tool = sum(1 for item in last3 if item["surface"] == "tool" and item["action_name"] == action_name) + trigger = same_tool >= 2 + else: + same_shell_reason = sum( + 1 + for item in last3 + if item["surface"] == "shell" and item["normalized_reason"] == normalized_reason + ) + trigger = same_shell_reason >= 2 + if not trigger and len(last5) >= 3: + counts: dict[tuple[str, str], int] = {} + for item in last5: + key = (item["surface"], item["normalized_reason"]) + counts[key] = counts.get(key, 0) + 1 + trigger = len(last5) >= 3 and any(count >= 2 for count in counts.values()) + if trigger: + self._surface_grounding_block = surface + self._surface_grounding_block_message = ( + f"Recent {surface} attempts keep failing. Reground with one read/inspect step before retrying risky {surface} mutations." + ) + self._pending_recovery_message = self._surface_grounding_block_message + + def _check_transient_action_guards(self, tool_name: str, tool_input: dict[str, Any]) -> str | None: + if self._surface_grounding_block: + if self._is_grounding_action(tool_name, tool_input): + self._surface_grounding_block = None + self._surface_grounding_block_message = "" + elif ( + self._surface_grounding_block == "tool" + and tool_name != "Bash" + and self._is_mutating_tool_call(tool_name, tool_input) + ): + return self._surface_grounding_block_message or "Recent tool attempts are failing; reground with one read/inspect action first." + elif ( + self._surface_grounding_block == "shell" + and tool_name == "Bash" + and self._is_mutating_tool_call(tool_name, tool_input) + ): + return self._surface_grounding_block_message or "Recent shell attempts are failing; reground with one read/inspect action first." + + if tool_name != "Bash": + if self._shell_retry_block_form: + self._shell_retry_block_form = None + self._recent_failed_shell_forms.clear() + return None + command = str(tool_input.get("command", "")) + normalized = self._normalize_shell_command(command) + near_duplicate = bool(normalized) and ( + normalized in self._recent_failed_shell_forms + or (self._shell_retry_block_form is not None and normalized == self._shell_retry_block_form) + ) + brittle = self._is_brittle_shell_command(command) + if self._shell_retry_block_form is not None and near_duplicate and brittle: + return ( + "Blocked brittle shell retry after recent failure. Use a simpler command, a temp script, a read step, or a non-shell edit path." + ) + if self._recent_failed_shell_forms and near_duplicate and brittle: + self._shell_retry_block_form = normalized + return ( + "Current shell form is brittle after recent failure. Use a simpler command, a temp script, a read step, or a non-shell edit path." + ) + if self._shell_retry_block_form and (not near_duplicate or not brittle): + self._shell_retry_block_form = None + self._recent_failed_shell_forms.clear() + return None + def _dispatch_event(self, event: dict[str, Any]) -> None: if "turn_index" not in event and isinstance(self._current_turn_index, int): event = {**event, "turn_index": self._current_turn_index} diff --git a/villani_code/state_tooling.py b/villani_code/state_tooling.py index c84c64ca..12bdaa5f 100644 --- a/villani_code/state_tooling.py +++ b/villani_code/state_tooling.py @@ -493,6 +493,17 @@ def execute_tool_with_policy( runner._tighten_tool_input(tool_name, tool_input) if tool_name in {"Write", "Patch"}: _normalize_mutation_payload(tool_name, tool_input) + guard_message = runner._check_transient_action_guards(tool_name, tool_input) + if guard_message: + runner.event_callback( + { + "type": "transient_action_guard_blocked", + "tool": tool_name, + "input": tool_input, + "message": guard_message, + } + ) + return {"content": guard_message, "is_error": True} policy = runner.permissions.evaluate_with_reason( tool_name,