#!/usr/bin/env python3
"""Strict TSPLIB/CVRPLIB reader for the RouteLogic public gap campaign."""

from __future__ import annotations

import argparse
import gzip
import hashlib
import json
import math
import re
from pathlib import Path
from typing import Any


class BenchmarkFormatError(RuntimeError):
    pass


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


def read_text(path: Path) -> str:
    if path.suffix == ".gz":
        with gzip.open(path, "rt", encoding="ascii") as handle:
            return handle.read()
    return path.read_text(encoding="ascii")


def parse_headers_and_sections(text: str) -> tuple[dict[str, str], dict[str, list[str]]]:
    headers: dict[str, str] = {}
    sections: dict[str, list[str]] = {}
    current: str | None = None
    known_sections = {"NODE_COORD_SECTION", "DEMAND_SECTION", "DEPOT_SECTION", "EDGE_WEIGHT_SECTION"}
    for raw_line in text.replace("\r\n", "\n").replace("\r", "\n").split("\n"):
        line = raw_line.strip()
        if not line:
            continue
        if line == "EOF":
            current = None
            break
        if line in known_sections:
            current = line
            sections[current] = []
            continue
        if current is not None:
            sections[current].append(line)
            continue
        match = re.match(r"^([^:]+?)\s*:\s*(.*?)\s*$", line)
        if match:
            headers[match.group(1).strip().upper()] = match.group(2).strip()
    return headers, sections


def parse_coords(lines: list[str], dimension: int) -> list[tuple[int, float, float]]:
    coords: list[tuple[int, float, float]] = []
    for line in lines:
        parts = line.split()
        if len(parts) != 3:
            raise BenchmarkFormatError(f"invalid NODE_COORD_SECTION row: {line!r}")
        coords.append((int(parts[0]), float(parts[1]), float(parts[2])))
    coords.sort(key=lambda item: item[0])
    expected_ids = list(range(1, dimension + 1))
    actual_ids = [item[0] for item in coords]
    if actual_ids != expected_ids:
        raise BenchmarkFormatError(f"node ids must be exactly 1..{dimension}, got {actual_ids[:8]}...")
    return coords


def euc_2d_matrix(coords: list[tuple[int, float, float]]) -> list[list[int]]:
    # TSPLIB95 EUC_2D, literally: (int)(sqrt(dx*dx + dy*dy) + 0.5).
    matrix: list[list[int]] = []
    for _, x1, y1 in coords:
        row: list[int] = []
        for _, x2, y2 in coords:
            dx = x1 - x2
            dy = y1 - y2
            row.append(int(math.sqrt(dx * dx + dy * dy) + 0.5))
        matrix.append(row)
    return matrix


def parse_demands(lines: list[str], dimension: int) -> dict[int, int]:
    demands: dict[int, int] = {}
    for line in lines:
        parts = line.split()
        if len(parts) != 2:
            raise BenchmarkFormatError(f"invalid DEMAND_SECTION row: {line!r}")
        node_id, demand = int(parts[0]), int(parts[1])
        if node_id in demands or demand < 0:
            raise BenchmarkFormatError(f"invalid demand for node {node_id}")
        demands[node_id] = demand
    if sorted(demands) != list(range(1, dimension + 1)):
        raise BenchmarkFormatError("DEMAND_SECTION must contain every node exactly once")
    return demands


def parse_depot(lines: list[str]) -> int:
    values = [int(line.split()[0]) for line in lines]
    depots = [value for value in values if value != -1]
    if len(depots) != 1:
        raise BenchmarkFormatError(f"exactly one depot is required, found {depots}")
    return depots[0]


def parse_tsplib_reference(path: Path, name: str) -> int:
    html = path.read_text(encoding="latin-1")
    match = re.search(rf"\b{re.escape(name)}\s*:\s*([0-9]+)\b", html, flags=re.IGNORECASE)
    if not match:
        raise BenchmarkFormatError(f"official TSPLIB reference missing for {name}")
    return int(match.group(1))


def parse_cvrplib_reference(path: Path) -> int:
    text = path.read_text(encoding="ascii")
    match = re.search(r"(?im)^Cost\s+([0-9]+(?:\.[0-9]+)?)\s*$", text)
    if not match:
        raise BenchmarkFormatError(f"CVRPLIB BKS cost missing in {path}")
    value = float(match.group(1))
    if not value.is_integer():
        raise BenchmarkFormatError(f"non-integer BKS is outside this EUC_2D campaign: {value}")
    return int(value)


def load_instance(base_dir: Path, spec: dict[str, Any]) -> dict[str, Any]:
    instance_path = base_dir / spec["instance"]["file"]
    reference_path = base_dir / spec["reference"]["file"]
    if sha256_file(instance_path) != spec["instance"]["sha256"]:
        raise BenchmarkFormatError(f"instance SHA-256 mismatch: {spec['name']}")
    if sha256_file(reference_path) != spec["reference"]["sha256"]:
        raise BenchmarkFormatError(f"reference SHA-256 mismatch: {spec['name']}")

    headers, sections = parse_headers_and_sections(read_text(instance_path))
    name = headers.get("NAME", "").strip()
    if name != spec["name"]:
        raise BenchmarkFormatError(f"NAME mismatch: expected {spec['name']!r}, got {name!r}")
    if headers.get("EDGE_WEIGHT_TYPE") != "EUC_2D":
        raise BenchmarkFormatError(
            f"{name}: EDGE_WEIGHT_TYPE is {headers.get('EDGE_WEIGHT_TYPE')!r}, not EUC_2D"
        )
    dimension = int(headers.get("DIMENSION", "0"))
    if dimension != spec["n"]:
        raise BenchmarkFormatError(f"{name}: DIMENSION {dimension} != locked n {spec['n']}")
    coords = parse_coords(sections.get("NODE_COORD_SECTION", []), dimension)
    matrix = euc_2d_matrix(coords)

    if spec["library"] == "TSPLIB":
        if headers.get("TYPE") != "TSP":
            raise BenchmarkFormatError(f"{name}: TYPE must be TSP")
        depot_node_id = 1
        demands = {node_id: 0 for node_id in range(1, dimension + 1)}
        capacity = None
        reference_value = parse_tsplib_reference(reference_path, name)
    elif spec["library"] == "CVRPLIB":
        if headers.get("TYPE") != "CVRP":
            raise BenchmarkFormatError(f"{name}: TYPE must be CVRP")
        depot_node_id = parse_depot(sections.get("DEPOT_SECTION", []))
        demands = parse_demands(sections.get("DEMAND_SECTION", []), dimension)
        capacity = int(headers.get("CAPACITY", "0"))
        if capacity <= 0:
            raise BenchmarkFormatError(f"{name}: CAPACITY must be positive")
        reference_value = parse_cvrplib_reference(reference_path)
    else:
        raise BenchmarkFormatError(f"unsupported library {spec['library']!r}")

    if reference_value != spec["referenceValue"]:
        raise BenchmarkFormatError(
            f"{name}: downloaded reference {reference_value} != locked value {spec['referenceValue']}"
        )
    if depot_node_id != 1:
        raise BenchmarkFormatError(f"{name}: this locked campaign requires depot node 1, got {depot_node_id}")

    return {
        "library": spec["library"],
        "name": name,
        "n": dimension,
        "k": spec["k"],
        "capacity": capacity,
        "depotNodeId": depot_node_id,
        "nodeIds": [node_id for node_id, _, _ in coords],
        "coords": coords,
        "demands": demands,
        "matrix": matrix,
        "referenceValue": reference_value,
        "referenceType": spec["referenceType"],
        "instanceSource": spec["instance"],
        "referenceSource": spec["reference"],
    }


def canonical_problem(instance: dict[str, Any], seed: int, campaign_id: str) -> dict[str, Any]:
    is_cvrp = instance["library"] == "CVRPLIB"
    depot_index = instance["nodeIds"].index(instance["depotNodeId"])
    dimensions = ["load"] if is_cvrp else []
    vehicles = []
    for index in range(instance["k"]):
        vehicle: dict[str, Any] = {
            "id": f"V{index + 1}",
            "startDepotId": "D1",
            "endDepotId": "D1",
            "capacity": {"load": instance["capacity"]} if is_cvrp else {},
            "fixedCost": 0,
            "perDistanceCost": 0,
            "perDurationCost": 0,
        }
        vehicles.append(vehicle)

    jobs = []
    for matrix_index, node_id in enumerate(instance["nodeIds"]):
        if node_id == instance["depotNodeId"]:
            continue
        jobs.append({
            "id": f"N{node_id}",
            "matrixIndex": matrix_index,
            "type": "delivery",
            "demand": {"load": instance["demands"][node_id]} if is_cvrp else {},
            "serviceMin": 0,
            "mandatory": True,
        })

    problem_id = f"{campaign_id}-{instance['name']}-seed-{seed}"
    return {
        "requestId": problem_id,
        "datasetId": instance["name"],
        "problemId": problem_id,
        "seed": seed,
        "capacityDimensions": dimensions,
        "matrix": {"distance": instance["matrix"], "duration": instance["matrix"]},
        "depots": [{"id": "D1", "matrixIndex": depot_index}],
        "vehicles": vehicles,
        "jobs": jobs,
        "pickupDeliveryPairs": [],
        "precedence": [],
        "objective": {"weights": {
            "fixedVehicle": 0,
            "distance": 1,
            "duration": 0,
            "driving": 0,
            "service": 0,
            "waiting": 0,
            "unassignedMandatory": 0,
            "unassignedOptional": 0,
        }},
    }


def inspect_config(config_path: Path) -> dict[str, Any]:
    config = json.loads(config_path.read_text(encoding="utf-8"))
    base_dir = config_path.parent
    rows = []
    for spec in config["instances"]:
        instance = load_instance(base_dir, spec)
        rows.append({
            "library": instance["library"],
            "name": instance["name"],
            "n": instance["n"],
            "k": instance["k"],
            "capacity": instance["capacity"],
            "edgeWeightType": "EUC_2D",
            "depotNodeId": instance["depotNodeId"],
            "referenceValue": instance["referenceValue"],
            "referenceType": instance["referenceType"],
        })
    return {"ok": True, "count": len(rows), "instances": rows}


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("config", type=Path)
    parser.add_argument("--inspect", action="store_true")
    args = parser.parse_args()
    payload = inspect_config(args.config)
    print(json.dumps(payload, indent=2, sort_keys=True))


if __name__ == "__main__":
    main()
