"""Certified first-escape counts for an explicit finite-cutoff counterexample family.

Run from the bundle root: python code/verify.py
The two enclosure methods share the recurrence and inputs, but not arithmetic providers.
"""
import json
import pathlib
from fractions import Fraction

from flint import arb, ctx

BITS = 192
BOUND = Fraction(17, 32)


def parameter(cutoff):
    return Fraction(1, 4) + Fraction(1, 16 * (cutoff + 1) ** 2)


def ceiling_div(numerator, denominator):
    return -(-numerator // denominator)


def integer_enclosure(cutoff, c):
    """Outward-rounded positive fixed-point arithmetic, using only integer operations."""
    scale = 1 << BITS
    c_low = c.numerator * scale // c.denominator
    c_high = ceiling_div(c.numerator * scale, c.denominator)
    low = high = 0
    cutoff_high = None
    exact = Fraction(0)
    controls = 0
    # The analytic growth bound in the paper gives this finite upper bound.
    upper_steps = 32 * (cutoff + 1) ** 2 + 1
    for iteration in range(1, upper_steps + 1):
        previous_high = high
        low = low * low // scale + c_low
        high = ceiling_div(high * high, scale) + c_high
        if iteration <= 7:
            exact = exact * exact + c
            assert low * exact.denominator <= exact.numerator * scale <= high * exact.denominator
            controls += 1
        if iteration <= cutoff:
            assert high * BOUND.denominator <= BOUND.numerator * scale
        if iteration == cutoff:
            cutoff_high = high
        if low > 2 * scale:
            assert previous_high <= 2 * scale
            assert iteration > cutoff and cutoff_high is not None
            return {
                'first_escape': iteration,
                'exact_rational_control_checks': controls,
                'cutoff_upper_numerator_hex': hex(cutoff_high),
                'pre_escape_upper_numerator_hex': hex(previous_high),
                'escape_lower_numerator_hex': hex(low),
                'escape_upper_numerator_hex': hex(high),
                'enclosure_denominator': f'2^{BITS}',
            }
        if high > 2 * scale:
            raise ArithmeticError('Integer enclosure cannot decide this iterate at the fixed precision')
    raise ArithmeticError('No certified escape within the proved upper bound')


def arb_enclosure(cutoff, c):
    """Arb ball arithmetic, with its own directed error propagation."""
    ctx.prec = BITS
    c_ball = arb(c.numerator) / arb(c.denominator)
    z = arb(0)
    bound = arb(17) / 32
    upper_steps = 32 * (cutoff + 1) ** 2 + 1
    for iteration in range(1, upper_steps + 1):
        previous_upper = z.upper()
        z = z * z + c_ball
        if iteration <= cutoff:
            assert z.upper() <= bound
        if z.lower() > 2:
            assert previous_upper <= 2
            assert iteration > cutoff
            return iteration
        if z.upper() > 2:
            raise ArithmeticError('Arb enclosure cannot decide this iterate at the fixed precision')
    raise ArithmeticError('No certified escape within the proved upper bound')


def main():
    inputs = json.loads(pathlib.Path('data/cases.json').read_text())
    cutoffs = inputs['cutoffs']
    assert cutoffs and all(type(t) is int and t >= 1 for t in cutoffs)
    assert len(set(cutoffs)) == len(cutoffs)
    rows = []
    for cutoff in cutoffs:
        c = parameter(cutoff)
        integer = integer_enclosure(cutoff, c)
        ball_count = arb_enclosure(cutoff, c)
        assert ball_count == integer['first_escape']
        rows.append({
            'cutoff': cutoff,
            'parameter_numerator': str(c.numerator),
            'parameter_denominator': str(c.denominator),
            'arb_first_escape': ball_count,
            'integer_certificate': integer,
        })
    result = {
        'cutoffs': cutoffs,
        'arithmetic_bits': BITS,
        'certified_cases': len(rows),
        'agreeing_cases': sum(r['arb_first_escape'] == r['integer_certificate']['first_escape'] for r in rows),
        'escape_iterations': [r['arb_first_escape'] for r in rows],
        'all_escaped_after_cutoff': all(r['arb_first_escape'] > r['cutoff'] for r in rows),
        'exact_rational_control_checks': sum(r['integer_certificate']['exact_rational_control_checks'] for r in rows),
        'cases': rows,
    }
    pathlib.Path('results').mkdir(exist_ok=True)
    pathlib.Path('results/R1.json').write_text(json.dumps(result, indent=2) + '\n')
    print(json.dumps({k:v for k,v in result.items() if k != 'cases'}))


if __name__ == '__main__':
    main()
