فهرست منبع

框架:贯通编排失败分类与历史反馈

将 FailureDetail 从 Runner 贯通到 LocalAgentExecutor、Worker/Validator结果、TaskAttempt、ValidationReport、TaskCycleResult 与 wire view。新增 TOOL_FAILURE 和 NO_PROGRESS 稳定分类,结构化失败优先于宽泛的 AGENT_FAILED。\n\n失败 Attempt 持久化 Worker 摘要,并为后续所有 Attempt 生成 prior_feedback;即使尚未提交 Artifact、没有 Validator 报告,也能看到上次工具错误和未冻结状态。repair_feedback 保留为兼容入口。\n\n补充真实错误不被 protocol_violation 覆盖、Ledger往返和无Validation反馈测试。
SamLee 12 ساعت پیش
والد
کامیت
6db2445f84

+ 106 - 17
agent/agent/orchestration/coordinator.py

@@ -11,6 +11,8 @@ from dataclasses import asdict, is_dataclass
 from typing import Any, Callable, Dict, Iterable, List, Optional, Sequence, Tuple
 from uuid import uuid4
 
+from agent.failures import FailureDetail
+
 from .config import OrchestrationConfig
 from .models import (
     AgentRole,
@@ -681,6 +683,7 @@ class TaskCoordinator:
                     validation_id=result.validation_id,
                     validation=result.validation,
                     error="Dispatch with this idempotency key is still in progress",
+                    failure=result.failure,
                 )
             if errors[index] and result.error != errors[index]:
                 result = TaskCycleResult(
@@ -690,6 +693,7 @@ class TaskCoordinator:
                     validation_id=result.validation_id,
                     validation=result.validation,
                     error=errors[index],
+                    failure=result.failure,
                 )
             results.append(result)
         return results
@@ -857,6 +861,7 @@ class TaskCoordinator:
                     validation_id = candidate_id
                     break
         error = validation.error if validation else (attempt.error if attempt else None)
+        failure = validation.failure if validation else (attempt.failure if attempt else None)
         if task.status in {
             TaskStatus.RUNNING,
             TaskStatus.AWAITING_VALIDATION,
@@ -870,6 +875,7 @@ class TaskCoordinator:
             validation_id=validation_id,
             validation=validation,
             error=error,
+            failure=failure,
         )
 
     async def _recover_cycle_failure(
@@ -1174,6 +1180,7 @@ class TaskCoordinator:
                     )
                 else:
                     worker_result = None
+            prior_feedback = self._prior_feedback(ledger, task, attempt_id)
             worker_context = {
                 "root_trace_id": root_trace_id,
                 "task_id": task_id,
@@ -1189,10 +1196,11 @@ class TaskCoordinator:
                     else None
                 ),
                 "repair_feedback": (
-                    self._repair_feedback(ledger, task)
+                    prior_feedback
                     if attempt.execution_mode == "repair"
                     else None
                 ),
+                "prior_feedback": prior_feedback,
                 "accepted_child_results": accepted_child_results,
                 "operation_id": operation_id,
                 "execution_epoch": execution_epoch,
@@ -1247,8 +1255,15 @@ class TaskCoordinator:
                 assert worker_result is not None
                 stats = _failure_stats(
                     worker_result.execution_stats,
-                    _failure_code(worker_result.status, protocol_failure=True),
-                    override=worker_result.status == "completed",
+                    _failure_code(
+                        worker_result.status,
+                        protocol_failure=True,
+                        failure=worker_result.failure,
+                    ),
+                    override=(
+                        worker_result.status == "completed"
+                        and worker_result.failure is None
+                    ),
                 )
                 await self._mark_worker_failure(
                     root_trace_id,
@@ -1259,16 +1274,24 @@ class TaskCoordinator:
                     stats,
                     operation_id,
                     execution_epoch,
+                    failure=worker_result.failure,
+                    worker_summary=worker_result.summary,
                 )
+            current = await self.task_store.load(root_trace_id)
+            current_attempt = current.attempts[attempt_id]
             return TaskCycleResult(
                 task_id=task_id,
-                task_status=(await self.task_store.load(root_trace_id)).tasks[task_id].status,
+                task_status=current.tasks[task_id].status,
                 attempt_id=attempt_id,
                 error=(
-                    attempt.error
+                    current_attempt.error
                     or (worker_result.error if worker_result else None)
                     or "Worker did not produce a submitted attempt"
                 ),
+                failure=(
+                    current_attempt.failure
+                    or (worker_result.failure if worker_result else None)
+                ),
             )
 
         validation_id = self._validation_for_attempt(ledger, task, attempt_id)
@@ -1330,6 +1353,7 @@ class TaskCoordinator:
                 validation_id=validation_id,
                 validation=validation,
                 error=validation.error,
+                failure=validation.failure,
             )
 
         try:
@@ -1438,8 +1462,15 @@ class TaskCoordinator:
         if validation.status != ValidationRunStatus.COMPLETED:
             stats = _failure_stats(
                 validator_result.execution_stats,
-                _failure_code(validator_result.status, protocol_failure=True),
-                override=validator_result.status == "completed",
+                _failure_code(
+                    validator_result.status,
+                    protocol_failure=True,
+                    failure=validator_result.failure,
+                ),
+                override=(
+                    validator_result.status == "completed"
+                    and validator_result.failure is None
+                ),
             )
             await self._mark_validation_error(
                 root_trace_id,
@@ -1450,6 +1481,7 @@ class TaskCoordinator:
                 stats,
                 operation_id,
                 execution_epoch,
+                failure=validator_result.failure,
             )
             ledger = await self.task_store.load(root_trace_id)
             validation = ledger.validations[validation_id]
@@ -1460,6 +1492,7 @@ class TaskCoordinator:
             validation_id=validation_id,
             validation=validation,
             error=validation.error,
+            failure=validation.failure,
         )
 
     async def _run_validation_preflight(
@@ -2567,6 +2600,9 @@ class TaskCoordinator:
         stats: ExecutionStats,
         operation_id: Optional[str],
         execution_epoch: int,
+        *,
+        failure: Optional[FailureDetail] = None,
+        worker_summary: str = "",
     ) -> None:
         def mutate(ledger: TaskLedger) -> Dict[str, Any]:
             task = _task(ledger, task_id)
@@ -2581,6 +2617,8 @@ class TaskCoordinator:
             status_map = {"stopped": AttemptStatus.STOPPED, "expired": AttemptStatus.EXPIRED}
             attempt.status = status_map.get(run_status, AttemptStatus.FAILED)
             attempt.error = error
+            attempt.failure = failure
+            attempt.worker_summary = worker_summary
             attempt.execution_stats = stats
             attempt.completed_at = utc_now()
             attempt.updated_at = utc_now()
@@ -2601,6 +2639,8 @@ class TaskCoordinator:
         stats: ExecutionStats,
         operation_id: Optional[str],
         execution_epoch: int,
+        *,
+        failure: Optional[FailureDetail] = None,
     ) -> None:
         def mutate(ledger: TaskLedger) -> Dict[str, Any]:
             task = _task(ledger, task_id)
@@ -2622,6 +2662,7 @@ class TaskCoordinator:
             }
             report.status = status_map.get(run_status, ValidationRunStatus.ERROR)
             report.error = error
+            report.failure = failure
             report.execution_stats = stats
             report.completed_at = utc_now()
             report.updated_at = utc_now()
@@ -2633,17 +2674,54 @@ class TaskCoordinator:
         await self._mutate(root_trace_id, "validation_error", mutate)
 
     @staticmethod
-    def _repair_feedback(ledger: TaskLedger, task: TaskRecord) -> Optional[Dict[str, Any]]:
-        if not task.validation_ids:
+    def _prior_feedback(
+        ledger: TaskLedger,
+        task: TaskRecord,
+        current_attempt_id: str,
+    ) -> Optional[Dict[str, Any]]:
+        previous_ids = [
+            attempt_id
+            for attempt_id in task.attempt_ids
+            if attempt_id != current_attempt_id
+        ]
+        if not previous_ids:
             return None
-        report = ledger.validations[task.validation_ids[-1]]
-        return {
-            "verdict": report.verdict.value if report.verdict else None,
-            "summary": report.summary,
-            "criterion_results": [_plain(asdict(x)) for x in report.criterion_results],
-            "risks": report.risks,
-            "recommendation": report.recommendation,
+        attempt = ledger.attempts[previous_ids[-1]]
+        report = next(
+            (
+                ledger.validations[validation_id]
+                for validation_id in reversed(task.validation_ids)
+                if ledger.validations[validation_id].attempt_id == attempt.attempt_id
+            ),
+            None,
+        )
+        feedback: Dict[str, Any] = {
+            "attempt": {
+                "attempt_id": attempt.attempt_id,
+                "status": attempt.status.value,
+                "error": attempt.error,
+                "failure": attempt.failure.to_dict() if attempt.failure else None,
+                "worker_trace_id": attempt.worker_trace_id,
+                "worker_summary": attempt.worker_summary,
+                "artifact_submitted": attempt.submission is not None,
+                "snapshot_id": attempt.snapshot_id,
+            },
         }
+        if report is not None:
+            feedback["validation"] = {
+                "validation_id": report.validation_id,
+                "status": report.status.value,
+                "verdict": report.verdict.value if report.verdict else None,
+                "summary": report.summary,
+                "criterion_results": [
+                    _plain(asdict(item)) for item in report.criterion_results
+                ],
+                "risks": report.risks,
+                "recommendation": report.recommendation,
+                "error": report.error,
+                "failure": report.failure.to_dict() if report.failure else None,
+            }
+        return feedback
 
     @staticmethod
     def _accepted_child_results(
@@ -2866,7 +2944,18 @@ def _cas_backoff(attempt_number: int) -> float:
     return min(0.005 * (2**attempt_number), 0.1)
 
 
-def _failure_code(run_status: str, *, protocol_failure: bool) -> FailureCode:
+def _failure_code(
+    run_status: str,
+    *,
+    protocol_failure: bool,
+    failure: Optional[FailureDetail] = None,
+) -> FailureCode:
+    if failure is not None:
+        if failure.code == "NO_PROGRESS":
+            return FailureCode.NO_PROGRESS
+        if failure.code == "PROTOCOL_VIOLATION":
+            return FailureCode.PROTOCOL_VIOLATION
+        return FailureCode.TOOL_FAILURE
     if protocol_failure and run_status == "completed":
         return FailureCode.PROTOCOL_VIOLATION
     return {

+ 45 - 7
agent/agent/orchestration/executor.py

@@ -11,6 +11,8 @@ from math import isfinite
 from types import MappingProxyType
 from typing import Any, Dict, Iterable, Optional, Tuple
 
+from agent.failures import FailureDetail, FailureDisposition
+
 from agent.core.runner import RunConfig
 from agent.core.knowledge_config import KnowledgeConfig
 
@@ -72,6 +74,7 @@ class LocalAgentExecutor:
             "task_spec": context["task_spec"],
             "attempt_id": context["attempt_id"],
             "repair_feedback": context.get("repair_feedback"),
+            "prior_feedback": context.get("prior_feedback"),
             "accepted_child_results": context.get("accepted_child_results", []),
             "instruction": "Execute this TaskSpec and finish with submit_attempt.",
         }
@@ -195,31 +198,41 @@ class LocalAgentExecutor:
             )
             result_trace_id = result.get("trace_id") or trace_id
             trace = await self._get_trace(result_trace_id)
+            failure = _parse_failure(result.get("failure"))
             return {
                 "trace_id": result_trace_id,
                 "status": result.get("status", "failed"),
                 "summary": result.get("summary", ""),
                 "error": result.get("error"),
+                "failure": failure,
                 "execution_stats": _execution_stats(
                     result,
                     trace,
                     usage_baseline,
                     config.model,
                     result.get("status", "failed"),
+                    failure,
                 ),
             }
         except Exception as exc:
             trace = await self._get_trace(trace_id)
+            failure = FailureDetail(
+                code="EXECUTOR_ERROR",
+                message=str(exc),
+                disposition=FailureDisposition.ABORT_RUN,
+            )
             return {
                 "trace_id": trace_id,
                 "status": "failed",
                 "error": str(exc),
+                "failure": failure,
                 "execution_stats": _execution_stats(
                     {},
                     trace,
                     usage_baseline,
                     config.model,
                     "executor_error",
+                    failure,
                 ),
             }
 
@@ -362,6 +375,7 @@ def _execution_stats(
     baseline: Tuple[int, float],
     fallback_model: str,
     run_status: str,
+    failure: Optional[FailureDetail] = None,
 ) -> ExecutionStats:
     raw = result.get("stats")
     raw = raw if isinstance(raw, dict) else {}
@@ -384,13 +398,24 @@ def _execution_stats(
         or getattr(trace, "model", None)
         or fallback_model
     )
-    failure_code = {
-        "executor_error": FailureCode.EXECUTOR_ERROR,
-        "failed": FailureCode.AGENT_FAILED,
-        "error": FailureCode.AGENT_FAILED,
-        "expired": FailureCode.TIMEOUT,
-        "stopped": FailureCode.STOPPED,
-    }.get(str(run_status))
+    if failure is not None:
+        failure_code = (
+            FailureCode.NO_PROGRESS
+            if failure.code == "NO_PROGRESS"
+            else FailureCode.EXECUTOR_ERROR
+            if failure.code == "EXECUTOR_ERROR"
+            else FailureCode.PROTOCOL_VIOLATION
+            if failure.code == "PROTOCOL_VIOLATION"
+            else FailureCode.TOOL_FAILURE
+        )
+    else:
+        failure_code = {
+            "executor_error": FailureCode.EXECUTOR_ERROR,
+            "failed": FailureCode.AGENT_FAILED,
+            "error": FailureCode.AGENT_FAILED,
+            "expired": FailureCode.TIMEOUT,
+            "stopped": FailureCode.STOPPED,
+        }.get(str(run_status))
     return ExecutionStats(
         primary_model=str(primary_model) if primary_model is not None else None,
         total_tokens=total_tokens,
@@ -399,6 +424,19 @@ def _execution_stats(
     )
 
 
+def _parse_failure(value: Any) -> Optional[FailureDetail]:
+    if value is None:
+        return None
+    try:
+        return FailureDetail.from_dict(value)
+    except (TypeError, ValueError):
+        return FailureDetail(
+            code="INVALID_FAILURE_ENVELOPE",
+            message="Runner returned an invalid failure envelope",
+            disposition=FailureDisposition.ABORT_RUN,
+        )
+
+
 def _nonnegative_int(value: Any) -> Optional[int]:
     if value is None or isinstance(value, bool):
         return None

+ 24 - 0
agent/agent/orchestration/models.py

@@ -13,6 +13,8 @@ from math import isfinite
 from typing import Any, Dict, List, Optional, Tuple
 from uuid import uuid4
 
+from agent.failures import FailureDetail
+
 
 def utc_now() -> str:
     return datetime.now(timezone.utc).isoformat()
@@ -196,6 +198,8 @@ class FailureCode(StrEnum):
     STOPPED = "stopped"
     INTERRUPTED = "interrupted"
     PROTOCOL_VIOLATION = "protocol_violation"
+    TOOL_FAILURE = "tool_failure"
+    NO_PROGRESS = "no_progress"
 
 
 @dataclass(frozen=True)
@@ -460,6 +464,8 @@ class TaskAttempt:
     submission: Optional[AttemptSubmission] = None
     execution_stats: Optional[ExecutionStats] = None
     error: Optional[str] = None
+    failure: Optional[FailureDetail] = None
+    worker_summary: str = ""
     started_at: Optional[str] = None
     completed_at: Optional[str] = None
     created_at: str = field(default_factory=utc_now)
@@ -501,6 +507,12 @@ class TaskAttempt:
                 ExecutionStats.from_dict(data["execution_stats"]) if data.get("execution_stats") is not None else None
             ),
             error=data.get("error"),
+            failure=(
+                FailureDetail.from_dict(data["failure"])
+                if data.get("failure") is not None
+                else None
+            ),
+            worker_summary=str(data.get("worker_summary", "")),
             started_at=data.get("started_at"),
             completed_at=data.get("completed_at"),
             created_at=data.get("created_at", utc_now()),
@@ -549,6 +561,7 @@ class ValidationReport:
     recommendation: str = ""
     execution_stats: Optional[ExecutionStats] = None
     error: Optional[str] = None
+    failure: Optional[FailureDetail] = None
     started_at: Optional[str] = None
     completed_at: Optional[str] = None
     created_at: str = field(default_factory=utc_now)
@@ -588,6 +601,11 @@ class ValidationReport:
                 ExecutionStats.from_dict(data["execution_stats"]) if data.get("execution_stats") is not None else None
             ),
             error=data.get("error"),
+            failure=(
+                FailureDetail.from_dict(data["failure"])
+                if data.get("failure") is not None
+                else None
+            ),
             started_at=data.get("started_at"),
             completed_at=data.get("completed_at"),
             created_at=data.get("created_at", utc_now()),
@@ -794,6 +812,7 @@ class TaskCycleResult:
     validation_id: Optional[str] = None
     validation: Optional[ValidationReport] = None
     error: Optional[str] = None
+    failure: Optional[FailureDetail] = None
 
     def to_dict(self) -> Dict[str, Any]:
         return json_values(asdict(self))
@@ -808,6 +827,11 @@ class TaskCycleResult:
             validation_id=data.get("validation_id"),
             validation=ValidationReport.from_dict(validation) if validation else None,
             error=data.get("error"),
+            failure=(
+                FailureDetail.from_dict(data["failure"])
+                if data.get("failure") is not None
+                else None
+            ),
         )
 
 

+ 4 - 0
agent/agent/orchestration/protocols.py

@@ -5,6 +5,8 @@ from __future__ import annotations
 from dataclasses import dataclass
 from typing import Any, Dict, List, Optional, Protocol, Sequence
 
+from agent.failures import FailureDetail
+
 from .models import (
     ArtifactSnapshot,
     AttemptSubmission,
@@ -30,6 +32,7 @@ class WorkerRunResult:
     status: str
     summary: str = ""
     error: Optional[str] = None
+    failure: Optional[FailureDetail] = None
     execution_stats: Optional[ExecutionStats] = None
 
 
@@ -39,6 +42,7 @@ class ValidatorRunResult:
     status: str
     summary: str = ""
     error: Optional[str] = None
+    failure: Optional[FailureDetail] = None
     execution_stats: Optional[ExecutionStats] = None
 
 

+ 17 - 0
agent/agent/orchestration/wire.py

@@ -7,6 +7,8 @@ from typing import Annotated, Any, Dict, List, Literal, Optional, Union
 
 from pydantic import BaseModel, ConfigDict, Field
 
+from agent.failures import FailureDetail
+
 from .models import (
     ArtifactSnapshot,
     BackgroundOperation,
@@ -130,6 +132,18 @@ class ExecutionStatsView(WireModel):
     failure_code: Optional[FailureCode] = None
 
 
+class FailureDetailView(WireModel):
+    code: str
+    message: str
+    disposition: str
+    source_tool: Optional[str] = None
+    details: Dict[str, Any]
+
+    @classmethod
+    def from_domain(cls, value: FailureDetail) -> "FailureDetailView":
+        return cls.model_validate(value.to_dict())
+
+
 class TaskView(WireModel):
     task_id: str
     # Deprecated compatibility field. Explicit ledgers no longer persist or
@@ -171,6 +185,8 @@ class AttemptView(WireModel):
     submission: Optional[AttemptSubmissionView] = None
     execution_stats: Optional[ExecutionStatsView] = None
     error: Optional[str] = None
+    failure: Optional[FailureDetailView] = None
+    worker_summary: str = ""
     started_at: Optional[str] = None
     completed_at: Optional[str] = None
     duration_ms: Optional[int] = None
@@ -206,6 +222,7 @@ class ValidationView(WireModel):
     recommendation: str
     execution_stats: Optional[ExecutionStatsView] = None
     error: Optional[str] = None
+    failure: Optional[FailureDetailView] = None
     started_at: Optional[str] = None
     completed_at: Optional[str] = None
     duration_ms: Optional[int] = None

+ 103 - 1
agent/tests/test_orchestration_v2_control.py

@@ -4,13 +4,14 @@ from types import SimpleNamespace
 
 import pytest
 
-from agent import AgentExecutor
+from agent import AgentExecutor, FailureDetail, FailureDisposition
 from agent.orchestration.config import OrchestrationConfig
 from agent.orchestration.coordinator import TaskCoordinator
 from agent.orchestration.executor import LocalAgentExecutor
 from agent.orchestration.models import (
     AttemptStatus,
     AttemptSubmission,
+    DecisionAction,
     ExecutionStats,
     FailureCode,
     TaskAttempt,
@@ -395,6 +396,107 @@ async def test_protocol_failure_overrides_executor_supplied_failure_code(tmp_pat
     assert attempt.execution_stats.failure_code == FailureCode.PROTOCOL_VIOLATION
 
 
+@pytest.mark.asyncio
+async def test_structured_worker_failure_persists_and_feeds_next_attempt(tmp_path):
+    failure = FailureDetail(
+        code="INPUT_SCOPE_MISMATCH",
+        message="Paragraph has no covering Structure",
+        disposition=FailureDisposition.REPLAN_TASK,
+        source_tool="save_structured_script_candidate",
+        details={"scope": "paragraph/3"},
+    )
+
+    class StructuredFailureExecutor:
+        def __init__(self):
+            self.contexts = []
+            self.coordinator = None
+
+        async def run_worker(self, context):
+            self.contexts.append(context)
+            return WorkerRunResult(
+                trace_id=context["worker_trace_id"],
+                status="failed",
+                summary="save failed before submit",
+                error=f"{failure.code}: {failure.message}",
+                failure=failure,
+                execution_stats=ExecutionStats(
+                    primary_model="worker-model",
+                    total_tokens=9,
+                    failure_code=FailureCode.TOOL_FAILURE,
+                ),
+            )
+
+        async def run_validator(self, _context):
+            raise AssertionError("validator must not run")
+
+        async def stop(self, _trace_id):
+            return True
+
+    executor = StructuredFailureExecutor()
+    coordinator, store, _ = await make_coordinator(tmp_path, executor)
+    task_id = await create_task(coordinator, "structured failure")
+
+    first = (await coordinator.dispatch_tasks("root", [task_id]))[0]
+    ledger = await store.load("root")
+    first_attempt = ledger.attempts[first.attempt_id]
+    assert first.failure == failure
+    assert first_attempt.failure == failure
+    assert first_attempt.worker_summary == "save failed before submit"
+    assert first_attempt.execution_stats.failure_code == FailureCode.TOOL_FAILURE
+
+    await coordinator.decide_task(
+        "root",
+        task_id,
+        None,
+        DecisionAction.RETRY,
+        {"reason": "retry after inspecting failure"},
+        "retry-structured-failure",
+    )
+    await coordinator.dispatch_tasks("root", [task_id])
+
+    feedback = executor.contexts[1]["prior_feedback"]
+    assert feedback["attempt"]["failure"] == failure.to_dict()
+    assert feedback["attempt"]["worker_summary"] == "save failed before submit"
+    assert feedback["attempt"]["artifact_submitted"] is False
+    assert feedback.get("validation") is None
+
+
+@pytest.mark.asyncio
+async def test_local_executor_uses_structured_failure_code_instead_of_agent_failed():
+    failure = FailureDetail(
+        code="INPUT_SCOPE_MISMATCH",
+        message="revise contract",
+        disposition=FailureDisposition.REPLAN_TASK,
+    )
+
+    class Runner:
+        trace_store = None
+
+        async def run_result(self, **_kwargs):
+            return {
+                "status": "failed",
+                "summary": "",
+                "error": "protocol_violation",
+                "failure": failure.to_dict(),
+            }
+
+    result = await LocalAgentExecutor(Runner()).run_worker(
+        {
+            "worker_preset": "worker",
+            "worker_trace_id": "worker-trace",
+            "root_trace_id": "root",
+            "task_id": "task",
+            "spec_version": 1,
+            "attempt_id": "attempt",
+            "task_spec": {},
+            "continue_trace_id": None,
+        }
+    )
+
+    assert result.failure == failure
+    assert result.execution_stats.failure_code == FailureCode.TOOL_FAILURE
+
+
 @pytest.mark.asyncio
 async def test_late_failure_cannot_overwrite_concurrent_submission(tmp_path):
     executor = FakeExecutor([])