#!/usr/bin/env python3
"""Strict finite metric validation, primal/dual audits and bounded exact k-median."""
import argparse
from fractions import Fraction
from itertools import combinations
import json
import math
from pathlib import Path

MODEL = dict(objective="unweighted_sum_nearest_distance", distance="strict_rational_metric",
             capacities="unlimited", facility_budget="at_most_k", extra_constraints="none")
MAX_POINTS = 48
MAX_SUBSETS = 100000
MAX_LOOKUPS = 5000000


def integer(x, lo, hi, label):
    if type(x) is not int or not lo <= x <= hi:
        raise ValueError(f"{label} must be an integer in [{lo},{hi}]")
    return x


def rational(x):
    if not isinstance(x, str) or len(x) > 160:
        raise ValueError("Require bounded canonical rational strings")
    try:
        q = Fraction(x)
    except (ValueError, ZeroDivisionError):
        raise ValueError("Invalid rational") from None
    if str(q) != x or max(abs(q.numerator).bit_length(), q.denominator.bit_length()) > 128:
        raise ValueError("Require canonical rational with at most 128-bit numerator/denominator")
    return q


def index_set(values, n, label):
    if not isinstance(values, list) or len(values) > n:
        raise ValueError(f"Invalid {label} list")
    for x in values:
        integer(x, 0, n-1, label)
    if len(set(values)) != len(values):
        raise ValueError(f"Duplicate {label} indices")
    return sorted(values)


def parse(record):
    if not isinstance(record, dict) or set(record) != {"distances", "clients", "facilities", "k", "model"}:
        raise ValueError("Require exactly distances, clients, facilities, k, model")
    if record["model"] != MODEL:
        raise ValueError("Unsupported model; weights, capacities, extra constraints and directed distances need separate bridges")
    raw = record["distances"]
    if not isinstance(raw, list):
        raise ValueError("Require a square complete distance matrix")
    n = integer(len(raw), 1, MAX_POINTS, "point count")
    if any(not isinstance(row, list) or len(row) != n for row in raw):
        raise ValueError("Require a square complete distance matrix")
    d = [[rational(x) for x in row] for row in raw]
    clients = index_set(record["clients"], n, "client")
    facilities = index_set(record["facilities"], n, "facility")
    k = integer(record["k"], 1, len(facilities), "k")
    for i in range(n):
        for j in range(n):
            if (i == j and d[i][j] != 0) or (i != j and d[i][j] <= 0):
                raise ValueError(f"Nonnegative strict metric/diagonal failure at ({i},{j})")
            if d[i][j] != d[j][i]:
                raise ValueError(f"Symmetry failure at ({i},{j})")
            for h in range(n):
                if d[i][h] > d[i][j] + d[j][h]:
                    raise ValueError(f"Triangle failure at ({i},{j},{h})")
    return d, clients, facilities, k


def primal(data, selected):
    d, clients, facilities, k = data
    chosen = index_set(selected, len(d), "selected")
    if not chosen or len(chosen) > k or not set(chosen) <= set(facilities):
        raise ValueError("Require one to k distinct candidate facilities")
    assignments = [min(chosen, key=lambda f: (d[c][f], f)) for c in clients]
    cost = sum((d[c][f] for c, f in zip(clients, assignments)), Fraction(0))
    return dict(valid=True, selected=chosen, clients=clients, assignments=assignments, exact_cost=str(cost))


def dual(data, certificate):
    """Verify LP weak duality: alpha_j <= d_jf + beta_jf; sum_j beta_jf <= lambda."""
    d, clients, facilities, k = data
    if not isinstance(certificate, dict) or set(certificate) != {"alpha", "beta", "lambda"}:
        raise ValueError("Require exactly alpha, beta, lambda, in sorted client/facility order")
    a, b = certificate["alpha"], certificate["beta"]
    if not isinstance(a, list) or len(a) != len(clients):
        raise ValueError("Wrong alpha shape")
    if not isinstance(b, list) or len(b) != len(clients) or any(not isinstance(row, list) or len(row) != len(facilities) for row in b):
        raise ValueError("Wrong beta shape")
    a, b, lam = [rational(x) for x in a], [[rational(x) for x in row] for row in b], rational(certificate["lambda"])
    if lam < 0 or any(x < 0 for row in b for x in row):
        raise ValueError("Require nonnegative beta and lambda")
    for j, c in enumerate(clients):
        for f, site in enumerate(facilities):
            if a[j] > d[c][site] + b[j][f]:
                raise ValueError("Client dual inequality failure")
    for f in range(len(facilities)):
        if sum((row[f] for row in b), Fraction(0)) > lam:
            raise ValueError("Facility dual inequality failure")
    return dict(valid=True, exact_lower_bound=str(sum(a, Fraction(0))-k*lam),
                ordering="sorted input client and facility indices", independent_of_source_theorem=True)


def calculate(record, selected=None, certificate=None, subset_budget=MAX_SUBSETS, lookup_budget=MAX_LOOKUPS):
    integer(subset_budget, 0, MAX_SUBSETS, "subset budget")
    integer(lookup_budget, 0, MAX_LOOKUPS, "lookup budget")
    data = parse(record)
    d, clients, facilities, k = data
    witness = primal(data, selected if selected is not None else facilities[:k])
    upper = Fraction(witness["exact_cost"])
    lower = sum((min(d[c][f] for f in facilities) for c in clients), Fraction(0))
    dual_audit = dual(data, certificate) if certificate is not None else None
    if dual_audit:
        lower = max(lower, Fraction(dual_audit["exact_lower_bound"]))
    subsets = lookups = 0
    complete = upper == lower
    reason = "primal_matches_lower_bound" if complete else "budget_exhausted"
    if not complete:
        for chosen in combinations(facilities, k):
            needed = len(clients)*k
            if subsets >= subset_budget or lookups + needed > lookup_budget:
                break
            cost = sum((min(d[c][f] for f in chosen) for c in clients), Fraction(0))
            subsets += 1
            lookups += needed
            if cost < upper:
                witness = primal(data, list(chosen))
                upper = cost
            if upper == lower:
                complete, reason = True, "primal_matches_lower_bound"
                break
        else:
            complete, reason = True, "all_k_subsets_completed"
    if complete:
        lower = upper
    if upper < lower:
        raise ArithmeticError("Checked lower bound exceeds checked primal")
    return dict(status="exact_optimum" if complete else "bounded_unknown_optimum", reason=reason,
                metric_valid=True, primal=witness, dual=dual_audit, exact_lower_bound=str(lower),
                exact_upper_bound=str(upper), independently_known_optimum=str(upper) if complete else None,
                certified_ratio_upper=str(upper/lower) if lower > 0 else ("1" if upper == 0 else None),
                absolute_gap=str(upper-lower), total_k_subsets=math.comb(len(facilities), k),
                examined_subsets=subsets, enumeration_distance_lookups=lookups,
                subset_budget=subset_budget, lookup_budget=lookup_budget,
                resource_scope="Counts cover enumeration subsets and distance lookups only; metric parsing, O(n^3) triangle checks, primal/dual audits, Fraction bit arithmetic, allocations and runtime are not charged. Input dimensions/rationals have separate caps.",
                enumeration_justification="An at-most-k nonempty set can be augmented to k candidates without increasing uncapacitated nearest-distance cost. Exhausting exactly-k subsets therefore suffices.",
                source_algorithm_implemented=False, source_proof_verification="not_run", formal_python_verification="not_run",
                input_scope="Exact frozen unweighted finite metric only; no rounded/geographic/traffic fidelity or buyer value is established.")


def main():
    p = argparse.ArgumentParser(description=__doc__)
    p.add_argument("input", type=Path)
    p.add_argument("--selected", type=Path)
    p.add_argument("--dual", type=Path)
    p.add_argument("--subset-budget", type=int, default=MAX_SUBSETS)
    p.add_argument("--lookup-budget", type=int, default=MAX_LOOKUPS)
    p.add_argument("--output", type=Path)
    a = p.parse_args()
    result = calculate(json.loads(a.input.read_text()), json.loads(a.selected.read_text()) if a.selected else None,
                       json.loads(a.dual.read_text()) if a.dual else None, a.subset_budget, a.lookup_budget)
    body = json.dumps(result, indent=2, allow_nan=False)+"\n"
    if a.output:
        with a.output.open("x") as f:
            f.write(body)
    else:
        print(body, end="")


if __name__ == "__main__":
    main()
