Ver código fonte

修正E2E报告中的规划器错误与耗时统计

从根规划器 Trace 的结构化工具结果中提取失败代码,补齐 error_history_by_code,同时保持已恢复错误不污染 current_domain_error。\n\n按根规划器 Trace 的真实起止时间统计 script_planner 耗时,并增加完整离线流程测试覆盖错误历史、耗时和成功态诊断。
SamLee 19 horas atrás
pai
commit
a3ff1d8611

+ 66 - 9
script_build_host/src/script_build_host/internal_e2e.py

@@ -596,12 +596,14 @@ async def _collect_diagnostics(host: Any, script_build_id: int) -> dict[str, Any
     task_seconds: dict[str, float] = {}
     task_seconds: dict[str, float] = {}
     preset_usage: dict[str, dict[str, int | float]] = {}
     preset_usage: dict[str, dict[str, int | float]] = {}
     domain_errors: list[dict[str, str]] = []
     domain_errors: list[dict[str, str]] = []
+    planner_tool_errors: list[dict[str, str]] = []
+    trace_store = host.composition.mission_service.runner.trace_store
     trace_presets = {binding.root_trace_id: "script_planner"}
     trace_presets = {binding.root_trace_id: "script_planner"}
     trace_presets.update({item.worker_trace_id: item.worker_preset for item in attempts})
     trace_presets.update({item.worker_trace_id: item.worker_preset for item in attempts})
     trace_presets.update({item.validator_trace_id: item.validator_preset for item in validations})
     trace_presets.update({item.validator_trace_id: item.validator_preset for item in validations})
     traces_with_usage: set[str] = set()
     traces_with_usage: set[str] = set()
     usage_reader = getattr(
     usage_reader = getattr(
-        host.composition.mission_service.runner.trace_store, "get_model_usage", None
+        trace_store, "get_model_usage", None
     )
     )
     if callable(usage_reader):
     if callable(usage_reader):
         for trace_id, preset in trace_presets.items():
         for trace_id, preset in trace_presets.items():
@@ -620,6 +622,10 @@ async def _collect_diagnostics(host: Any, script_build_id: int) -> dict[str, Any
                 usage["completion_tokens"] += int(model.get("completion_tokens") or 0)
                 usage["completion_tokens"] += int(model.get("completion_tokens") or 0)
                 usage["cache_read_tokens"] += int(model.get("cache_read_tokens") or 0)
                 usage["cache_read_tokens"] += int(model.get("cache_read_tokens") or 0)
                 usage["tokens"] += int(model.get("total_tokens") or 0)
                 usage["tokens"] += int(model.get("total_tokens") or 0)
+    planner_trace = await trace_store.get_trace(binding.root_trace_id)
+    preset_usage.setdefault("script_planner", _empty_preset_usage())["seconds"] += (
+        _trace_elapsed_seconds(planner_trace)
+    )
     for item in [*attempts, *validations]:
     for item in [*attempts, *validations]:
         task = ledger.tasks[item.task_id]
         task = ledger.tasks[item.task_id]
         kind = next(
         kind = next(
@@ -670,21 +676,24 @@ async def _collect_diagnostics(host: Any, script_build_id: int) -> dict[str, Any
     tool_calls: dict[str, int] = {}
     tool_calls: dict[str, int] = {}
     broker_events = [
     broker_events = [
         item
         item
-        for item in await host.composition.mission_service.runner.trace_store.get_events(
-            binding.root_trace_id
-        )
+        for item in await trace_store.get_events(binding.root_trace_id)
         if item.get("event") == "context_broker_receipt"
         if item.get("event") == "context_broker_receipt"
     ]
     ]
     for trace_id in trace_ids:
     for trace_id in trace_ids:
-        for message in await host.composition.mission_service.runner.trace_store.get_trace_messages(
-            trace_id
-        ):
+        for message in await trace_store.get_trace_messages(trace_id):
             content = message.content
             content = message.content
             if not isinstance(content, dict):
             if not isinstance(content, dict):
                 continue
                 continue
+            if trace_id == binding.root_trace_id:
+                failure = _tool_failure_from_message(message)
+                if failure is not None:
+                    if not failure["task_id"]:
+                        failure["task_id"] = root.task_id
+                    planner_tool_errors.append(failure)
             for call in content.get("tool_calls") or []:
             for call in content.get("tool_calls") or []:
                 name = str(call.get("function", {}).get("name") or "unknown")
                 name = str(call.get("function", {}).get("name") or "unknown")
                 tool_calls[name] = tool_calls.get(name, 0) + 1
                 tool_calls[name] = tool_calls.get(name, 0) + 1
+    historical_errors = [*domain_errors, *planner_tool_errors]
     last_task = max(ledger.tasks.values(), key=lambda item: item.updated_at)
     last_task = max(ledger.tasks.values(), key=lambda item: item.updated_at)
     last_attempt = max(attempts, key=lambda item: item.updated_at) if attempts else None
     last_attempt = max(attempts, key=lambda item: item.updated_at) if attempts else None
     last_validation = max(validations, key=lambda item: item.updated_at) if validations else None
     last_validation = max(validations, key=lambda item: item.updated_at) if validations else None
@@ -722,8 +731,8 @@ async def _collect_diagnostics(host: Any, script_build_id: int) -> dict[str, Any
             None,
             None,
         ),
         ),
         "error_history_by_code": {
         "error_history_by_code": {
-            code: sum(item["code"] == code for item in domain_errors)
-            for code in sorted({item["code"] for item in domain_errors})
+            code: sum(item["code"] == code for item in historical_errors)
+            for code in sorted({item["code"] for item in historical_errors})
         },
         },
         "phase_timing_seconds": task_seconds,
         "phase_timing_seconds": task_seconds,
         "preset_usage": preset_usage,
         "preset_usage": preset_usage,
@@ -739,6 +748,54 @@ async def _collect_diagnostics(host: Any, script_build_id: int) -> dict[str, Any
     }
     }
 
 
 
 
+def _tool_failure_from_message(message: Any) -> dict[str, str] | None:
+    if getattr(message, "role", None) != "tool" or not isinstance(message.content, dict):
+        return None
+    result = message.content.get("result")
+    if isinstance(result, str):
+        try:
+            payload = json.loads(result)
+        except (TypeError, ValueError):
+            return None
+    elif isinstance(result, dict):
+        payload = result
+    else:
+        return None
+    if not isinstance(payload, dict):
+        return None
+    failure = payload.get("failure")
+    if not isinstance(failure, dict):
+        control = payload.get("_control")
+        failure = control.get("failure") if isinstance(control, dict) else None
+    if not isinstance(failure, dict):
+        return None
+    code = str(failure.get("code") or "").strip()
+    if not code:
+        return None
+    details = failure.get("details")
+    task_id = str(details.get("task_id") or "") if isinstance(details, dict) else ""
+    return {
+        "code": code,
+        "message": str(failure.get("message") or ""),
+        "task_id": task_id,
+        "source_tool": str(
+            failure.get("source_tool") or message.content.get("tool_name") or ""
+        ),
+    }
+
+
+def _trace_elapsed_seconds(trace: Any) -> float:
+    if trace is None:
+        return 0.0
+    started_at = getattr(trace, "created_at", None)
+    completed_at = getattr(trace, "completed_at", None) or getattr(
+        trace, "last_activity_at", None
+    )
+    if isinstance(started_at, datetime) and isinstance(completed_at, datetime):
+        return max((completed_at - started_at).total_seconds(), 0.0)
+    return max(float(getattr(trace, "total_duration_ms", 0) or 0) / 1000, 0.0)
+
+
 def _context_broker_diagnostics(events: list[dict[str, Any]]) -> dict[str, Any]:
 def _context_broker_diagnostics(events: list[dict[str, Any]]) -> dict[str, Any]:
     sizes = sorted(int(item.get("estimated_tokens") or 0) for item in events)
     sizes = sorted(int(item.get("estimated_tokens") or 0) for item in events)
     detail_pages = [
     detail_pages = [

+ 42 - 1
script_build_host/tests/test_phase_two_real_runner_e2e.py

@@ -3,6 +3,7 @@ from __future__ import annotations
 import asyncio
 import asyncio
 import json
 import json
 from collections.abc import Mapping, Sequence
 from collections.abc import Mapping, Sequence
+from datetime import UTC, datetime, timedelta
 from pathlib import Path
 from pathlib import Path
 from types import SimpleNamespace
 from types import SimpleNamespace
 from typing import Any, cast
 from typing import Any, cast
@@ -10,7 +11,13 @@ from uuid import uuid4
 
 
 import httpx
 import httpx
 import pytest
 import pytest
-from agent import AgentRunner, FileSystemArtifactStore, FileSystemTaskStore, FileSystemTraceStore
+from agent import (
+    AgentRunner,
+    FileSystemArtifactStore,
+    FileSystemTaskStore,
+    FileSystemTraceStore,
+    Message,
+)
 from agent.orchestration import ArtifactRef, DecisionAction, TaskStatus, ValidationVerdict
 from agent.orchestration import ArtifactRef, DecisionAction, TaskStatus, ValidationVerdict
 from sqlalchemy import select
 from sqlalchemy import select
 
 
@@ -2204,6 +2211,38 @@ async def test_real_runner_phase_one_to_three_final_publication_and_legacy_readb
         prompt_tokens=123,
         prompt_tokens=123,
         completion_tokens=7,
         completion_tokens=7,
     )
     )
+    planner_trace = await trace_store.get_trace(ROOT)
+    assert planner_trace is not None
+    await trace_store.add_message(
+        Message.create(
+            trace_id=ROOT,
+            role="tool",
+            sequence=planner_trace.last_sequence + 1,
+            parent_sequence=planner_trace.head_sequence,
+            tool_call_id="diagnostic-planner-failure",
+            content={
+                "tool_name": "plan_script_tasks",
+                "result": json.dumps(
+                    {
+                        "failure": {
+                            "code": "INPUT_SCOPE_MISMATCH",
+                            "message": "scope is outside the active direction",
+                            "disposition": "replan_task",
+                            "source_tool": "plan_script_tasks",
+                            "details": {"task_id": root.task_id},
+                        }
+                    }
+                ),
+            },
+        )
+    )
+    planner_started_at = datetime(2026, 7, 22, tzinfo=UTC)
+    await trace_store.update_trace(
+        ROOT,
+        created_at=planner_started_at,
+        completed_at=planner_started_at + timedelta(seconds=12.5),
+        last_activity_at=planner_started_at + timedelta(seconds=12.5),
+    )
     diagnostics = await _collect_diagnostics(
     diagnostics = await _collect_diagnostics(
         SimpleNamespace(composition=composition), script_build_id
         SimpleNamespace(composition=composition), script_build_id
     )
     )
@@ -2211,6 +2250,8 @@ async def test_real_runner_phase_one_to_three_final_publication_and_legacy_readb
     assert "last_domain_error" not in diagnostics
     assert "last_domain_error" not in diagnostics
     assert diagnostics["preset_usage"]["script_planner"]["calls"] >= 1
     assert diagnostics["preset_usage"]["script_planner"]["calls"] >= 1
     assert diagnostics["preset_usage"]["script_planner"]["tokens"] >= 130
     assert diagnostics["preset_usage"]["script_planner"]["tokens"] >= 130
+    assert diagnostics["preset_usage"]["script_planner"]["seconds"] == pytest.approx(12.5)
+    assert diagnostics["error_history_by_code"]["INPUT_SCOPE_MISMATCH"] == 1
 
 
 
 
 def _phase_two_policy_count(messages: Sequence[Any]) -> int:
 def _phase_two_policy_count(messages: Sequence[Any]) -> int: