#!/usr/bin/env python3
"""Run only campaign (a): 100 ILS iterations, 300 second wall timeout."""

from __future__ import annotations

import argparse
import datetime as dt
import hashlib
import json
import os
import subprocess
import sys
import tempfile
import time
from pathlib import Path
from typing import Any

from public_benchmark_independent_verifier import IndependentVerificationError, verify
from public_benchmark_reader import BenchmarkFormatError, canonical_problem, load_instance


EXPECTED_NAMES = [
    "eil51", "berlin52", "eil76", "kroA100", "eil101", "ch150", "a280",
    "A-n32-k5", "A-n44-k6", "A-n60-k9", "A-n80-k10", "B-n50-k7",
    "P-n76-k4", "X-n101-k25", "X-n200-k36",
]


def canonical_input_active_families(library: str) -> list[str]:
    if library == "TSPLIB":
        return ["mandatory_jobs"]
    if library == "CVRPLIB":
        return [
            "multiple_vehicles",
            "multidimensional_capacity_and_demand",
            "delivery_and_pickup_load_propagation",
            "mandatory_jobs",
        ]
    raise RuntimeError(f"unsupported benchmark library {library!r}")


def canonical_input_benchmark_constraints(library: str) -> list[str]:
    constraints = ["single_depot", "closed_routes", "all_jobs_mandatory", "exact_once_service"]
    if library == "CVRPLIB":
        constraints.extend(["exact_vehicle_supply_k", "single_dimension_capacity_Q"])
    return constraints


def sha256_file(file_path: Path) -> str:
    digest = hashlib.sha256()
    with file_path.open("rb") as handle:
        for chunk in iter(lambda: handle.read(1024 * 1024), b""):
            digest.update(chunk)
    return digest.hexdigest()


def source_inventory(config: dict[str, Any]) -> list[dict[str, Any]]:
    unique: dict[tuple[str, str, str], dict[str, Any]] = {}
    for spec in config["instances"]:
        for role in ("instance", "reference"):
            item = dict(spec[role])
            item["role"] = role
            key = (item["file"], item["url"], item["sha256"])
            unique[key] = item
    for source in config.get("auxiliarySources", []):
        item = dict(source)
        key = (item["file"], item["url"], item["sha256"])
        unique[key] = item
    return sorted(unique.values(), key=lambda item: (item["file"], item["role"]))


def normalize_result_rows(
    rows: list[dict[str, Any]], repo_root: Path
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    runner_manifest_statuses = {"VALID_GAP", "ENGINE_ERROR"}
    manifests = [
        row["engineFiles"]
        for row in rows
        if row.get("engineFiles") and row.get("status") in runner_manifest_statuses
    ]
    if not manifests:
        raise RuntimeError("no completed execution returned the actual engine require-cache manifest")
    canonical_manifest = manifests[0]
    if any(manifest != canonical_manifest for manifest in manifests[1:]):
        raise RuntimeError("engine file manifests diverged between executions")

    git_heads = {row["engineGitHead"] for row in rows if row.get("engineGitHead")}
    identities = [row["engineProductionIdentity"] for row in rows if row.get("engineProductionIdentity")]
    solver_ids = {row["solverId"] for row in rows if row.get("solverId")}
    solver_versions = {row["solverVersion"] for row in rows if row.get("solverVersion")}
    if len(git_heads) != 1 or len({json.dumps(item, sort_keys=True) for item in identities}) != 1:
        raise RuntimeError("engine Git head or production identity diverged between executions")
    if len(solver_ids) != 1 or len(solver_versions) != 1:
        raise RuntimeError("solver id or version diverged between executions")

    for item in canonical_manifest:
        engine_file = repo_root / item["path"]
        actual_sha256 = sha256_file(engine_file)
        if actual_sha256 != item["sha256"]:
            raise RuntimeError(
                f"engine SHA-256 changed after execution: {item['path']} "
                f"recorded={item['sha256']} actual={actual_sha256}"
            )

    canonical_git_head = next(iter(git_heads))
    canonical_identity = identities[0]
    canonical_solver_id = next(iter(solver_ids))
    canonical_solver_version = next(iter(solver_versions))
    normalized: list[dict[str, Any]] = []
    for original in rows:
        row = dict(original)
        runner_manifest_present = bool(row.get("engineFiles")) and row.get("status") in runner_manifest_statuses
        row["engineFiles"] = canonical_manifest
        row["engineFilesEvidence"] = (
            "runner_require_cache"
            if runner_manifest_present
            else "identical_campaign_executable_environment_verified_against_disk_after_timeout"
        )
        row["engineGitHead"] = canonical_git_head
        row["engineProductionIdentity"] = canonical_identity
        row["solverId"] = canonical_solver_id
        row["solverVersion"] = canonical_solver_version
        row["activeFamilies"] = canonical_input_active_families(row["library"])
        row["activeFamiliesEvidence"] = "canonical_problem_input"
        row["benchmarkConstraints"] = canonical_input_benchmark_constraints(row["library"])
        normalized.append(row)

    engine_environment = {
        "gitHead": canonical_git_head,
        "productionIdentity": canonical_identity,
        "solverId": canonical_solver_id,
        "solverVersion": canonical_solver_version,
        "files": canonical_manifest,
        "manifestReturnedByExecutions": len(manifests),
        "manifestIdenticalAcrossReturnedExecutions": True,
        "allFilesRehashedAgainstCurrentDiskAfterCampaign": True,
    }
    return normalized, engine_environment


def utc_now() -> str:
    return dt.datetime.now(dt.timezone.utc).isoformat().replace("+00:00", "Z")


def atomic_json(path: Path, payload: dict[str, Any]) -> None:
    temporary = path.with_suffix(path.suffix + ".tmp")
    temporary.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n", encoding="utf-8")
    os.replace(temporary, path)


def summary(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
    output = []
    for name in EXPECTED_NAMES:
        selected = [row for row in rows if row["instance"] == name]
        completed = [row for row in selected if row["status"] == "VALID_GAP"]
        gap_per_seed: list[dict[str, Any]] = []
        for seed in (11, 29):
            seed_rows = [row for row in completed if row["seed"] == seed]
            gaps = sorted(set(row["gapPercent"] for row in seed_rows))
            costs = sorted(set(row["cost"] for row in seed_rows))
            fingerprints = sorted(set(row["solutionFingerprint"] for row in seed_rows))
            gap_per_seed.append({
                "seed": seed,
                "gapPercent": gaps[0] if len(gaps) == 1 else None,
                "cost": costs[0] if len(costs) == 1 else None,
                "solutionFingerprint": fingerprints[0] if len(fingerprints) == 1 else None,
                "repetitionsObserved": len(seed_rows),
                "repeatResultIdentical": len(seed_rows) < 2 or (
                    len(gaps) == 1 and len(costs) == 1 and len(fingerprints) == 1
                ),
            })
        seconds = [row["seconds"] for row in selected if isinstance(row.get("seconds"), (int, float))]
        output.append({
            "instance": name,
            "rows": len(selected),
            "statusCounts": {
                status: sum(1 for row in selected if row["status"] == status)
                for status in sorted(set(row["status"] for row in selected))
            },
            "gapPerSeed": gap_per_seed,
            "secondsMinimumAcrossFourExecutions": min(seconds) if seconds else None,
            "secondsMaximumAcrossFourExecutions": max(seconds) if seconds else None,
        })
    return output


def build_document(config: dict[str, Any], rows: list[dict[str, Any]], started_at: str) -> dict[str, Any]:
    return {
        "contract": "routelogic_public_quality_gap_results_v1",
        "campaignId": config["campaignId"],
        "configuration": config["configuration"],
        "startedAtUtc": started_at,
        "updatedAtUtc": utc_now(),
        "completedAtUtc": utc_now() if len(rows) == 60 else None,
        "requiredRowCount": 60,
        "actualRowCount": len(rows),
        "campaignComplete": len(rows) == 60,
        "configurationBStarted": False,
        "rows": rows,
        "summary": summary(rows),
    }


def finalize_existing_document(
    config: dict[str, Any], document: dict[str, Any], repo_root: Path
) -> dict[str, Any]:
    if document.get("actualRowCount") != 60 or document.get("campaignComplete") is not True:
        raise RuntimeError("only a complete 60-row campaign (a) may be finalized")
    if document.get("configurationBStarted") is not False:
        raise RuntimeError("configuration (b) must not have started")
    rows, engine_environment = normalize_result_rows(document["rows"], repo_root)
    finalized = dict(document)
    finalized["rows"] = rows
    finalized["summary"] = summary(rows)
    finalized["sourceFiles"] = source_inventory(config)
    finalized["engineEnvironment"] = engine_environment
    finalized["finalizedAtUtc"] = utc_now()
    finalized["finalization"] = {
        "method": "metadata_only_no_solver_rerun",
        "measurementRowsChanged": False,
        "timeoutRowsRetained": True,
        "engineErrorRowsRetained": True,
        "configurationBStarted": False,
    }
    return finalized


def validate_config(config: dict[str, Any]) -> None:
    settings = config.get("configuration", {})
    if settings.get("iterationLimit") != 100 or settings.get("timeoutSeconds") != 300:
        raise RuntimeError("this program is locked to campaign (a): 100 iterations, 300 seconds")
    if settings.get("seeds") != [11, 29] or settings.get("repetitionsPerSeed") != 2:
        raise RuntimeError("campaign (a) requires seeds 11/29 and two repetitions")
    names = [item["name"] for item in config.get("instances", [])]
    if names != EXPECTED_NAMES:
        raise RuntimeError(f"locked instance list/order mismatch: {names}")


def row_base(spec: dict[str, Any], seed: int, repetition: int, config: dict[str, Any]) -> dict[str, Any]:
    return {
        "campaignId": config["campaignId"],
        "configuration": "a",
        "instance": spec["name"],
        "library": spec["library"],
        "n": spec["n"],
        "k": spec["k"],
        "referenceValue": spec["referenceValue"],
        "referenceType": spec["referenceType"],
        "iterationLimit": 100,
        "timeoutSeconds": 300,
        "seed": seed,
        "repetition": repetition,
        "instanceSource": spec["instance"],
        "referenceSource": spec["reference"],
        "activeFamilies": canonical_input_active_families(spec["library"]),
        "activeFamiliesEvidence": "canonical_problem_input",
        "benchmarkConstraints": canonical_input_benchmark_constraints(spec["library"]),
    }


def run_case(
    root: Path,
    node_runner: Path,
    instance: dict[str, Any],
    spec: dict[str, Any],
    seed: int,
    repetition: int,
    config: dict[str, Any],
    temp_dir: Path,
) -> dict[str, Any]:
    row = row_base(spec, seed, repetition, config)
    problem = canonical_problem(instance, seed, config["campaignId"])
    input_path = temp_dir / f"{spec['name']}-{seed}-{repetition}.json"
    input_path.write_text(json.dumps({"problem": problem, "iterationLimit": 100}), encoding="utf-8")
    started_at = utc_now()
    started = time.monotonic()
    try:
        completed = subprocess.run(
            ["node", str(node_runner), str(input_path)],
            cwd=root,
            stdin=subprocess.DEVNULL,
            stdout=subprocess.PIPE,
            stderr=subprocess.PIPE,
            text=True,
            timeout=300,
            check=False,
        )
    except subprocess.TimeoutExpired as error:
        elapsed = time.monotonic() - started
        row.update({
            "status": "TIMEOUT",
            "startedAtUtc": started_at,
            "finishedAtUtc": utc_now(),
            "seconds": elapsed,
            "iterations": None,
            "cost": None,
            "gapPercent": None,
            "solutionFingerprint": None,
            "engineGitHead": None,
            "engineFiles": None,
            "error": {"code": "TIMEOUT", "message": f"process exceeded 300 seconds", "stderr": error.stderr},
        })
        return row

    elapsed = time.monotonic() - started
    try:
        runner = json.loads(completed.stdout.strip())
    except json.JSONDecodeError:
        row.update({
            "status": "PROCESS_ERROR",
            "startedAtUtc": started_at,
            "finishedAtUtc": utc_now(),
            "seconds": elapsed,
            "iterations": None,
            "cost": None,
            "gapPercent": None,
            "solutionFingerprint": None,
            "engineGitHead": None,
            "engineFiles": None,
            "error": {
                "code": "INVALID_RUNNER_OUTPUT",
                "message": "node runner did not emit one JSON object",
                "returnCode": completed.returncode,
                "stdout": completed.stdout[-4000:],
                "stderr": completed.stderr[-4000:],
            },
        })
        return row

    result = runner.get("result") or {}
    common = {
        "startedAtUtc": started_at,
        "finishedAtUtc": utc_now(),
        "seconds": elapsed,
        "solverElapsedSeconds": runner.get("elapsedSeconds"),
        "engineGitHead": runner.get("gitHead"),
        "engineProductionIdentity": runner.get("productionIdentity"),
        "engineFiles": runner.get("engineFiles"),
        "solverId": result.get("solverId"),
        "solverVersion": result.get("solverVersion"),
    }
    if runner.get("runnerStatus") != "COMPLETED" or completed.returncode != 0:
        row.update(common)
        row.update({
            "status": "PROCESS_ERROR",
            "iterations": None,
            "cost": None,
            "gapPercent": None,
            "solutionFingerprint": None,
            "error": runner.get("error") or {
                "code": "RUNNER_NONZERO", "returnCode": completed.returncode, "stderr": completed.stderr[-4000:]
            },
        })
        return row
    if result.get("status") != "FEASIBLE":
        row.update(common)
        row.update({
            "status": "ENGINE_ERROR",
            "iterations": (result.get("telemetry") or {}).get("iterations"),
            "cost": None,
            "gapPercent": None,
            "solutionFingerprint": None,
            "error": result.get("error"),
        })
        return row

    verified = verify(instance, problem, result)
    cost = verified["recomputedDistance"]
    gap = (cost - spec["referenceValue"]) * 100.0 / spec["referenceValue"]
    row.update(common)
    row.update({
        "status": "VALID_GAP",
        "iterations": (result.get("telemetry") or {}).get("iterations"),
        "cost": cost,
        "gapPercent": gap,
        "solutionFingerprint": result.get("solutionHash"),
        "independentSolutionFingerprint": verified["independentSolutionFingerprint"],
        "independentVerification": verified,
        "error": None,
    })
    return row


def assert_repeat_identity(rows: list[dict[str, Any]], instance: str, seed: int) -> None:
    selected = [
        row for row in rows
        if row["instance"] == instance and row["seed"] == seed and row["status"] == "VALID_GAP"
    ]
    if len(selected) != 2:
        return
    fields = ("cost", "gapPercent", "solutionFingerprint", "independentSolutionFingerprint")
    divergence = {field: [row[field] for row in selected] for field in fields if selected[0][field] != selected[1][field]}
    if divergence:
        raise IndependentVerificationError(
            "FIXED_SEED_REPEAT_DIVERGENCE",
            f"{instance} seed {seed} produced different solutions across repetitions",
            divergence,
        )


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", type=Path, default=Path(__file__).with_name("campaign-a-config.json"))
    parser.add_argument("--output", type=Path, default=Path(__file__).with_name("campaign-a-results.json"))
    parser.add_argument("--finalize-existing", action="store_true")
    args = parser.parse_args()
    config_path = args.config.resolve()
    output_path = args.output.resolve()
    benchmark_root = config_path.parent
    repo_root = benchmark_root.parents[2]
    node_runner = benchmark_root / "public_quality_case.js"
    config = json.loads(config_path.read_text(encoding="utf-8"))
    validate_config(config)

    if args.finalize_existing:
        document = json.loads(output_path.read_text(encoding="utf-8"))
        finalized = finalize_existing_document(config, document, repo_root)
        atomic_json(output_path, finalized)
        print(
            f"CAMPAIGN_A_FINALIZED rows={finalized['actualRowCount']}/60 "
            f"engineFiles={len(finalized['engineEnvironment']['files'])} "
            "configurationBStarted=false solverRerun=false",
            flush=True,
        )
        return

    print("CAMPAIGN_A_LOCKED instances=15 iterations=100 timeoutSeconds=300 seeds=11,29 repetitions=2", flush=True)
    loaded: dict[str, dict[str, Any]] = {}
    try:
        for spec in config["instances"]:
            loaded[spec["name"]] = load_instance(benchmark_root, spec)
            item = loaded[spec["name"]]
            print(
                f"PREFLIGHT instance={item['name']} library={item['library']} n={item['n']} "
                f"k={item['k']} edgeWeightType=EUC_2D reference={item['referenceValue']} "
                f"referenceType={item['referenceType']} PASS",
                flush=True,
            )
    except BenchmarkFormatError as error:
        print(f"PREFLIGHT_FAIL {error}", flush=True)
        raise SystemExit(2)

    rows: list[dict[str, Any]] = []
    started_at = utc_now()
    atomic_json(output_path, build_document(config, rows, started_at))
    total = len(config["instances"]) * 4
    with tempfile.TemporaryDirectory(prefix="routelogic-public-gap-a-") as temporary:
        temp_dir = Path(temporary)
        for spec in config["instances"]:
            instance = loaded[spec["name"]]
            for seed in config["configuration"]["seeds"]:
                for repetition in range(config["configuration"]["repetitionsPerSeed"]):
                    ordinal = len(rows) + 1
                    print(
                        f"RUN {ordinal:02d}/{total} instance={spec['name']} seed={seed} repetition={repetition} START",
                        flush=True,
                    )
                    try:
                        row = run_case(
                            repo_root, node_runner, instance, spec, seed, repetition, config, temp_dir
                        )
                    except IndependentVerificationError as error:
                        fatal = row_base(spec, seed, repetition, config)
                        fatal.update({
                            "status": "INVALID_STOP",
                            "seconds": None,
                            "iterations": None,
                            "cost": None,
                            "gapPercent": None,
                            "solutionFingerprint": None,
                            "error": {"code": error.code, "message": str(error), "details": error.details},
                        })
                        rows.append(fatal)
                        atomic_json(output_path, build_document(config, rows, started_at))
                        print(
                            f"RUN {ordinal:02d}/{total} instance={spec['name']} seed={seed} "
                            f"repetition={repetition} INVALID_STOP code={error.code}",
                            flush=True,
                        )
                        raise SystemExit(3)

                    rows.append(row)
                    atomic_json(output_path, build_document(config, rows, started_at))
                    if row["status"] == "VALID_GAP":
                        print(
                            f"RUN {ordinal:02d}/{total} instance={spec['name']} seed={seed} "
                            f"repetition={repetition} VALID_GAP cost={row['cost']} "
                            f"gapPercent={row['gapPercent']:.9f} seconds={row['seconds']:.6f} "
                            f"iterations={row['iterations']} fingerprint={row['solutionFingerprint']}",
                            flush=True,
                        )
                    else:
                        code = (row.get("error") or {}).get("code", "UNKNOWN")
                        print(
                            f"RUN {ordinal:02d}/{total} instance={spec['name']} seed={seed} "
                            f"repetition={repetition} status={row['status']} code={code} "
                            f"seconds={row['seconds']:.6f}",
                            flush=True,
                        )
                    assert_repeat_identity(rows, spec["name"], seed)

    document = finalize_existing_document(
        config,
        build_document(config, rows, started_at),
        repo_root,
    )
    atomic_json(output_path, document)
    print(
        f"CAMPAIGN_A_COMPLETE rows={len(rows)}/{total} output={output_path} configurationBStarted=false",
        flush=True,
    )


if __name__ == "__main__":
    main()
