#!/usr/bin/env python3
"""Independent raw-matrix verifier for RouteLogic public quality gaps."""

from __future__ import annotations

import argparse
import collections
import hashlib
import json
from pathlib import Path
from typing import Any


class IndependentVerificationError(RuntimeError):
    def __init__(self, code: str, message: str, details: dict[str, Any] | None = None):
        super().__init__(message)
        self.code = code
        self.details = details or {}


def exact_number(value: Any, label: str) -> int:
    if isinstance(value, bool) or not isinstance(value, (int, float)):
        raise IndependentVerificationError("NON_NUMERIC_OBJECTIVE", f"{label} is not numeric", {"value": value})
    if int(value) != value:
        raise IndependentVerificationError("NON_INTEGER_OBJECTIVE", f"{label} is not integer", {"value": value})
    return int(value)


def active_families(problem: dict[str, Any], result: dict[str, Any]) -> list[str]:
    families: list[str] = []
    if len(problem["vehicles"]) > 1 and len(result.get("routes", [])) > 1:
        families.append("multiple_vehicles")
    if problem["jobs"] and all(job.get("mandatory", True) for job in problem["jobs"]):
        families.append("mandatory_jobs")
    return families


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


def verify(instance: dict[str, Any], problem: dict[str, Any], result: dict[str, Any]) -> dict[str, Any]:
    if result.get("status") != "FEASIBLE" or result.get("feasible") is not True:
        raise IndependentVerificationError(
            "RESULT_NOT_FEASIBLE",
            "only FEASIBLE results can be assigned a quality gap",
            {"status": result.get("status"), "error": result.get("error")},
        )

    expected_jobs = [job["id"] for job in problem["jobs"]]
    expected_set = set(expected_jobs)
    if len(expected_set) != len(expected_jobs):
        raise IndependentVerificationError("DUPLICATE_EXPECTED_JOB", "canonical benchmark jobs are not unique")

    unassigned_job_ids = result.get("unassignedJobIds")
    unassigned_service_ids = result.get("unassignedServiceIds")
    mandatory_unassigned = result.get("mandatoryUnassigned")
    if unassigned_job_ids != [] or unassigned_service_ids != [] or mandatory_unassigned != []:
        raise IndependentVerificationError(
            "UNASSIGNED_JOB",
            "the solver returned at least one unassigned mandatory job/service",
            {
                "unassignedJobIds": unassigned_job_ids,
                "unassignedServiceIds": unassigned_service_ids,
                "mandatoryUnassigned": mandatory_unassigned,
            },
        )

    routes = result.get("routes")
    if not isinstance(routes, list):
        raise IndependentVerificationError("MISSING_ROUTES", "result.routes is not an array")
    expected_active_routes = 1 if instance["library"] == "TSPLIB" else instance["k"]
    if len(problem["vehicles"]) != instance["k"]:
        raise IndependentVerificationError(
            "WRONG_VEHICLE_SUPPLY", "canonical problem does not supply exactly k vehicles",
            {"expected": instance["k"], "actual": len(problem["vehicles"])},
        )
    if len(routes) != expected_active_routes:
        raise IndependentVerificationError(
            "WRONG_ACTIVE_VEHICLE_COUNT",
            "solution does not use the same number of routes as the locked reference perimeter",
            {"expected": expected_active_routes, "actual": len(routes)},
        )

    visits: list[str] = []
    for route in routes:
        service_ids = route.get("serviceIds")
        if not isinstance(service_ids, list):
            raise IndependentVerificationError("MISSING_SERVICE_SEQUENCE", "route serviceIds is not an array")
        visits.extend(service_ids)
    counts = collections.Counter(visits)
    missing = sorted(expected_set - set(counts))
    unknown = sorted(set(counts) - expected_set)
    duplicated = sorted(job_id for job_id, count in counts.items() if count != 1)
    if missing or unknown or duplicated or len(visits) != len(expected_jobs):
        raise IndependentVerificationError(
            "EXACT_ONCE_COVERAGE_FAILURE",
            "every non-depot node must occur exactly once across all routes",
            {
                "expectedCount": len(expected_jobs),
                "actualVisitCount": len(visits),
                "missing": missing,
                "unknown": unknown,
                "duplicatedOrWrongCount": {job_id: counts[job_id] for job_id in duplicated},
            },
        )

    job_by_id = {job["id"]: job for job in problem["jobs"]}
    vehicle_by_id = {vehicle["id"]: vehicle for vehicle in problem["vehicles"]}
    depot_by_id = {depot["id"]: depot for depot in problem["depots"]}
    matrix = problem["matrix"]["distance"]
    total_distance = 0
    route_distances: dict[str, int] = {}
    route_loads: dict[str, int] = {}
    for route in routes:
        vehicle_id = route.get("vehicleId")
        vehicle = vehicle_by_id.get(vehicle_id)
        if vehicle is None:
            raise IndependentVerificationError("UNKNOWN_VEHICLE", f"unknown vehicle {vehicle_id!r}")
        start = depot_by_id[vehicle["startDepotId"]]["matrixIndex"]
        end = depot_by_id[vehicle["endDepotId"]]["matrixIndex"]
        indices = [job_by_id[service_id]["matrixIndex"] for service_id in route["serviceIds"]]
        sequence = [start, *indices, end]
        distance = sum(matrix[sequence[index]][sequence[index + 1]] for index in range(len(sequence) - 1))
        declared_route_distance = exact_number(route.get("distance"), f"route {vehicle_id} distance")
        if declared_route_distance != distance:
            raise IndependentVerificationError(
                "ROUTE_DISTANCE_MISMATCH",
                f"route {vehicle_id} distance differs from raw matrix recomputation",
                {"declared": declared_route_distance, "recomputed": distance},
            )
        route_distances[vehicle_id] = distance
        total_distance += distance

        if instance["library"] == "CVRPLIB":
            load = sum(int(job_by_id[service_id]["demand"]["load"]) for service_id in route["serviceIds"])
            capacity = int(vehicle["capacity"]["load"])
            if load > capacity:
                raise IndependentVerificationError(
                    "CAPACITY_FAILURE", f"route {vehicle_id} exceeds capacity Q",
                    {"load": load, "capacity": capacity},
                )
            route_loads[vehicle_id] = load

    objective = result.get("objective") or {}
    components = objective.get("components") or {}
    declared_values = {
        "objective.economic": exact_number(objective.get("economic"), "objective.economic"),
        "objective.components.rawTotalDistance": exact_number(
            components.get("rawTotalDistance"), "objective.components.rawTotalDistance"
        ),
        "objective.components.weightedDistance": exact_number(
            components.get("weightedDistance"), "objective.components.weightedDistance"
        ),
        "objective.objectiveValue": exact_number(objective.get("objectiveValue"), "objective.objectiveValue"),
        "verification.independentEconomic": exact_number(
            (result.get("verification") or {}).get("independentEconomic"),
            "verification.independentEconomic",
        ),
    }
    non_distance_components = {
        key: exact_number(components.get(key), f"objective.components.{key}")
        for key in (
            "weightedFixedVehicle", "weightedDuration", "weightedDriving", "weightedService",
            "weightedWaiting", "vehicleVariableCost", "optionalUnassignedPenalty",
        )
    }
    mismatched = {label: value for label, value in declared_values.items() if value != total_distance}
    nonzero = {label: value for label, value in non_distance_components.items() if value != 0}
    if mismatched or nonzero:
        raise IndependentVerificationError(
            "OBJECTIVE_MISMATCH",
            "solver objective is not exactly the independently recomputed raw distance",
            {"recomputedDistance": total_distance, "mismatched": mismatched, "nonzeroComponents": nonzero},
        )
    if (result.get("verification") or {}).get("matchesSolver") is not True:
        raise IndependentVerificationError(
            "BUILTIN_VERIFIER_DISAGREES", "production independent verifier does not match solver"
        )

    canonical_fingerprint_payload = [
        {"vehicleId": route["vehicleId"], "serviceIds": list(route["serviceIds"])}
        for route in sorted(routes, key=lambda item: item["vehicleId"])
    ]
    independent_fingerprint = hashlib.sha256(
        json.dumps(canonical_fingerprint_payload, separators=(",", ":"), sort_keys=True).encode("utf-8")
    ).hexdigest()

    return {
        "status": "PASS",
        "allJobsMandatory": True,
        "unassignedJobCount": 0,
        "exactOnceCoverage": True,
        "visitedJobCount": len(visits),
        "activeRouteCount": len(routes),
        "routeDistances": route_distances,
        "routeLoads": route_loads,
        "recomputedDistance": total_distance,
        "objectiveEquality": True,
        "independentSolutionFingerprint": independent_fingerprint,
        "activeFamilies": active_families(problem, result),
        "benchmarkConstraints": benchmark_constraints(instance),
    }


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("payload", type=Path, help="JSON with instance, problem and result")
    args = parser.parse_args()
    payload = json.loads(args.payload.read_text(encoding="utf-8"))
    try:
        output = verify(payload["instance"], payload["problem"], payload["result"])
    except IndependentVerificationError as error:
        output = {"status": "FAIL", "code": error.code, "message": str(error), "details": error.details}
        print(json.dumps(output, indent=2, sort_keys=True))
        raise SystemExit(2)
    print(json.dumps(output, indent=2, sort_keys=True))


if __name__ == "__main__":
    main()
