diff --git a/agent_platform/runtime/loop.py b/agent_platform/runtime/loop.py index 14d118c..04b4f8a 100644 --- a/agent_platform/runtime/loop.py +++ b/agent_platform/runtime/loop.py @@ -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: diff --git a/tests/test_agent_loop.py b/tests/test_agent_loop.py index 07b67c7..5e34d8c 100644 --- a/tests/test_agent_loop.py +++ b/tests/test_agent_loop.py @@ -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": [