#!/usr/bin/env python3
"""Reproduce the three-atom Brenier sharpness example with rational geometry.

Integrates intersections of the two affine-max partitions of [-1,1]^2.
It is a benchmark reference, not an optimal-transport solver or proof checker.
"""
import argparse
import hashlib
import json
from fractions import Fraction
from math import isqrt
from pathlib import Path

MAX_WORK = 10_000


def exact_sqrt(value):
    n, d = isqrt(value.numerator), isqrt(value.denominator)
    return Fraction(n, d) if n*n == value.numerator and d*d == value.denominator else None


def calculate(record, work_budget=MAX_WORK):
    if not isinstance(record, dict) or set(record) != {'a'}:
        raise ValueError('Provide a only')
    raw = record['a']
    if not isinstance(raw, str) or len(raw) > 180:
        raise ValueError('a must be a bounded rational string')
    try:
        a = Fraction(raw)
    except (ValueError, ZeroDivisionError) as exc:
        raise ValueError('Invalid rational a') from exc
    if max(abs(a.numerator).bit_length(), a.denominator.bit_length()) > 256 or not 0 < a < Fraction(1, 2):
        raise ValueError('Require 0 < a < 1/2 and at most 256-bit rational input')
    if isinstance(work_budget, bool) or not isinstance(work_budget, int) or not 1 <= work_budget <= MAX_WORK:
        raise ValueError('Invalid work budget')
    b = a/2
    maps = [([(-1, 0), (1, 0), (0, 0)], [0, 0, a]),
            ([(-1, 0), (1, 0), (0, b)], [0, 0, a])]
    square = [(Fraction(-1), Fraction(-1)), (Fraction(1), Fraction(-1)),
              (Fraction(1), Fraction(1)), (Fraction(-1), Fraction(1))]
    work = 0

    def charge():
        nonlocal work
        if work >= work_budget:
            raise OverflowError('Polygon-vertex examination budget exhausted')
        work += 1

    def clip(polygon, normal, bound):
        if not polygon:
            return []
        output = []
        previous = polygon[-1]
        previous_value = sum((normal[h]*previous[h] for h in range(2)), Fraction(0)) - bound
        for current in polygon:
            charge()
            value = sum((normal[h]*current[h] for h in range(2)), Fraction(0)) - bound
            if (value <= 0) != (previous_value <= 0):
                t = previous_value/(previous_value-value)
                output.append(tuple(previous[h]+t*(current[h]-previous[h]) for h in range(2)))
            if value <= 0:
                output.append(current)
            previous, previous_value = current, value
        return output

    def cell(polygon, model, index):
        slopes, intercepts = model
        for other in range(3):
            if other != index:
                normal = tuple(slopes[other][h]-slopes[index][h] for h in range(2))
                polygon = clip(polygon, normal, intercepts[index]-intercepts[other])
        return polygon

    def area(polygon):
        if not polygon:
            return Fraction(0)
        total = Fraction(0)
        for i, point in enumerate(polygon):
            charge()
            following = polygon[(i+1) % len(polygon)]
            total += point[0]*following[1] - point[1]*following[0]
        return abs(total)/2

    result = dict(
        input_record_sha256=hashlib.sha256(json.dumps(record, sort_keys=True, separators=(',', ':')).encode()).hexdigest(),
        source_connection=dict(family='374', source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb',
                               source_proof_verification='not_run'),
        formal_program_verification='not_run', arithmetic='exact_rational',
        source_model='uniform_square_two_dimensions', parameter_a=str(a), parameter_b=str(b),
        method='exact_halfplane_clipping_and_polygon_area',
        limits=dict(parameter_bits=256, work_budget=work_budget, wall_time_deadline='none'),
        work_unit_definition='Vertices examined by clipping and area routines, not bit operations or elapsed time.')
    try:
        intersections = []
        first_masses, second_masses = [Fraction(0)]*3, [Fraction(0)]*3
        squared_map_distance = Fraction(0)
        changed_cell_mass = Fraction(0)
        for i in range(3):
            first = cell(square, maps[0], i)
            for j in range(3):
                polygon = cell(first, maps[1], j)
                mass = area(polygon)/4
                squared_gap = sum(((maps[0][0][i][h]-maps[1][0][j][h])**2 for h in range(2)), Fraction(0))
                first_masses[i] += mass
                second_masses[j] += mass
                squared_map_distance += mass*squared_gap
                if (i == 2) != (j == 2):
                    changed_cell_mass += mass
                intersections.append(dict(first_cell=i, second_cell=j, mass=str(mass),
                                          squared_map_gap=str(squared_gap),
                                          vertices=[[str(x), str(y)] for x, y in polygon]))
        assert sum(first_masses) == sum(second_masses) == 1
        expected_masses = [(1-a)/2, (1-a)/2, a]
        # Independent geometry result is compared with the paper's formulas.
        assert first_masses == second_masses == expected_masses
        assert squared_map_distance == b/2+a*b*b
        assert changed_cell_mass == b/2
    except OverflowError as exc:
        return dict(result, status='budget_exceeded', reason=str(exc), work_units=work,
                    squared_map_distance=None, geometric_cells=None,
                    interpretation='The complete geometric comparison is unknown.')
    squared_w2 = a*b*b
    w2 = exact_sqrt(squared_w2)
    return dict(result, status='complete', work_units=work,
                first_atom_masses=[str(x) for x in first_masses],
                second_atom_masses=[str(x) for x in second_masses],
                geometric_cells=intersections, central_cell_symmetric_difference_mass=str(changed_cell_mass),
                squared_map_distance=str(squared_map_distance), squared_target_w2=str(squared_w2),
                target_w1=str(a*b), target_w2=str(w2) if w2 is not None else None,
                target_distance_justification='Matching equal-mass atoms attains the second-coordinate lower bound; no general transport optimization was run.',
                required_half_exponent_constant_squared=str(squared_map_distance/w2) if w2 is not None else None,
                one_third_exponent_ratio_sixth_power=str(squared_map_distance**3/squared_w2),
                source_constant_squared_rational_upper_bound='9464',
                constant_bound_justification='For K=Y=[-1,1]^2 choose R^2=L^2=2, P_K/volume(K)=1 and sqrt(162)<=13 in the paper formula.',
                interpretation='Exact finite reproduction of the continuous two-dimensional counterexample formulas; extra cube coordinates factor out as described by the source. It neither proves the general stability theorem nor certifies a learned, discrete or entropic transport solver.')


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--work-budget', type=int, default=MAX_WORK)
    parser.add_argument('--output', type=Path)
    args = parser.parse_args()
    body = json.dumps(calculate(json.loads(args.input.read_text()), args.work_budget), indent=2, allow_nan=False) + '\n'
    if args.output:
        with args.output.open('x') as stream:
            stream.write(body)
    else:
        print(body, end='')


if __name__ == '__main__':
    main()
