Ver Fonte

框架:让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 há 13 horas atrás
pai
commit
cd2734a46f

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

@@ -31,6 +31,7 @@ if TYPE_CHECKING:
     from agent.core.dream import DreamReport
     from agent.core.dream import DreamReport
 
 
 from agent.trace.models import Trace, Message
 from agent.trace.models import Trace, Message
+from agent.failures import FailureDetail, FailureDisposition
 from agent.trace.protocols import TraceAttachmentStore, TraceStore
 from agent.trace.protocols import TraceAttachmentStore, TraceStore
 from agent.trace.goal_models import GoalTree
 from agent.trace.goal_models import GoalTree
 from agent.trace.compaction import (
 from agent.trace.compaction import (
@@ -59,6 +60,13 @@ from agent.core.prompts import (
 
 
 logger = logging.getLogger(__name__)
 logger = logging.getLogger(__name__)
 
 
+_FAILURE_DISPOSITION_PRIORITY = {
+    FailureDisposition.RETRY_CALL: 0,
+    FailureDisposition.REPAIR_ATTEMPT: 1,
+    FailureDisposition.REPLAN_TASK: 2,
+    FailureDisposition.ABORT_RUN: 3,
+}
+
 
 
 @dataclass
 @dataclass
 class ContextUsage:
 class ContextUsage:
@@ -490,6 +498,7 @@ class AgentRunner:
 
 
         status = final_trace.status if final_trace else "unknown"
         status = final_trace.status if final_trace else "unknown"
         error = final_trace.error_message if final_trace else None
         error = final_trace.error_message if final_trace else None
+        failure = final_trace.failure if final_trace else None
         summary = last_assistant_text or (
         summary = last_assistant_text or (
             final_trace.result_summary if final_trace else ""
             final_trace.result_summary if final_trace else ""
         )
         )
@@ -507,6 +516,7 @@ class AgentRunner:
             "summary": summary,
             "summary": summary,
             "trace_id": trace_id,
             "trace_id": trace_id,
             "error": error,
             "error": error,
+            "failure": failure.to_dict() if failure else None,
             "saved_knowledge_ids": saved_knowledge_ids,  # 新增:返回保存的知识 ID
             "saved_knowledge_ids": saved_knowledge_ids,  # 新增:返回保存的知识 ID
             "stats": {
             "stats": {
                 "total_messages": final_trace.total_messages if final_trace else 0,
                 "total_messages": final_trace.total_messages if final_trace else 0,
@@ -1215,6 +1225,8 @@ class AgentRunner:
         )
         )
         explicit_terminal_submitted = False
         explicit_terminal_submitted = False
         terminal_summary: Optional[str] = None
         terminal_summary: Optional[str] = None
+        terminal_failure: Optional[FailureDetail] = None
+        last_failure: Optional[FailureDetail] = None
 
 
         # 当前主路径头节点的 sequence(用于设置 parent_sequence)
         # 当前主路径头节点的 sequence(用于设置 parent_sequence)
         head_seq = trace.head_sequence
         head_seq = trace.head_sequence
@@ -2023,6 +2035,12 @@ class AgentRunner:
                     for res in results:
                     for res in results:
                         tc, tool_args, tool_result = res
                         tc, tool_args, tool_result = res
                         tool_name = tc["function"]["name"]
                         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:
                         if tool_args is None:
                             history.append(
                             history.append(
                                 {
                                 {
@@ -2154,6 +2172,30 @@ class AgentRunner:
                                     )
                                     )
                             except Exception as e:
                             except Exception as e:
                                 self.log.warning(f"[Skill 指定注入] 记录追踪失败: {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:
                 else:
                     for tc in tool_calls:
                     for tc in tool_calls:
                         current_goal_id = (
                         current_goal_id = (
@@ -2408,7 +2450,16 @@ class AgentRunner:
                             except Exception as e:
                             except Exception as e:
                                 self.log.warning(f"[Skill 指定注入] 记录追踪失败: {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_run = True
                             terminal_summary = (
                             terminal_summary = (
                                 terminal_control.get("result_summary") or tool_text
                                 terminal_control.get("result_summary") or tool_text
@@ -2422,7 +2473,10 @@ class AgentRunner:
                             break
                             break
 
 
                 if terminal_run:
                 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
                         explicit_terminal_submitted = True
                     break
                     break
 
 
@@ -2531,8 +2585,13 @@ class AgentRunner:
         final_status = "completed"
         final_status = "completed"
         final_error: Optional[str] = None
         final_error: Optional[str] = None
         final_summary: Optional[str] = terminal_summary
         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 = await self._get_mission_completion(trace_id)
             mission_status = self._status_value(mission.get("status"))
             mission_status = self._status_value(mission.get("status"))
             if mission_status == "completed":
             if mission_status == "completed":
@@ -2549,7 +2608,20 @@ class AgentRunner:
         elif explicit_role in (AgentRole.WORKER, AgentRole.VALIDATOR):
         elif explicit_role in (AgentRole.WORKER, AgentRole.VALIDATOR):
             if not explicit_terminal_submitted:
             if not explicit_terminal_submitted:
                 final_status = "failed"
                 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 终态
         # 更新 head_sequence 和 Trace 终态
         if self.trace_store:
         if self.trace_store:
@@ -2557,6 +2629,7 @@ class AgentRunner:
                 "status": final_status,
                 "status": final_status,
                 "head_sequence": head_seq,
                 "head_sequence": head_seq,
                 "error_message": final_error,
                 "error_message": final_error,
+                "failure": final_failure,
                 "completed_at": datetime.now(),
                 "completed_at": datetime.now(),
             }
             }
             if final_summary:
             if final_summary:
@@ -2566,6 +2639,50 @@ class AgentRunner:
             if trace_obj:
             if trace_obj:
                 yield 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]:
     async def _get_mission_completion(self, root_trace_id: str) -> Dict[str, Any]:
         """Read the coordinator-owned mission completion snapshot."""
         """Read the coordinator-owned mission completion snapshot."""
         method = getattr(self.task_coordinator, "root_completion", None)
         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
 from typing import Dict, Any, List, Optional, Literal, Union
 import uuid
 import uuid
 
 
+from agent.failures import FailureDetail
+
 # ===== 消息线格式类型别名 =====
 # ===== 消息线格式类型别名 =====
 # 轻量 wire-format 类型,用于工具参数和 runner/LLM API 接口。
 # 轻量 wire-format 类型,用于工具参数和 runner/LLM API 接口。
 # 内部存储使用下方的 Message dataclass。
 # 内部存储使用下方的 Message dataclass。
@@ -76,6 +78,7 @@ class Trace:
     tools: Optional[List[Dict]] = None       # 工具定义(整个 trace 共享)
     tools: Optional[List[Dict]] = None       # 工具定义(整个 trace 共享)
     llm_params: Dict[str, Any] = field(default_factory=dict)  # LLM 参数(temperature 等)
     llm_params: Dict[str, Any] = field(default_factory=dict)  # LLM 参数(temperature 等)
     context: Dict[str, Any] = field(default_factory=dict)     # 其他元数据
     context: Dict[str, Any] = field(default_factory=dict)     # 其他元数据
+    runtime_state: Dict[str, Any] = field(default_factory=dict)  # 框架内部续跑状态
 
 
     # 当前焦点 goal
     # 当前焦点 goal
     current_goal_id: Optional[str] = None
     current_goal_id: Optional[str] = None
@@ -88,6 +91,7 @@ class Trace:
     # 结果
     # 结果
     result_summary: Optional[str] = None     # 执行结果摘要
     result_summary: Optional[str] = None     # 执行结果摘要
     error_message: Optional[str] = None      # 错误信息
     error_message: Optional[str] = None      # 错误信息
+    failure: Optional[FailureDetail] = None  # 结构化终态失败
 
 
     # 时间
     # 时间
     created_at: datetime = field(default_factory=datetime.now)
     created_at: datetime = field(default_factory=datetime.now)
@@ -112,6 +116,8 @@ class Trace:
         """从字典创建 Trace(处理日期字段反序列化)"""
         """从字典创建 Trace(处理日期字段反序列化)"""
         from dateutil import parser
         from dateutil import parser
 
 
+        data = dict(data)
+
         # 处理日期字段
         # 处理日期字段
         if "created_at" in data and isinstance(data["created_at"], str):
         if "created_at" in data and isinstance(data["created_at"], str):
             data["created_at"] = parser.isoparse(data["created_at"])
             data["created_at"] = parser.isoparse(data["created_at"])
@@ -119,6 +125,8 @@ class Trace:
             data["completed_at"] = parser.isoparse(data["completed_at"])
             data["completed_at"] = parser.isoparse(data["completed_at"])
         if "last_activity_at" in data and isinstance(data["last_activity_at"], str):
         if "last_activity_at" in data and isinstance(data["last_activity_at"], str):
             data["last_activity_at"] = parser.isoparse(data["last_activity_at"])
             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)
         return cls(**data)
 
 
@@ -151,10 +159,12 @@ class Trace:
             "tools": self.tools,
             "tools": self.tools,
             "llm_params": self.llm_params,
             "llm_params": self.llm_params,
             "context": self.context,
             "context": self.context,
+            "runtime_state": self.runtime_state,
             "current_goal_id": self.current_goal_id,
             "current_goal_id": self.current_goal_id,
             "reflected_at_sequence": self.reflected_at_sequence,
             "reflected_at_sequence": self.reflected_at_sequence,
             "result_summary": self.result_summary,
             "result_summary": self.result_summary,
             "error_message": self.error_message,
             "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,
             "created_at": self.created_at.isoformat() if self.created_at else None,
             "completed_at": self.completed_at.isoformat() if self.completed_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,
             "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.presets import AgentPreset, register_preset
 from agent.core.runner import AgentRunner, RunConfig
 from agent.core.runner import AgentRunner, RunConfig
+from agent.failures import FailureDetail, FailureDisposition, ToolExecutionError
 from agent.orchestration.models import AgentRole, CompletionPolicy
 from agent.orchestration.models import AgentRole, CompletionPolicy
 from agent.trace.models import Trace
 from agent.trace.models import Trace
 from agent.trace.store import FileSystemTraceStore
 from agent.trace.store import FileSystemTraceStore
+from agent.tools.models import ToolCapability
+from agent.tools.registry import ToolRegistry
 
 
 
 
 class StubCoordinator:
 class StubCoordinator:
@@ -329,3 +332,124 @@ async def test_explicit_subtrace_without_terminal_submit_is_protocol_failure(
 
 
     assert result["status"] == "failed"
     assert result["status"] == "failed"
     assert result["error"] == "protocol_violation"
     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