"""Exact dyadic rounding experiment. Python 3.12+, standard library only.

All input fractions are generated here. No expected trajectory or escape count
is embedded in the program. Every ambiguous escape enclosure fails the run.
"""
from fractions import Fraction
import json
import math
from pathlib import Path
import struct
import sys


ZERO = Fraction(0)
QUARTER = Fraction(1, 4)
HALF = Fraction(1, 2)


def power_two(exponent):
    return Fraction(1 << exponent) if exponent >= 0 else Fraction(1, 1 << -exponent)


def nearest_even(value, precision):
    """Round a nonnegative rational to p significand bits, unbounded exponent."""
    assert value >= 0 and precision >= 3
    if value == 0:
        return ZERO
    exponent = value.numerator.bit_length() - value.denominator.bit_length()
    if value < power_two(exponent):
        exponent -= 1
    quantum = power_two(exponent - precision + 1)
    scaled = value / quantum
    integer, remainder = divmod(scaled.numerator, scaled.denominator)
    twice = 2 * remainder
    if twice > scaled.denominator or (twice == scaled.denominator and integer % 2):
        integer += 1
    return integer * quantum


def parameter(precision):
    return QUARTER + power_two(-precision - 1)


def step(value, c, precision, mode):
    if mode == "separate":
        return nearest_even(nearest_even(value * value, precision) + c, precision)
    assert mode == "fused"
    return nearest_even(value * value + c, precision)


def fraction_text(value):
    return f"{value.numerator}/{value.denominator}"


def rounded_trajectory(precision, mode):
    c = parameter(precision)
    current = ZERO
    # A conservative finite resource limit, not an escape or stationarity claim.
    limit = 100_000
    for iteration in range(1, limit + 1):
        following = step(current, c, precision, mode)
        assert current <= following <= HALF
        if following == current:
            return {"stationary_transition": iteration,
                    "first_stationary_value_iteration": iteration - 1,
                    "stationary_value": fraction_text(current)}
        current = following
    raise RuntimeError(f"Stationarity not established within limit at p={precision}, {mode}")


def certified_exact_escape(c):
    """Integer interval enclosure of the real orbit, with directed rounding."""
    scale = 1 << 320
    scaled_c = c * scale
    lower_c = scaled_c.numerator // scaled_c.denominator
    upper_c = -(-scaled_c.numerator // scaled_c.denominator)
    lower = upper = 0
    previous_upper = 0
    # Increment >= epsilon gives the independent finite termination bound.
    epsilon = c - QUARTER
    bound_fraction = 2 / epsilon
    bound = bound_fraction.numerator // bound_fraction.denominator + 1
    cap = min(bound, 100_000)
    for iteration in range(1, cap + 1):
        previous_upper = upper
        lower = lower * lower // scale + lower_c
        upper = -(-(upper * upper) // scale) + upper_c
        if lower > 2 * scale:
            assert previous_upper <= 2 * scale
            return {"first_escape_iteration": iteration,
                    "scale_bits": 320,
                    "previous_upper_numerator_hex": hex(previous_upper),
                    "escape_lower_numerator_hex": hex(lower),
                    "escape_upper_numerator_hex": hex(upper)}
        if upper > 2 * scale:
            raise RuntimeError("Interval straddles escape threshold before certification")
    raise RuntimeError("Exact escape not established within resource limit")


def native_round(value, precision):
    if precision == 53:
        return float(value)
    code = {11: "e", 24: "f"}[precision]
    return struct.unpack(">" + code, struct.pack(">" + code, float(value)))[0]


def native_separate_check(precision):
    """Compare native storage rounding against every exact simulated step.

    For p=11,24, the product of two p-bit values and the subsequent sum
    of the rounded product with c on these trajectories fit in binary64.
    Thus binary64 evaluation followed by e/f storage gives the specified
    two-rounding trajectory. p=53 uses Python binary64 operations directly.
    """
    c = parameter(precision)
    native_c = native_round(c, precision)
    assert Fraction.from_float(native_c) == c
    exact = ZERO
    native = 0.0
    limit = 100_000 if precision != 53 else 1_000
    stationary_at = None
    for iteration in range(1, limit + 1):
        expected = step(exact, c, precision, "separate")
        observed = native_round(native_round(native * native, precision) + native_c, precision)
        assert Fraction.from_float(observed) == expected
        assert observed <= 0.5
        if observed == native:
            stationary_at = iteration
            break
        exact, native = expected, observed
    # At p=53 only a prefix and the boundary are tested, not stationarity.
    boundary = native_round(native_round(0.5 * 0.5, precision) + native_c, precision)
    assert boundary == 0.5
    if precision == 53:
        assert native_c == math.nextafter(0.25, math.inf)
    else:
        assert stationary_at is not None
    return {"precision_bits": precision, "compared_transitions": iteration,
            "stationary_transition": stationary_at,
            "parameter_hex": native_c.hex(), "boundary_next_hex": boundary.hex()}


def exhaustive_binade_check(precision):
    """Check zero and every p-bit value in [1/4,1/2] for both maps.

    Nonzero orbit states are in this range after the first transition.
    There are 2**(p-1)+1 such values, including the upper binade boundary.
    """
    c = parameter(precision)
    quantum = power_two(-precision - 1)
    count = 0
    for numerator in range((1 << (precision - 1)) + 1):
        value = QUARTER + numerator * quantum
        for mode in ["separate", "fused"]:
            assert QUARTER <= step(value, c, precision, mode) <= HALF
            count += 1
    for mode in ["separate", "fused"]:
        assert step(ZERO, c, precision, mode) == c
        count += 1
    return count


def main():
    assert sys.float_info.radix == 2 and sys.float_info.mant_dig == 53
    cases = []
    for precision in range(3, 25):
        c = parameter(precision)
        assert nearest_even(c, precision) == c
        assert c > QUARTER
        assert nearest_even(HALF + power_two(-precision - 1), precision) == HALF
        cases.append({"precision_bits": precision, "parameter": fraction_text(c),
                      "separate": rounded_trajectory(precision, "separate"),
                      "fused": rounded_trajectory(precision, "fused"),
                      "exact": certified_exact_escape(c)})
    boundary_cases = []
    for precision in [11, 24, 53]:
        c = parameter(precision)
        boundary_cases.append({"precision_bits": precision, "parameter": fraction_text(c),
                               "separate_boundary": fraction_text(step(HALF, c, precision, "separate")),
                               "fused_boundary": fraction_text(step(HALF, c, precision, "fused"))})
    exhaustive_comparisons = sum(exhaustive_binade_check(p) for p in range(3, 13))
    native = [native_separate_check(p) for p in [11, 24, 53]]
    payload = {"case_count": len(cases), "precision_min": 3, "precision_max": 24,
               "all_stationary_values_half": all(
                   case[mode]["stationary_value"] == "1/2"
                   for case in cases for mode in ["separate", "fused"]),
               "exhaustive_transition_comparisons": exhaustive_comparisons,
               "cases": cases, "boundary_cases": boundary_cases, "native": native,
               "binary32_separate_stationary_transition": cases[-1]["separate"]["stationary_transition"],
               "binary32_fused_stationary_transition": cases[-1]["fused"]["stationary_transition"],
               "binary32_exact_first_escape": cases[-1]["exact"]["first_escape_iteration"]}
    results = Path("results")
    results.mkdir(exist_ok=True)
    (results / "R1.json").write_text(json.dumps(payload, indent=2) + "\n")
    print(json.dumps({k: payload[k] for k in ["case_count", "all_stationary_values_half",
                                             "binary32_separate_stationary_transition",
                                             "binary32_fused_stationary_transition",
                                             "binary32_exact_first_escape"]}))


if __name__ == "__main__":
    main()
