125 lines
4.5 KiB
Python
125 lines
4.5 KiB
Python
|
|
"""Report generator skill v1 (docs/05 §2.4).
|
||
|
|
|
||
|
|
Produces the *numeric skeleton* of a report: every figure is a
|
||
|
|
``{tool_call_id, path}`` reference into a recorded tool output, with the value
|
||
|
|
copied from that path — never typed, never computed here. The LLM's role
|
||
|
|
(docs/07 D-1 08:00 step 3) is prose around these references and arrives with
|
||
|
|
the runtime in M3; this skill is what makes "数字原样引用求解器输出" checkable.
|
||
|
|
"""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import hashlib
|
||
|
|
import json
|
||
|
|
import re
|
||
|
|
from datetime import datetime, timezone
|
||
|
|
from typing import Any, Callable
|
||
|
|
|
||
|
|
from vpp_contracts.report_request import ReportRequest
|
||
|
|
from vpp_contracts.skill_report import SkillReport
|
||
|
|
|
||
|
|
from . import SKILL_VERSIONS
|
||
|
|
|
||
|
|
SKILL = "report-generator"
|
||
|
|
DECIMAL_RE = re.compile(r"^-?\d+(\.\d+)?$")
|
||
|
|
|
||
|
|
UNIT_SUFFIXES = (
|
||
|
|
("_yuan_per_mwh", "yuan_per_mwh"),
|
||
|
|
("_mwh", "mwh"),
|
||
|
|
("_mw", "mw"),
|
||
|
|
("_yuan", "yuan"),
|
||
|
|
("_pct", "pct"),
|
||
|
|
("_ms", "ms"),
|
||
|
|
)
|
||
|
|
|
||
|
|
SECTION_TITLE = {
|
||
|
|
"DAY_AHEAD_BID_SUMMARY": "Day-ahead bid",
|
||
|
|
"FORECAST_EVAL": "Forecast evaluation",
|
||
|
|
"BID_BACKTEST": "Bid backtest",
|
||
|
|
}
|
||
|
|
|
||
|
|
|
||
|
|
def resolve(output: Any, path: str) -> Any:
|
||
|
|
"""Dereference a dotted path (list indices as integers) into a tool output."""
|
||
|
|
node = output
|
||
|
|
for part in path.split("."):
|
||
|
|
if isinstance(node, list):
|
||
|
|
node = node[int(part)]
|
||
|
|
elif isinstance(node, dict):
|
||
|
|
node = node[part]
|
||
|
|
else:
|
||
|
|
raise KeyError(path)
|
||
|
|
return node
|
||
|
|
|
||
|
|
|
||
|
|
def _unit_for(name: str, parent: str) -> str:
|
||
|
|
for suffix, unit in UNIT_SUFFIXES:
|
||
|
|
if name.endswith(suffix) or parent.endswith(suffix):
|
||
|
|
return unit
|
||
|
|
return "1"
|
||
|
|
|
||
|
|
|
||
|
|
def _numeric_leaves(node: Any, prefix: str = "", parent: str = "") -> list[tuple[str, str, str]]:
|
||
|
|
"""(path, value, unit) for every decimal-string scalar; curves/lists are skipped
|
||
|
|
(a 96-point curve is cited by snapshot ref, not inlined into a report)."""
|
||
|
|
out: list[tuple[str, str, str]] = []
|
||
|
|
if isinstance(node, dict):
|
||
|
|
for key, val in node.items():
|
||
|
|
path = f"{prefix}.{key}" if prefix else key
|
||
|
|
if isinstance(val, str) and DECIMAL_RE.match(val):
|
||
|
|
out.append((path, val, _unit_for(key, parent)))
|
||
|
|
elif isinstance(val, dict) and "values" not in val:
|
||
|
|
out.extend(_numeric_leaves(val, path, key))
|
||
|
|
return out
|
||
|
|
|
||
|
|
|
||
|
|
def generate_report(
|
||
|
|
req: ReportRequest,
|
||
|
|
clock: Callable[[], datetime] = lambda: datetime.now(timezone.utc),
|
||
|
|
) -> SkillReport:
|
||
|
|
kind = str(req.kind.value)
|
||
|
|
sections = []
|
||
|
|
for src in req.sources:
|
||
|
|
metrics = [
|
||
|
|
{"name": path, "value": value, "unit": unit, "ref": {"tool_call_id": src.tool_call_id, "path": path}}
|
||
|
|
for path, value, unit in _numeric_leaves(src.output)
|
||
|
|
]
|
||
|
|
notes = [f"source: {src.tool} v{src.version} (tool call {src.tool_call_id})"]
|
||
|
|
if isinstance(src.output, dict):
|
||
|
|
solver = src.output.get("solver")
|
||
|
|
if isinstance(solver, dict) and "status" in solver:
|
||
|
|
notes.append(f"solver status: {solver['status']}")
|
||
|
|
for key in ("binding_constraints",):
|
||
|
|
if isinstance(src.output.get(key), list) and src.output[key]:
|
||
|
|
notes.append(f"{key}: {', '.join(map(str, src.output[key]))}")
|
||
|
|
sections.append({"title": f"{SECTION_TITLE[kind]} · {src.tool}", "metrics": metrics, "notes": notes})
|
||
|
|
|
||
|
|
digest = hashlib.sha256(
|
||
|
|
json.dumps([s.model_dump() for s in req.sources], sort_keys=True, separators=(",", ":")).encode()
|
||
|
|
).hexdigest()[:12]
|
||
|
|
report = {
|
||
|
|
"id": f"rep-{kind.lower()}-{req.market_date}-{digest}",
|
||
|
|
"kind": kind,
|
||
|
|
"market_date": req.market_date,
|
||
|
|
"sections": sections,
|
||
|
|
"skill_version": SKILL_VERSIONS[SKILL],
|
||
|
|
"generated_at": clock().isoformat().replace("+00:00", "Z"),
|
||
|
|
}
|
||
|
|
return SkillReport.model_validate(report)
|
||
|
|
|
||
|
|
|
||
|
|
def verify_report(report: SkillReport, req: ReportRequest) -> list[str]:
|
||
|
|
"""Judge (docs/12 L1 数字一致性): every metric value must equal the referenced output."""
|
||
|
|
by_id = {s.tool_call_id: s.output for s in req.sources}
|
||
|
|
problems = []
|
||
|
|
for section in report.sections:
|
||
|
|
for m in section.metrics:
|
||
|
|
try:
|
||
|
|
actual = resolve(by_id[m.ref.tool_call_id], m.ref.path)
|
||
|
|
except (KeyError, IndexError, ValueError):
|
||
|
|
problems.append(f"{m.name}: unresolvable ref {m.ref.tool_call_id}#{m.ref.path}")
|
||
|
|
continue
|
||
|
|
if str(actual) != m.value:
|
||
|
|
problems.append(f"{m.name}: value {m.value} != referenced {actual}")
|
||
|
|
return problems
|