Jelajahi Sumber

框架:让Runner按失败处置语义交还控制权

Trace 与 run_result 持久化 FailureDetail。Runner 统一识别工具控制字段:replan_task 会立即结束 Worker 或 Validator,但 Planner仍可读取完整错误继续修订;abort_run 会终止当前角色。\n\n只有完全没有结构化失败且缺少 terminal submit 时才生成 PROTOCOL_VIOLATION。并行工具调用按处置强度聚合,成功 terminal 与失败 terminal 不再混用。\n\n补充 Worker立即返还 Planner、Planner保留控制权和协议失败兜底测试。
SamLee 8 jam lalu
induk
melakukan
cd2734a46f

+ 121 - 4
agent/agent/core/runner.py

@@ -31,6 +31,7 @@ if TYPE_CHECKING:
     from agent.core.dream import DreamReport
 
 from agent.trace.models import Trace, Message
+from agent.failures import FailureDetail, FailureDisposition
 from agent.trace.protocols import TraceAttachmentStore, TraceStore
 from agent.trace.goal_models import GoalTree
 from agent.trace.compaction import (
@@ -59,6 +60,13 @@ from agent.core.prompts import (
 
 logger = logging.getLogger(__name__)
 
+_FAILURE_DISPOSITION_PRIORITY = {
+    FailureDisposition.RETRY_CALL: 0,
+    FailureDisposition.REPAIR_ATTEMPT: 1,
+    FailureDisposition.REPLAN_TASK: 2,
+    FailureDisposition.ABORT_RUN: 3,
+}
+
 
 @dataclass
 class ContextUsage:
@@ -490,6 +498,7 @@ class AgentRunner:
 
         status = final_trace.status if final_trace else "unknown"
         error = final_trace.error_message if final_trace else None
+        failure = final_trace.failure if final_trace else None
         summary = last_assistant_text or (
             final_trace.result_summary if final_trace else ""
         )
@@ -507,6 +516,7 @@ class AgentRunner:
             "summary": summary,
             "trace_id": trace_id,
             "error": error,
+            "failure": failure.to_dict() if failure else None,
             "saved_knowledge_ids": saved_knowledge_ids,  # 新增:返回保存的知识 ID
             "stats": {
                 "total_messages": final_trace.total_messages if final_trace else 0,
@@ -1215,6 +1225,8 @@ class AgentRunner:
         )
         explicit_terminal_submitted = False
         terminal_summary: Optional[str] = None
+        terminal_failure: Optional[FailureDetail] = None
+        last_failure: Optional[FailureDetail] = None
 
         # 当前主路径头节点的 sequence(用于设置 parent_sequence)
         head_seq = trace.head_sequence
@@ -2023,6 +2035,12 @@ class AgentRunner:
                     for res in results:
                         tc, tool_args, tool_result = res
                         tool_name = tc["function"]["name"]
+                        terminal_control = None
+                        if isinstance(tool_result, dict):
+                            terminal_control = tool_result.get("_control")
+                            if terminal_control:
+                                tool_result = dict(tool_result)
+                                tool_result.pop("_control", None)
                         if tool_args is None:
                             history.append(
                                 {
@@ -2154,6 +2172,30 @@ class AgentRunner:
                                     )
                             except Exception as e:
                                 self.log.warning(f"[Skill 指定注入] 记录追踪失败: {e}")
+
+                        observed_failure = self._control_failure(terminal_control)
+                        if observed_failure is not None:
+                            last_failure = self._stronger_failure(
+                                last_failure, observed_failure
+                            )
+                            if self._failure_terminates_role(
+                                observed_failure, explicit_role
+                            ):
+                                terminal_failure = self._stronger_failure(
+                                    terminal_failure, observed_failure
+                                )
+                                terminal_run = True
+                        elif terminal_control and terminal_control.get("terminate_run"):
+                            terminal_summary = (
+                                terminal_control.get("result_summary") or tool_text
+                            )
+                            terminal_run = True
+                    if terminal_summary and terminal_failure is None and self.trace_store:
+                        await self.trace_store.update_trace(
+                            trace_id,
+                            result_summary=terminal_summary,
+                            head_sequence=head_seq,
+                        )
                 else:
                     for tc in tool_calls:
                         current_goal_id = (
@@ -2408,7 +2450,16 @@ class AgentRunner:
                             except Exception as e:
                                 self.log.warning(f"[Skill 指定注入] 记录追踪失败: {e}")
 
-                        if terminal_control and terminal_control.get("terminate_run"):
+                        observed_failure = self._control_failure(terminal_control)
+                        if observed_failure is not None:
+                            last_failure = observed_failure
+                            if self._failure_terminates_role(
+                                observed_failure, explicit_role
+                            ):
+                                terminal_failure = observed_failure
+                                terminal_run = True
+                                break
+                        elif terminal_control and terminal_control.get("terminate_run"):
                             terminal_run = True
                             terminal_summary = (
                                 terminal_control.get("result_summary") or tool_text
@@ -2422,7 +2473,10 @@ class AgentRunner:
                             break
 
                 if terminal_run:
-                    if explicit_role in (AgentRole.WORKER, AgentRole.VALIDATOR):
+                    if (
+                        terminal_failure is None
+                        and explicit_role in (AgentRole.WORKER, AgentRole.VALIDATOR)
+                    ):
                         explicit_terminal_submitted = True
                     break
 
@@ -2531,8 +2585,13 @@ class AgentRunner:
         final_status = "completed"
         final_error: Optional[str] = None
         final_summary: Optional[str] = terminal_summary
+        final_failure: Optional[FailureDetail] = None
 
-        if explicit_role == AgentRole.PLANNER:
+        if terminal_failure is not None:
+            final_status = "failed"
+            final_failure = terminal_failure
+            final_error = self._failure_text(terminal_failure)
+        elif explicit_role == AgentRole.PLANNER:
             mission = await self._get_mission_completion(trace_id)
             mission_status = self._status_value(mission.get("status"))
             if mission_status == "completed":
@@ -2549,7 +2608,20 @@ class AgentRunner:
         elif explicit_role in (AgentRole.WORKER, AgentRole.VALIDATOR):
             if not explicit_terminal_submitted:
                 final_status = "failed"
-                final_error = "protocol_violation"
+                if last_failure is not None:
+                    final_failure = last_failure
+                    final_error = self._failure_text(last_failure)
+                else:
+                    final_failure = FailureDetail(
+                        code="PROTOCOL_VIOLATION",
+                        message=(
+                            "Worker ended without submit_attempt"
+                            if explicit_role == AgentRole.WORKER
+                            else "Validator ended without submit_validation"
+                        ),
+                        disposition=FailureDisposition.ABORT_RUN,
+                    )
+                    final_error = "protocol_violation"
 
         # 更新 head_sequence 和 Trace 终态
         if self.trace_store:
@@ -2557,6 +2629,7 @@ class AgentRunner:
                 "status": final_status,
                 "head_sequence": head_seq,
                 "error_message": final_error,
+                "failure": final_failure,
                 "completed_at": datetime.now(),
             }
             if final_summary:
@@ -2566,6 +2639,50 @@ class AgentRunner:
             if trace_obj:
                 yield trace_obj
 
+    @staticmethod
+    def _control_failure(control: Any) -> Optional[FailureDetail]:
+        if not isinstance(control, dict) or control.get("failure") is None:
+            return None
+        try:
+            return FailureDetail.from_dict(control["failure"])
+        except (TypeError, ValueError):
+            logger.error("Runner received an invalid structured tool failure")
+            return FailureDetail(
+                code="INVALID_FAILURE_ENVELOPE",
+                message="The tool returned an invalid failure envelope",
+                disposition=FailureDisposition.ABORT_RUN,
+            )
+
+    @staticmethod
+    def _failure_terminates_role(
+        failure: FailureDetail,
+        role: Optional[AgentRole],
+    ) -> bool:
+        if failure.disposition == FailureDisposition.ABORT_RUN:
+            return True
+        return (
+            failure.disposition == FailureDisposition.REPLAN_TASK
+            and role in (AgentRole.WORKER, AgentRole.VALIDATOR)
+        )
+
+    @staticmethod
+    def _stronger_failure(
+        current: Optional[FailureDetail],
+        candidate: FailureDetail,
+    ) -> FailureDetail:
+        if current is None:
+            return candidate
+        if (
+            _FAILURE_DISPOSITION_PRIORITY[candidate.disposition]
+            > _FAILURE_DISPOSITION_PRIORITY[current.disposition]
+        ):
+            return candidate
+        return current
+
+    @staticmethod
+    def _failure_text(failure: FailureDetail) -> str:
+        return f"{failure.code}: {failure.message}"
+
     async def _get_mission_completion(self, root_trace_id: str) -> Dict[str, Any]:
         """Read the coordinator-owned mission completion snapshot."""
         method = getattr(self.task_coordinator, "root_completion", None)

+ 10 - 0
agent/agent/trace/models.py

@@ -10,6 +10,8 @@ from datetime import datetime
 from typing import Dict, Any, List, Optional, Literal, Union
 import uuid
 
+from agent.failures import FailureDetail
+
 # ===== 消息线格式类型别名 =====
 # 轻量 wire-format 类型,用于工具参数和 runner/LLM API 接口。
 # 内部存储使用下方的 Message dataclass。
@@ -76,6 +78,7 @@ class Trace:
     tools: Optional[List[Dict]] = None       # 工具定义(整个 trace 共享)
     llm_params: Dict[str, Any] = field(default_factory=dict)  # LLM 参数(temperature 等)
     context: Dict[str, Any] = field(default_factory=dict)     # 其他元数据
+    runtime_state: Dict[str, Any] = field(default_factory=dict)  # 框架内部续跑状态
 
     # 当前焦点 goal
     current_goal_id: Optional[str] = None
@@ -88,6 +91,7 @@ class Trace:
     # 结果
     result_summary: Optional[str] = None     # 执行结果摘要
     error_message: Optional[str] = None      # 错误信息
+    failure: Optional[FailureDetail] = None  # 结构化终态失败
 
     # 时间
     created_at: datetime = field(default_factory=datetime.now)
@@ -112,6 +116,8 @@ class Trace:
         """从字典创建 Trace(处理日期字段反序列化)"""
         from dateutil import parser
 
+        data = dict(data)
+
         # 处理日期字段
         if "created_at" in data and isinstance(data["created_at"], str):
             data["created_at"] = parser.isoparse(data["created_at"])
@@ -119,6 +125,8 @@ class Trace:
             data["completed_at"] = parser.isoparse(data["completed_at"])
         if "last_activity_at" in data and isinstance(data["last_activity_at"], str):
             data["last_activity_at"] = parser.isoparse(data["last_activity_at"])
+        if data.get("failure") is not None:
+            data["failure"] = FailureDetail.from_dict(data["failure"])
 
         return cls(**data)
 
@@ -151,10 +159,12 @@ class Trace:
             "tools": self.tools,
             "llm_params": self.llm_params,
             "context": self.context,
+            "runtime_state": self.runtime_state,
             "current_goal_id": self.current_goal_id,
             "reflected_at_sequence": self.reflected_at_sequence,
             "result_summary": self.result_summary,
             "error_message": self.error_message,
+            "failure": self.failure.to_dict() if self.failure else None,
             "created_at": self.created_at.isoformat() if self.created_at else None,
             "completed_at": self.completed_at.isoformat() if self.completed_at else None,
             "last_activity_at": self.last_activity_at.isoformat() if self.last_activity_at else None,

+ 124 - 0
agent/tests/test_runner_completion_gate.py

@@ -2,9 +2,12 @@ import pytest
 
 from agent.core.presets import AgentPreset, register_preset
 from agent.core.runner import AgentRunner, RunConfig
+from agent.failures import FailureDetail, FailureDisposition, ToolExecutionError
 from agent.orchestration.models import AgentRole, CompletionPolicy
 from agent.trace.models import Trace
 from agent.trace.store import FileSystemTraceStore
+from agent.tools.models import ToolCapability
+from agent.tools.registry import ToolRegistry
 
 
 class StubCoordinator:
@@ -329,3 +332,124 @@ async def test_explicit_subtrace_without_terminal_submit_is_protocol_failure(
 
     assert result["status"] == "failed"
     assert result["error"] == "protocol_violation"
+    assert result["failure"]["code"] == "PROTOCOL_VIOLATION"
+
+
+@pytest.mark.asyncio
+async def test_worker_replan_failure_terminates_without_protocol_overwrite(tmp_path):
+    preset = "test_structured_failure_worker"
+    register_preset(
+        preset,
+        AgentPreset(
+            role=AgentRole.WORKER,
+            allowed_tools=["reject_contract"],
+            max_iterations=4,
+            skills=[],
+        ),
+    )
+    registry = ToolRegistry()
+
+    async def reject_contract() -> str:
+        raise ToolExecutionError(
+            FailureDetail(
+                code="INPUT_SCOPE_MISMATCH",
+                message="Paragraph has no covering Structure",
+                disposition=FailureDisposition.REPLAN_TASK,
+            )
+        )
+
+    registry.register(
+        reject_contract,
+        capabilities=[ToolCapability.READ],
+    )
+    calls = 0
+
+    async def llm_call(**_kwargs):
+        nonlocal calls
+        calls += 1
+        return {
+            "content": "",
+            "tool_calls": [
+                {
+                    "id": "call-reject",
+                    "type": "function",
+                    "function": {"name": "reject_contract", "arguments": "{}"},
+                }
+            ],
+            "finish_reason": "tool_calls",
+        }
+
+    store = FileSystemTraceStore(str(tmp_path))
+    runner = AgentRunner(
+        trace_store=store,
+        tool_registry=registry,
+        llm_call=llm_call,
+        task_coordinator=StubCoordinator([{"status": "completed"}]),
+    )
+    config = _explicit_config(preset, 4)
+    config.tools = ["reject_contract"]
+    result = await runner.run_result([{"role": "user", "content": "work"}], config)
+
+    assert calls == 1
+    assert result["status"] == "failed"
+    assert result["error"].startswith("INPUT_SCOPE_MISMATCH:")
+    assert result["failure"]["code"] == "INPUT_SCOPE_MISMATCH"
+    trace = await store.get_trace(result["trace_id"])
+    assert trace.failure.code == "INPUT_SCOPE_MISMATCH"
+
+
+@pytest.mark.asyncio
+async def test_planner_observes_replan_failure_and_keeps_control(tmp_path):
+    preset = "test_structured_failure_planner"
+    register_preset(
+        preset,
+        AgentPreset(
+            role=AgentRole.PLANNER,
+            allowed_tools=["reject_contract"],
+            max_iterations=3,
+            skills=[],
+        ),
+    )
+    registry = ToolRegistry()
+
+    async def reject_contract() -> str:
+        raise ToolExecutionError(
+            FailureDetail(
+                code="TASK_CONTRACT_NOT_EXECUTABLE",
+                message="revise the frozen contract",
+                disposition=FailureDisposition.REPLAN_TASK,
+            )
+        )
+
+    registry.register(reject_contract, capabilities=[ToolCapability.TASK_CONTROL])
+    calls = 0
+
+    async def llm_call(**_kwargs):
+        nonlocal calls
+        calls += 1
+        if calls == 1:
+            return {
+                "content": "",
+                "tool_calls": [
+                    {
+                        "id": "call-reject",
+                        "type": "function",
+                        "function": {"name": "reject_contract", "arguments": "{}"},
+                    }
+                ],
+                "finish_reason": "tool_calls",
+            }
+        return {"content": "replanned", "tool_calls": None, "finish_reason": "stop"}
+
+    config = _explicit_config(preset, 3)
+    config.tools = ["reject_contract"]
+    result = await AgentRunner(
+        trace_store=FileSystemTraceStore(str(tmp_path)),
+        tool_registry=registry,
+        llm_call=llm_call,
+        task_coordinator=StubCoordinator([{"status": "completed"}]),
+    ).run_result([{"role": "user", "content": "plan"}], config)
+
+    assert calls == 2
+    assert result["status"] == "completed"
+    assert result["failure"] is None