#!/usr/bin/env python3
"""Audit coverage of benchmark_manifest entries across the 6 required charter classes."""

import argparse
import json
import shlex
import sys
import traceback
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Dict, Iterable, List, Optional, Tuple


@dataclass(frozen=True)
class CharterClass:
    class_id: int
    name: str


CLASSES: List[CharterClass] = [
    CharterClass(1, "Perception & Grounding"),
    CharterClass(2, "Reasoning & Inference"),
    CharterClass(3, "Planning & Control"),
    CharterClass(4, "Memory & State Tracking"),
    CharterClass(5, "Tool Use & Environment Interaction"),
    CharterClass(6, "Reliability, Safety & Maintainability"),
]

BENCHMARK_CLASS_MAP: Dict[str, int] = {
    "agi_repo_main_pycompile": 6,
    "mk2_test_imports_pycompile": 6,
    "qapla_health_pycompile": 6,
    "mk4_mk3_tool_executor_pycompile": 5,
}

SCRIPT_DIR = Path(__file__).resolve().parent
AI_RESEARCH_ROOT = SCRIPT_DIR.parent
REPO_ROOT = AI_RESEARCH_ROOT.parent

AUDIT_DIR = AI_RESEARCH_ROOT / "audit"
ALT_AUDIT_DIR = AI_RESEARCH_ROOT / "benchmarks" / "audit"
EVIDENCE_FILES = (
    AUDIT_DIR / "coverage_gap_audit_evidence.txt",
    ALT_AUDIT_DIR / "coverage_gap_audit_evidence.txt",
)
COVERAGE_GAP_MAP = AI_RESEARCH_ROOT / "benchmarks" / "COVERAGE_GAP_MAP.md"


def _dedupe_paths(paths: Iterable[Path]) -> List[Path]:
    deduped: List[Path] = []
    seen = set()
    for path in paths:
        try:
            key = str(path.expanduser().resolve())
        except OSError:
            key = str(path)
        if key not in seen:
            seen.add(key)
            deduped.append(Path(key))
    return deduped


def _normalized_parts(manifest_arg: Path) -> Tuple[str, ...]:
    return tuple(part for part in manifest_arg.parts if part not in {"", "."})


def manifest_candidates(manifest_arg: Path) -> List[Path]:
    """Return deterministic candidate paths for manifest resolution."""
    candidates: List[Path] = []
    if manifest_arg.is_absolute():
        candidates.append(manifest_arg.resolve())
    else:
        # 1) Preserve cwd-relative behavior (legacy).
        candidates.append((Path.cwd() / manifest_arg).resolve())

        parts = _normalized_parts(manifest_arg)
        if parts and parts[0] == "ai_research":
            # 2) Resolve repo-root-prefixed paths from repository root.
            candidates.append((REPO_ROOT / Path(*parts)).resolve())
            # 3) Resolve `ai_research/...` under the ai_research root.
            if len(parts) > 1:
                candidates.append((AI_RESEARCH_ROOT / Path(*parts[1:])).resolve())
        else:
            # 4) Fallback relative to ai_research root.
            candidates.append((AI_RESEARCH_ROOT / manifest_arg).resolve())
            # 5) Fallback relative to repository root.
            candidates.append((REPO_ROOT / manifest_arg).resolve())
            # 6) Explicit repo/ai_research prefix fallback.
            candidates.append((REPO_ROOT / "ai_research" / manifest_arg).resolve())

    # 7) Known manifest basename fallback for compatibility.
    if manifest_arg.name == "benchmark_manifest.json":
        candidates.append((AI_RESEARCH_ROOT / "benchmark_manifest.json").resolve())
        candidates.append((REPO_ROOT / "ai_research" / "benchmark_manifest.json").resolve())

    return _dedupe_paths(candidates)


def resolve_manifest_path(manifest_arg: Path) -> Path:
    for candidate in manifest_candidates(manifest_arg):
        if candidate.is_file():
            return candidate
    msg = "Cannot resolve benchmark manifest. Tried: " + ", ".join(str(c) for c in manifest_candidates(manifest_arg))
    raise FileNotFoundError(msg)


def load_manifest(manifest_path: Path) -> List[Dict[str, Any]]:
    raw = json.loads(manifest_path.read_text(encoding="utf-8"))
    if isinstance(raw, dict) and "benchmarks" in raw:
        entries = raw["benchmarks"]
    elif isinstance(raw, list):
        entries = raw
    else:
        raise ValueError("Unexpected benchmark_manifest format")

    if not isinstance(entries, list):
        raise ValueError("benchmark list must be a JSON array")

    return entries


def build_class_summary(
    entries: List[Dict[str, Any]],
) -> Tuple[int, Dict[str, Any], List[str], List[int], List[Dict[str, Any]]]:
    class_to_benchmarks: Dict[int, List[str]] = {c.class_id: [] for c in CLASSES}
    unmapped: List[str] = []

    for entry in entries:
        if not isinstance(entry, dict):
            unmapped.append("<missing-or-invalid-entry>")
            continue

        bid = entry.get("id")
        if not isinstance(bid, str):
            unmapped.append("<missing-or-invalid-id>")
            continue

        cid = BENCHMARK_CLASS_MAP.get(bid)
        if cid is None:
            unmapped.append(bid)
            continue
        class_to_benchmarks.setdefault(cid, []).append(bid)

    class_summary: List[Dict[str, Any]] = []
    for c in CLASSES:
        mapped = sorted(class_to_benchmarks.get(c.class_id, []))
        class_summary.append(
            {
                "class_id": c.class_id,
                "class_name": c.name,
                "mapped_benchmark_ids": mapped,
                "coverage_status": "covered" if mapped else "uncovered",
            }
        )

    covered_classes = sum(1 for row in class_summary if row["coverage_status"] == "covered")
    total_required_classes = len(CLASSES)
    uncovered_class_ids = [row["class_id"] for row in class_summary if row["coverage_status"] == "uncovered"]

    output = {
        "manifest": "",
        "manifest_benchmark_count": len(entries),
        "total_required_classes": total_required_classes,
        "covered_class_count": covered_classes,
        "uncovered_class_count": len(uncovered_class_ids),
        "covered_ratio": f"{covered_classes}/{total_required_classes}",
        "uncovered_class_ids": uncovered_class_ids,
        "class_summary": class_summary,
        "unmapped_manifest_ids": sorted(unmapped),
    }

    return covered_classes, output, unmapped, uncovered_class_ids, class_summary


def _escape_md(value: Any) -> str:
    if value is None:
        return "<missing>"
    return str(value).replace("|", "\\|")


def write_coverage_gap_map(
    manifest_path: Path,
    output: Dict[str, Any],
    entries: List[Dict[str, Any]],
) -> None:
    entry_by_id = {
        e.get("id"): e for e in entries if isinstance(e, dict) and isinstance(e.get("id"), str)
    }

    lines: List[str] = []
    lines.append("# Benchmark Coverage Gap Map (6 Charter Capability Classes)")
    lines.append("")
    lines.append("## Scope and evidence")
    lines.append(f"- Source manifest: `{manifest_path.as_posix()}`")
    lines.append("- This map file: `ai_research/benchmarks/COVERAGE_GAP_MAP.md`")
    lines.append("- Audit script: `ai_research/benchmarks/coverage_gap_audit.py`")
    lines.append(f"- Snapshot timestamp: {datetime.now(timezone.utc).isoformat().replace('+00:00', 'Z')}")

    lines.extend(["", "## Required charter classes"])
    for c in CLASSES:
        lines.append(f"{c.class_id}. {c.name}")

    lines.extend(
        [
            "",
            "## Explicit mapping from `benchmark_manifest.json` entries",
            "",
            "| benchmark_id | benchmark_name | mapped_class_id | mapped_class_name | command |",
            "|---|---|---:|---|---|",
        ]
    )

    for row in output.get("class_summary", []):
        for bid in sorted(row.get("mapped_benchmark_ids", [])):
            ent = entry_by_id.get(bid, {})
            if not isinstance(ent, dict):
                ent = {}
            name = _escape_md(ent.get("name", "<missing_name>"))
            cmd = _escape_md(ent.get("command", "<missing_command>"))
            lines.append(
                f"| `{bid}` | `{name}` | {row['class_id']} | {row['class_name']} | `{cmd}` |"
            )

    lines.extend(
        [
            "",
            "## Class-level coverage map",
            "",
            "| class_id | class_name | coverage_status | mapped_benchmark_ids | interpretation |",
            "|---|---|---|---|---|",
        ]
    )

    for row in output.get("class_summary", []):
        mapped_ids = row.get("mapped_benchmark_ids", []) or []
        if mapped_ids:
            mapped_text = "`" + "`, `".join(mapped_ids) + "`"
        else:
            mapped_text = "`[]`"

        if row.get("coverage_status") == "covered":
            interpretation = "Current manifest includes benchmark coverage for this class."
        else:
            interpretation = "No dedicated benchmark currently validates this charter class."

        lines.append(
            f"| {row.get('class_id')} | {row.get('class_name')} | {row.get('coverage_status')} | {mapped_text} | {interpretation} |"
        )

    uncovered = output.get("uncovered_class_ids", [])
    lines.extend(
        [
            "",
            "## Uncovered class IDs",
            f"Explicit uncovered class IDs: {uncovered}",
            "",
            "## Shell-executable proposals for uncovered classes",
        ]
    )

    for row in output.get("class_summary", []):
        if row.get("coverage_status") != "covered":
            lines.extend(
                [
                    f"### Shell-Executable Proposal for Class {row.get('class_id')}",
                    "```bash",
                    f"# TODO: add benchmark coverage command for class {row.get('class_id')}.",
                    f"python3 ai_research/benchmarks/coverage_gap_audit.py --manifest ai_research/benchmark_manifest.json",
                    "```",
                    "",
                ]
            )

    lines.extend(
        [
            "## Shell-checkable summary (reproducible)",
            "",
            "Run:",
            "```bash",
            "python3 ai_research/benchmarks/coverage_gap_audit.py --manifest ai_research/benchmark_manifest.json",
            "```",
            "",
            f"- `total_required_classes: {output.get('total_required_classes')}`",
            f"- `covered_class_count: {output.get('covered_class_count')}`",
            f"- `uncovered_class_count: {output.get('uncovered_class_count')}`",
            f"- `covered_ratio: {output.get('covered_ratio')}`",
            f"- `uncovered_class_ids: {output.get('uncovered_class_ids')}`",
            "",
            "Script source:",
            "- `python3 ai_research/benchmarks/coverage_gap_audit.py`",
            "",
            "## Decision quality note",
            "Only class 6 has benchmark coverage in current manifest; classes 1-5 are coverage gaps and need new benchmark entries.",
        ]
    )

    COVERAGE_GAP_MAP.parent.mkdir(parents=True, exist_ok=True)
    COVERAGE_GAP_MAP.write_text("\n".join(lines).rstrip() + "\n", encoding="utf-8")


def write_evidence(
    command: str,
    exit_code: int,
    resolved_manifest_path: Optional[Path],
    output: Optional[Dict[str, Any]] = None,
) -> None:
    resolved = str(resolved_manifest_path) if resolved_manifest_path is not None else "<unresolved>"

    lines = [
        f"timestamp={datetime.now(timezone.utc).isoformat().replace('+00:00', 'Z')}",
        f"command={command}",
        f"exit_code={exit_code}",
        f"resolved_manifest_path={resolved}",
    ]
    if isinstance(output, dict):
        if "covered_ratio" in output:
            lines.append(f"covered_ratio={output.get('covered_ratio')}")
        if "covered_class_count" in output:
            lines.append(f"covered_class_count={output.get('covered_class_count')}")
        if "uncovered_class_count" in output:
            lines.append(f"uncovered_class_count={output.get('uncovered_class_count')}")
        if "uncovered_class_ids" in output:
            lines.append(f"uncovered_class_ids={output.get('uncovered_class_ids')}")
        if "total_required_classes" in output:
            lines.append(f"total_required_classes={output.get('total_required_classes')}")
        if "manifest" in output:
            lines.append(f"manifest={output.get('manifest')}")
        if "manifest_benchmark_count" in output:
            lines.append(f"manifest_benchmark_count={output.get('manifest_benchmark_count')}")

    evidence_content = "\n".join(lines) + "\n"
    for evidence_file in EVIDENCE_FILES:
        evidence_file.parent.mkdir(parents=True, exist_ok=True)
        evidence_file.write_text(evidence_content, encoding="utf-8")


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser()
    parser.add_argument("--manifest", default="ai_research/benchmark_manifest.json", type=Path)
    parser.add_argument("--fail-on-unmapped", action="store_true")
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    command = " ".join(shlex.quote(arg) for arg in sys.argv)
    resolved_manifest_path: Optional[Path] = None
    exit_code = 0
    output: Dict[str, Any] = {}

    try:
        resolved_manifest_path = resolve_manifest_path(args.manifest)
        _, output, unmapped, _, _ = build_class_summary(load_manifest(resolved_manifest_path))
        output["manifest"] = str(resolved_manifest_path)

        if args.fail_on_unmapped and unmapped:
            exit_code = 2

        print(json.dumps(output, indent=2, sort_keys=True))
        write_coverage_gap_map(resolved_manifest_path, output, load_manifest(resolved_manifest_path))
    except Exception as exc:
        exit_code = 1
        output = {"error": str(exc), "traceback": traceback.format_exc()}
        print(str(exc), file=sys.stderr)
    finally:
        write_evidence(
            command=command,
            exit_code=exit_code,
            resolved_manifest_path=resolved_manifest_path,
            output=output,
        )

    sys.exit(exit_code)


if __name__ == "__main__":
    main()
