fix: prevent unchanged execution retry loops
This commit is contained in:
@@ -70,6 +70,7 @@ EXECUTION_REQUEST = re.compile(
|
||||
MUTATION_TOOLS = {"write_file", "apply_patch"}
|
||||
VERIFICATION_TOOLS = {"read_file", "git_status", "git_diff", "exec", "poll_process"}
|
||||
EXECUTION_TOOLS = {"exec", "start_process", "poll_process"}
|
||||
REPEAT_GUARDED_EXECUTION_TOOLS = {"exec", "start_process"}
|
||||
MAX_CONSECUTIVE_CHECKPOINT_REJECTIONS = 3
|
||||
|
||||
|
||||
@@ -133,6 +134,28 @@ def recovery_message(failure: str) -> str:
|
||||
)
|
||||
|
||||
|
||||
def execution_recovery_message(failures: list[str]) -> str:
|
||||
details = "\n".join(f"- {failure}" for failure in failures)
|
||||
return (
|
||||
"Execution recovery is required before you may finish the task.\n"
|
||||
f"{details}\n"
|
||||
"Do not repeat an unchanged failed command. Text in assistant content does not edit a file. "
|
||||
"Use read_file or search_files to inspect the actual source, then use write_file or apply_patch "
|
||||
"to repair it (or use a materially different diagnostic command), and finally execute the relevant "
|
||||
"check successfully. Keep benchmark inputs small enough for quadratic algorithms to finish."
|
||||
)
|
||||
|
||||
|
||||
def execution_call_key(name: str, arguments: dict[str, Any]) -> str | None:
|
||||
if name not in REPEAT_GUARDED_EXECUTION_TOOLS:
|
||||
return None
|
||||
command = " ".join(str(arguments.get("command", "")).split())
|
||||
if not command:
|
||||
return None
|
||||
cwd = str(arguments.get("cwd", ".")).strip() or "."
|
||||
return f"{name}\0{cwd}\0{command}"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class RunRecorder:
|
||||
store: RuntimeStore
|
||||
@@ -246,6 +269,7 @@ class AgentLoop:
|
||||
successful_tools = 0
|
||||
successful_executions = 0
|
||||
unresolved_execution_failure: str | None = None
|
||||
failed_execution_revisions: dict[str, int] = {}
|
||||
consecutive_checkpoint_rejections = 0
|
||||
for iteration in range(max_iterations):
|
||||
await recorder.emit(
|
||||
@@ -321,6 +345,20 @@ class AgentLoop:
|
||||
"tool_calls": tool_calls,
|
||||
}
|
||||
messages.append(assistant_message)
|
||||
blocked_call_ids: set[str] = set()
|
||||
mutation_planned = False
|
||||
for call in tool_calls:
|
||||
function = call.get("function") or {}
|
||||
name = str(function.get("name", ""))
|
||||
try:
|
||||
arguments = json.loads(function.get("arguments") or "{}")
|
||||
except json.JSONDecodeError:
|
||||
arguments = {}
|
||||
if name in MUTATION_TOOLS:
|
||||
mutation_planned = True
|
||||
key = execution_call_key(name, arguments) if isinstance(arguments, dict) else None
|
||||
if key and not mutation_planned and failed_execution_revisions.get(key) == mutation_count:
|
||||
blocked_call_ids.add(str(call.get("id") or ""))
|
||||
results = await self._execute_calls(
|
||||
calls=tool_calls,
|
||||
spec=spec,
|
||||
@@ -328,9 +366,16 @@ class AgentLoop:
|
||||
context=tool_context,
|
||||
depth=depth,
|
||||
read_only=read_only,
|
||||
blocked_call_ids=blocked_call_ids,
|
||||
)
|
||||
execution_failures: list[str] = []
|
||||
for call, result in zip(tool_calls, results, strict=True):
|
||||
name = str((call.get("function") or {}).get("name", ""))
|
||||
function = call.get("function") or {}
|
||||
name = str(function.get("name", ""))
|
||||
try:
|
||||
arguments = json.loads(function.get("arguments") or "{}")
|
||||
except json.JSONDecodeError:
|
||||
arguments = {}
|
||||
if result.get("ok", False):
|
||||
successful_tools += 1
|
||||
if name in MUTATION_TOOLS:
|
||||
@@ -342,6 +387,10 @@ class AgentLoop:
|
||||
unresolved_execution_failure = None
|
||||
elif name in EXECUTION_TOOLS:
|
||||
unresolved_execution_failure = f"{name}: {public_tool_summary(result, 1000)}"
|
||||
execution_failures.append(unresolved_execution_failure)
|
||||
key = execution_call_key(name, arguments) if isinstance(arguments, dict) else None
|
||||
if key:
|
||||
failed_execution_revisions[key] = mutation_count
|
||||
messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
@@ -349,6 +398,13 @@ class AgentLoop:
|
||||
"content": tool_result_text(result, self.max_tool_output_chars),
|
||||
}
|
||||
)
|
||||
if execution_failures:
|
||||
messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": execution_recovery_message(execution_failures),
|
||||
}
|
||||
)
|
||||
|
||||
messages.append(
|
||||
{
|
||||
@@ -388,7 +444,9 @@ class AgentLoop:
|
||||
context: ToolContext,
|
||||
depth: int,
|
||||
read_only: bool,
|
||||
blocked_call_ids: set[str] | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
blocked_call_ids = blocked_call_ids or set()
|
||||
parsed: list[tuple[int, dict[str, Any], str, dict[str, Any], bool]] = []
|
||||
for index, call in enumerate(calls):
|
||||
function = call.get("function") or {}
|
||||
@@ -419,6 +477,7 @@ class AgentLoop:
|
||||
context=context,
|
||||
depth=depth,
|
||||
read_only=read_only,
|
||||
blocked_repeat=str(call.get("id") or "") in blocked_call_ids,
|
||||
)
|
||||
|
||||
index = 0
|
||||
@@ -445,6 +504,7 @@ class AgentLoop:
|
||||
context: ToolContext,
|
||||
depth: int,
|
||||
read_only: bool,
|
||||
blocked_repeat: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
call_id = str(call.get("id") or uuid.uuid4().hex)
|
||||
public_args = {
|
||||
@@ -456,6 +516,15 @@ class AgentLoop:
|
||||
)
|
||||
if "__parse_error__" in arguments:
|
||||
result = {"ok": False, "error": arguments["__parse_error__"]}
|
||||
elif blocked_repeat:
|
||||
result = {
|
||||
"ok": False,
|
||||
"error": (
|
||||
"Execution policy blocked this unchanged retry because the same command already failed "
|
||||
"and no repair was applied afterward. Inspect and edit the underlying files first, or run "
|
||||
"a materially different diagnostic command."
|
||||
),
|
||||
}
|
||||
elif name not in TOOL_METADATA:
|
||||
result = {"ok": False, "error": f"Unknown tool: {name}"}
|
||||
elif read_only and name not in READ_ONLY_TOOL_NAMES:
|
||||
|
||||
+100
-2
@@ -268,12 +268,23 @@ async def test_failed_execution_must_be_repaired_before_completion(settings) ->
|
||||
]
|
||||
}
|
||||
premature = {"choices": [{"message": {"role": "assistant", "content": "已完成"}, "finish_reason": "stop"}]}
|
||||
repair = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("4", "write_file", {"path": "sort.py", "content": "fixed"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
repaired_exec = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("4", "exec", {"command": "python3 sort.py"})],
|
||||
"tool_calls": [tool_call("5", "exec", {"command": "python3 sort.py"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
@@ -282,7 +293,7 @@ async def test_failed_execution_must_be_repaired_before_completion(settings) ->
|
||||
final = {
|
||||
"choices": [{"message": {"role": "assistant", "content": "脚本运行和报告验证完成"}, "finish_reason": "stop"}]
|
||||
}
|
||||
provider = ScriptedProvider([write, read, failed_exec, premature, repaired_exec, final])
|
||||
provider = ScriptedProvider([write, read, failed_exec, premature, repair, repaired_exec, final])
|
||||
registry = FailingExecutionRegistry()
|
||||
store = RuntimeStore(settings.database_url)
|
||||
await store.initialize()
|
||||
@@ -303,6 +314,9 @@ async def test_failed_execution_must_be_repaired_before_completion(settings) ->
|
||||
assert answer == "脚本运行和报告验证完成"
|
||||
assert registry.exec_calls == 2
|
||||
assert events.count("completion.rejected") == 1
|
||||
direct_recovery = provider.requests[3]["messages"][-1]
|
||||
assert direct_recovery["role"] == "user"
|
||||
assert "Execution recovery is required" in direct_recovery["content"]
|
||||
recovery = provider.requests[4]["messages"][-1]
|
||||
assert recovery["role"] == "user"
|
||||
assert "SyntaxError" in recovery["content"]
|
||||
@@ -310,6 +324,90 @@ async def test_failed_execution_must_be_repaired_before_completion(settings) ->
|
||||
await store.close()
|
||||
|
||||
|
||||
async def test_unchanged_failed_command_is_blocked_until_a_repair(settings) -> None:
|
||||
write = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("1", "write_file", {"path": "sort.py", "content": "broken"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
failed_exec = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("2", "exec", {"command": "python3 sort.py"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
repeated_exec = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("3", "exec", {"command": "python3 sort.py"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
repair = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("4", "write_file", {"path": "sort.py", "content": "fixed"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
successful_exec = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("5", "exec", {"command": "python3 sort.py"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
final = {"choices": [{"message": {"role": "assistant", "content": "验证完成"}, "finish_reason": "stop"}]}
|
||||
provider = ScriptedProvider([write, failed_exec, repeated_exec, repair, successful_exec, final])
|
||||
registry = FailingExecutionRegistry()
|
||||
store = RuntimeStore(settings.database_url)
|
||||
await store.initialize()
|
||||
|
||||
async def callback(event_type, payload):
|
||||
return None
|
||||
|
||||
try:
|
||||
answer = await AgentLoop(provider, registry, store, max_tool_output_chars=10_000).run(
|
||||
spec=get_model_spec("work-light"),
|
||||
messages=[{"role": "user", "content": "写脚本并运行"}],
|
||||
identity=UserIdentity("u1", "", "", "user"),
|
||||
raw_user_jwt="jwt",
|
||||
chat_id="c1",
|
||||
callback=callback,
|
||||
)
|
||||
assert answer == "验证完成"
|
||||
assert registry.exec_calls == 2
|
||||
assert [name for name, _ in registry.calls] == ["write_file", "exec", "write_file", "exec"]
|
||||
blocked_result = provider.requests[3]["messages"][-2]
|
||||
assert blocked_result["role"] == "tool"
|
||||
assert "unchanged retry" in blocked_result["content"]
|
||||
finally:
|
||||
await store.close()
|
||||
|
||||
|
||||
async def test_script_delivery_requires_successful_execution(settings) -> None:
|
||||
write = {
|
||||
"choices": [
|
||||
|
||||
Reference in New Issue
Block a user