Parcourir la source

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

从根规划器 Trace 的结构化工具结果中提取失败代码,补齐 error_history_by_code,同时保持已恢复错误不污染 current_domain_error。\n\n按根规划器 Trace 的真实起止时间统计 script_planner 耗时,并增加完整离线流程测试覆盖错误历史、耗时和成功态诊断。
SamLee il y a 20 heures
Parent
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] = {}
     preset_usage: dict[str, dict[str, int | float]] = {}
     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.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})
     traces_with_usage: set[str] = set()
     usage_reader = getattr(
-        host.composition.mission_service.runner.trace_store, "get_model_usage", None
+        trace_store, "get_model_usage", None
     )
     if callable(usage_reader):
         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["cache_read_tokens"] += int(model.get("cache_read_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]:
         task = ledger.tasks[item.task_id]
         kind = next(
@@ -670,21 +676,24 @@ async def _collect_diagnostics(host: Any, script_build_id: int) -> dict[str, Any
     tool_calls: dict[str, int] = {}
     broker_events = [
         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"
     ]
     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
             if not isinstance(content, dict):
                 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 []:
                 name = str(call.get("function", {}).get("name") or "unknown")
                 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_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
@@ -722,8 +731,8 @@ async def _collect_diagnostics(host: Any, script_build_id: int) -> dict[str, Any
             None,
         ),
         "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,
         "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]:
     sizes = sorted(int(item.get("estimated_tokens") or 0) for item in events)
     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 json
 from collections.abc import Mapping, Sequence
+from datetime import UTC, datetime, timedelta
 from pathlib import Path
 from types import SimpleNamespace
 from typing import Any, cast
@@ -10,7 +11,13 @@ from uuid import uuid4
 
 import httpx
 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 sqlalchemy import select
 
@@ -2204,6 +2211,38 @@ async def test_real_runner_phase_one_to_three_final_publication_and_legacy_readb
         prompt_tokens=123,
         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(
         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 diagnostics["preset_usage"]["script_planner"]["calls"] >= 1
     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: