#!/usr/bin/env python3
"""Exact educational replay; no third-party packages and no hidden random seed.
Core functions take a promised metric matrix. Call validate_metric at the boundary.
Continuous beta integration enumerates exact breakpoint intervals, not a grid.
"""
from fractions import Fraction as F
from itertools import combinations, permutations
from math import factorial
import json


def check(ok, message):
    if not ok:
        raise ValueError(message)


def rational(x):
    check(type(x) in (int, F), 'exact int or Fraction required')
    return F(x)


def integer(x, low, high):
    check(type(x) is int and low <= x <= high, 'integer outside interface')
    return x


def validate_metric(raw):
    n = len(raw)
    check(all(len(row) == n for row in raw), 'square matrix required')
    d = tuple(tuple(rational(x) for x in row) for row in raw)
    for u in range(n):
        for v in range(n):
            check((d[u][v] == 0 if u == v else d[u][v] > 0), 'positive separation required')
            check(d[u][v] == d[v][u], 'symmetry required')
            for w in range(n):
                check(d[u][v] <= d[u][w] + d[w][v], 'triangle inequality required')
    return d


def tape(n, pi, beta):
    pi = tuple(pi)
    check(len(pi) == n and all(type(v) is int for v in pi)
          and sorted(pi) == list(range(n)), 'permutation required')
    beta = rational(beta)
    check(1 <= beta < 2, 'beta must lie in [1,2)')
    return pi, beta


def farthest_first(d, k, start=0):
    n = len(d)
    integer(k, 0 if n == 0 else 1, n)
    if not n:
        return {'centers': [], 'owner': [], 'nearest': [], 'radius': F(0),
                'packing': [], 'lower': F(0), 'trace': [], 'queries': 0}
    integer(start, 0, n - 1)
    nearest, owner = [None] * n, [None] * n
    centers, trace, queries = [], [], 0
    c = start
    for _ in range(k):
        centers.append(c)
        for x in range(n):
            queries += 1
            if nearest[x] is None or d[x][c] < nearest[x]:
                nearest[x], owner[x] = d[x][c], c
        radius = max(nearest)
        trace.append({'added': c, 'radius': radius})
        # Positive separation ensures the maximizer is new whenever k < n.
        c = max(range(n), key=lambda x: (nearest[x], -x))
    witness = centers + [c] if k < n else []
    return {'centers': centers, 'owner': owner, 'nearest': nearest,
            'radius': radius, 'packing': witness, 'lower': radius / 2,
            'trace': trace, 'queries': queries}


def verify_cover(d, k, result):
    n = len(d)
    integer(k, 0 if n == 0 else 1, n)
    centers, owner = result['centers'], result['owner']
    check(len(centers) == k and all(type(c) is int and 0 <= c < n for c in centers)
          and len(set(centers)) == k, 'bad center identities')
    r = rational(result['radius'])
    check(r >= 0 and len(owner) == n, 'bad cover dimensions')
    selected = [False] * n
    for c in centers:
        selected[c] = True
    check(all(type(c) is int and 0 <= c < n and selected[c] for c in owner), 'bad assignments')
    check(all(d[x][owner[x]] <= r for x in range(n)), 'uncovered point')
    if n and k < n:
        p = result['packing']
        check(len(p) == k + 1 and all(type(x) is int and 0 <= x < n for x in p)
              and len(set(p)) == len(p), 'bad packing identities')
        check(all(d[x][y] >= r for x, y in combinations(p, 2)), 'packing fails')
    else:
        check(r == 0 and result['packing'] == [], 'zero-radius branch required')
    return True


def optimum_centers(d, k):
    n = len(d)
    integer(k, 0 if n == 0 else 1, n)
    if not n:
        return F(0), ()
    return min((max(min(d[x][c] for c in cs) for x in range(n)), cs)
               for cs in combinations(range(n), k))


def partition(d, delta, pi, beta):
    n = len(d)
    pi, beta = tape(n, pi, beta)
    delta = rational(delta)
    check(delta > 0, 'positive scale required')
    radius = beta * delta / 4
    owner, blocks, queries = [None] * n, [], 0
    for center in pi:  # Never skip an already assigned center.
        block = []
        for x in range(n):
            if owner[x] is None:
                queries += 1
                if d[center][x] <= radius:
                    owner[x] = center
                    block.append(x)
        if block:
            blocks.append({'center': center, 'points': block})
    check(all(c is not None for c in owner), 'partition incomplete')
    return {'radius': radius, 'owner': owner, 'blocks': blocks, 'queries': queries}


def pow2(i):
    return F(1 << i) if i >= 0 else F(1, 1 << -i)


def floor_log2(x):
    check(x > 0, 'log needs positive input')
    e = x.numerator.bit_length() - x.denominator.bit_length()
    return e if pow2(e) <= x else e - 1


def scales(d):
    n = len(d)
    if n < 2:
        return []
    distances = [d[x][y] for x, y in combinations(range(n), 2)]
    low = floor_log2(min(distances))
    high = floor_log2(max(distances))
    if pow2(high) < max(distances):
        high += 1
    return list(range(high + 2, low - 1, -1))


def hierarchy(d, pi, beta):
    n = len(d)
    pi, beta = tape(n, pi, beta)
    levels = scales(d)
    if n < 2:
        return {'levels': [], 'nodes': [] if not n else
                [{'id': 0, 'level': None, 'label': F(0), 'points': [0], 'parent': None}],
                'edges': [], 'partitions': [], 'distance': [[F(0)]] if n else [],
                'queries': 0}
    nodes, history, edges, queries = [], [], [], 0
    parent = [None] * n
    for i in levels:
        q = partition(d, pow2(i), pi, beta)
        queries += q['queries']
        # Refine with the old parent identity, never use Q_i alone.
        groups = {}
        for x in range(n):
            key = (parent[x], q['owner'][x])
            if key not in groups:
                groups[key] = []
            groups[key].append(x)
        current, ids = [None] * n, []
        for (p, _), points in groups.items():
            node = len(nodes)
            label = F(0) if i == levels[-1] else pow2(i)
            nodes.append({'id': node, 'level': i, 'label': label,
                          'points': points, 'parent': p})
            ids.append(node)
            if p is not None:
                edges.append([p, node, (nodes[p]['label'] - label) / 2])
            for x in points:
                current[x] = node
        history.append({'level': i, 'radius': q['radius'], 'raw_owner': q['owner'],
                        'node_of': current, 'nodes': ids})
        parent = current
    dt = [[F(0) for _ in range(n)] for _ in range(n)]
    for u, v in combinations(range(n), 2):
        common = None
        for row in history:
            if row['node_of'][u] != row['node_of'][v]:
                break
            common = row['node_of'][u]
        check(common is not None, 'missing common ancestor')
        dt[u][v] = dt[v][u] = nodes[common]['label']
    return {'levels': levels, 'nodes': nodes, 'edges': edges,
            'partitions': history, 'distance': dt, 'queries': queries}


def beta_intervals(d, levels):
    cuts = {F(1), F(2)}
    for i in levels:
        for row in d:
            for distance in row:
                threshold = 4 * distance / pow2(i)
                if 1 < threshold < 2:
                    cuts.add(threshold)
    cuts = sorted(cuts)
    return list(zip(cuts, cuts[1:]))


def exact_law(d, delta, ball_center, ball_radius):
    """Factorial diagnostic: full uniform permutation + continuous uniform beta."""
    n = len(d)
    check(n > 0, 'law example needs a point')
    integer(ball_center, 0, n - 1)
    ball_radius = rational(ball_radius)
    check(ball_radius >= 0, 'negative ball radius')
    delta = rational(delta)
    check(delta > 0, 'positive delta')
    levels = scales(d)
    intervals = beta_intervals(d, levels)
    # Also split at the independent fixed-scale partition's thresholds.
    cuts = {x for interval in intervals for x in interval} or {F(1), F(2)}
    cuts.update(4 * distance / delta for row in d for distance in row
                if 1 < 4 * distance / delta < 2)
    cuts = sorted(cuts)
    intervals = list(zip(cuts, cuts[1:]))
    expected = [[F(0) for _ in range(n)] for _ in range(n)]
    separation = [[F(0) for _ in range(n)] for _ in range(n)]
    ball = [v for v in range(n) if d[ball_center][v] <= ball_radius]
    fail, mass, tapes = F(0), F(0), 0
    for pi in permutations(range(n)):
        for a, b in intervals:
            beta, weight = (a + b) / 2, (b - a) / factorial(n)
            tree = hierarchy(d, pi, beta)
            part = partition(d, delta, pi, beta)
            mass += weight
            tapes += 1
            fail += weight * (len({part['owner'][v] for v in ball}) > 1)
            for u in range(n):
                for v in range(n):
                    check(tree['distance'][u][v] >= d[u][v], 'tree contraction')
                    expected[u][v] += weight * tree['distance'][u][v]
                    separation[u][v] += weight * (part['owner'][u] != part['owner'][v])
    check(mass == 1, 'probability mass')
    harmonic = sum((F(1, j) for j in range(1, n + 1)), F(0))
    check(all(expected[u][v] <= 8 * harmonic * d[u][v]
              for u in range(n) for v in range(n)), 'stretch bound')
    check(all(separation[u][v] <= 4 * harmonic * d[u][v] / delta
              for u in range(n) for v in range(n)), 'cut bound')
    check(fail <= 8 * harmonic * ball_radius / delta, 'padding bound')
    return {'intervals': intervals, 'tapes': tapes, 'mass': mass,
            'expected_tree_distance': expected, 'separation_probability': separation,
            'ball': ball, 'ball_failure': fail, 'harmonic': harmonic}


def line_metric(points):
    return validate_metric([[abs(x - y) for y in points] for x in points])


def main():
    d = line_metric([0, 1, 2, 3, 4])
    centers = farthest_first(d, 2)
    check(verify_cover(d, 2, centers), 'cover')
    optimum = optimum_centers(d, 2)
    check(centers['radius'] == 2 and optimum[0] == 1, 'main centers')
    pi, beta, delta = (2, 0, 4, 1, 3), F(3, 2), F(4)
    part = partition(d, delta, pi, beta)
    tree = hierarchy(d, pi, beta)
    law = exact_law(d, delta, 2, F(1))
    weighted_original = sum(d[i][i + 1] for i in range(4))
    weighted_expected = sum(law['expected_tree_distance'][i][i + 1] for i in range(4))
    # Star center already assigned: skipping it changes the partition.
    star = validate_metric([[0, 1, 2, 2], [1, 0, 1, 1], [2, 1, 0, 2], [2, 1, 2, 0]])
    weak = partition(star, F(8, 3), (0, 1, 2, 3), F(3, 2))
    check(weak['owner'] == [0, 0, 1, 1], 'assigned-center regression')
    enlarged = line_metric([0, 1, 2, 3, 1024])
    large_law = exact_law(enlarged, delta, 1, F(1))
    bad_tree = hierarchy(enlarged, (4, 0, 1, 2, 3), F(2047, 1024))
    check(bad_tree['distance'][0][1] == 4096, 'rare large stretch')
    crossing = hierarchy(d, (0, 4, 2, 1, 3), F(3, 2))
    check(verify_cover((), 0, farthest_first((), 0)), 'empty')
    check(verify_cover(((F(0),),), 1, farthest_first(((F(0),),), 1)), 'singleton')
    check(verify_cover(d, 5, farthest_first(d, 5)), 'all centers')
    rejected = 0
    bad_calls = [lambda: validate_metric([[0, 1], [2, 0]]),
                 lambda: validate_metric([[0, 0], [0, 0]]),
                 lambda: validate_metric([[0, 1, 4], [1, 0, 1], [4, 1, 0]]),
                 lambda: validate_metric([[0, 1], [1]]),
                 lambda: validate_metric([[False]]),
                 lambda: farthest_first(d, True), lambda: farthest_first(d, 0),
                 lambda: partition(d, 0, pi, beta),
                 lambda: hierarchy(d, (True, 0, 2, 3, 4), beta),
                 lambda: hierarchy(d, pi, F(2)),
                 lambda: hierarchy(d, pi, 1.5)]
    for call in bad_calls:
        try:
            call()
        except ValueError:
            rejected += 1
        else:
            raise ValueError('invalid input accepted')
    return {'points': [0, 1, 2, 3, 4], 'metric': d, 'cover': centers, 'optimum': optimum,
            'permutation': pi, 'beta': beta, 'delta': delta, 'partition': part,
            'tree': tree, 'law': law, 'fixed_adjacent_cost': [weighted_original, weighted_expected],
            'weak_not_strong': weak, 'large_aspect_ratio': {'points': [0, 1, 2, 3, 1024],
               'levels': scales(enlarged), 'law': large_law, 'bad_tape': {'pi': [4, 0, 1, 2, 3],
                   'beta': F(2047, 1024), 'distance_0_1': bad_tree['distance'][0][1]}},
            'crossing_raw_partitions': crossing['partitions'], 'rejected': rejected}


if __name__ == '__main__':
    print(json.dumps(main(), ensure_ascii=False, indent=2,
                     default=lambda x: str(x) if isinstance(x, F) else x))
