#!/usr/bin/env python3
"""Exact, dependency-free teaching replay for coherent delayed diagonal quadratics.

Python 3.10+. Default: solved examples and exhaustive finite self-tests.
Custom: --input FILE; rational JSON values are integers or strings such as "1/24".
This reader verifies a specified positive diagonal quadratic, not arbitrary f.
Rows retain the full teaching trace (O(T*d) log space); the replay state itself
keeps at most tau+1 snapshots and tau steps. No stochastic/worker simulation.
"""
import argparse
from collections import deque
from fractions import Fraction as F
import itertools
import json
import re
import sys


def need(condition, message):
    if not condition:
        raise ValueError(message)


def rat(value):
    need(not isinstance(value, bool) and isinstance(value, (int, str)),
         'rational must be an integer or fraction string; floats/bools are rejected')
    if isinstance(value, int):
        need(value.bit_length() <= 4096, 'input rational exceeds 4096-bit teaching-reader limit')
    else:
        need(len(value) <= 4096, 'input rational exceeds bounded teaching-reader token length')
        value = value.strip()
        need(re.fullmatch(r'[+-]?(?:[0-9]+(?:/[0-9]+)?|[0-9]+\.[0-9]*|\.[0-9]+)', value) is not None,
             'rational string must be an integer, decimal, or fraction without exponent notation')
    try:
        result = F(value)
    except (ValueError, ZeroDivisionError, OverflowError) as error:
        raise ValueError('invalid finite rational: ' + str(value)) from error
    need(result.numerator.bit_length() <= 4096 and result.denominator.bit_length() <= 4096,
         'input rational exceeds 4096-bit teaching-reader limit')
    return result


def integer(value, name):
    need(type(value) is int and value >= 0, name + ' must be a nonnegative integer')
    return value


def norm2(x):
    return sum((v*v for v in x), F(0))


def objective(x, diagonal):
    return sum((v*v*l for v, l in zip(x, diagonal)), F(0))/2


def encode(value):
    if isinstance(value, F):
        return str(value)
    if isinstance(value, dict):
        return {key: encode(v) for key, v in value.items()}
    if isinstance(value, (list, tuple)):
        return [encode(v) for v in value]
    return value


def budget(B0, epsilon, q, limit=100000):
    need(B0 >= 0 and epsilon > 0 and 0 < q < 1, 'invalid budget parameters')
    bound, n = B0, 0
    while bound > epsilon and n < limit:
        bound *= q
        n += 1
    return {'updates': n if bound <= epsilon else None,
            'status': 'certified' if bound <= epsilon else 'reader_budget_limit',
            'bound': bound, 'previous_bound': bound/q if n else None,
            'query_cost_excludes_initial_and_final_certificates': True}


def replay(spec):
    allowed = {'diagonal', 'x0', 'alpha', 'tau', 'read_versions', 'epsilon', 'verify_current'}
    need(isinstance(spec, dict) and not (set(spec)-allowed), 'unknown input fields')
    need({'diagonal', 'x0', 'alpha', 'tau', 'read_versions'} <= set(spec), 'missing input fields')
    need(isinstance(spec['diagonal'], list) and isinstance(spec['x0'], list), 'vectors must be lists')
    diagonal = tuple(map(rat, spec['diagonal']))
    x0 = tuple(map(rat, spec['x0']))
    need(0 < len(x0) == len(diagonal) <= 100, 'vector dimensions must match, in 1..100')
    need(min(diagonal) > 0, 'diagonal must be positive: exact strong convexity check')
    alpha = rat(spec['alpha'])
    need(alpha > 0, 'alpha must be positive')
    tau = integer(spec['tau'], 'tau')
    need(tau <= 10000, 'tau exceeds teaching-reader limit')
    reads = spec['read_versions']
    need(isinstance(reads, list) and len(reads) <= 10000, 'read_versions must be a list of at most 10000')
    for k, r in enumerate(reads):
        integer(r, 'read version')
        need(max(0, k-tau) <= r <= k, 'invalid read version at commit ' + str(k))
    verify = spec.get('verify_current', True)
    need(type(verify) is bool, 'verify_current must be boolean')
    epsilon = rat(spec['epsilon']) if 'epsilon' in spec else None
    if epsilon is not None:
        need(epsilon > 0, 'epsilon must be positive')
    mu, L = min(diagonal), max(diagonal)
    certified = alpha <= F(1, 2)/(L*(tau+1))
    q, c = 1-alpha*mu, alpha*L*L*tau
    x, snapshots, steps = x0, {0: x0}, deque(maxlen=tau or 1)
    delta0 = objective(x0, diagonal)
    previous_energy = delta0
    rows = []
    peak_snapshots = 1
    for k, r in enumerate(reads):
        need(r in snapshots, 'snapshot retention invariant failed')
        old = snapshots[r]
        current = x
        gradient = tuple(l*v for l, v in zip(diagonal, old))
        step = tuple(-alpha*v for v in gradient)
        x = tuple(v+s for v, s in zip(x, step))
        need(max((v.numerator.bit_length()+v.denominator.bit_length() for v in x), default=0) <= 200000,
             'computed rational exceeds teaching-reader arithmetic limit')
        steps.appendleft(step)
        H = c*sum(((tau-i)*norm2(s) for i, s in enumerate(steps) if i < tau), F(0))
        delta = objective(x, diagonal)
        energy = delta+H
        envelope = q**(k+1)*delta0 if certified else None
        if certified:
            need(energy <= q*previous_energy, 'energy contraction invariant failed')
            need(delta <= envelope, 'objective envelope invariant failed')
        rows.append({'commit': k, 'read_version': r, 'delay': k-r,
                     'x_current': current, 'read_point': old, 'gradient': gradient, 'step': step, 'x_next': x,
                     'gap': delta, 'history': H, 'energy': energy,
                     'energy_contraction': energy <= q*previous_energy if certified else None,
                     'theorem_envelope': envelope})
        previous_energy = energy
        # Once k+1 is committed, versions before k+1-tau cannot be read next.
        snapshots.pop(k-tau, None)
        snapshots[k+1] = x
        peak_snapshots = max(peak_snapshots, len(snapshots))
        need(len(snapshots) <= tau+1, 'snapshot bound failed')
    current_gradient = tuple(l*v for l, v in zip(diagonal, x)) if verify else None
    residual_bound = norm2(current_gradient)/(2*mu) if verify else None
    envelope = q**len(reads)*delta0 if certified else None
    status = 'budget_complete_with_apriori_bound' if certified else 'parameters_not_certified_by_delay_theorem'
    if verify and epsilon is not None and residual_bound <= epsilon:
        status = 'current_point_residual_certified'
    return {'model': 'exact_positive_diagonal_quadratic', 'mu': mu, 'L': L,
            'alpha': alpha, 'tau': tau, 'q': q if certified else None,
            'delay_theorem_parameters_certified': certified, 'status': status,
            'x_final': x, 'gap_final': objective(x, diagonal), 'initial_gap': delta0,
            'theorem_bound': envelope, 'current_gradient': current_gradient,
            'current_point_gap_certificate': residual_bound,
            'accepted_update_gradient_queries': len(reads),
            'final_current_gradient_queries': int(verify),
            'total_full_gradient_queries': len(reads)+int(verify),
            'diagnostic_objective_evaluations': len(reads)+1,
            'history_diagnostics_vector_work': 'O(T*tau*d), additional to update work',
            'initial_gap_known_from_explicit_quadratic': True,
            'peak_retained_snapshots': peak_snapshots,
            'teaching_log_rows': len(rows), 'rows': rows}


def scalar(a, tau, n):
    return replay({'diagonal': [1], 'x0': [1], 'alpha': str(a), 'tau': tau,
                   'read_versions': [max(0, k-tau) for k in range(n)], 'verify_current': False})


def examples():
    base = replay({'diagonal': [1, 4], 'x0': [1, 1], 'alpha': '1/24', 'tau': 2,
                   'read_versions': [0, 0, 0, 3, 2, 4], 'epsilon': '1/100'})
    need(base['x_final'] == (F(3527, 4608), F(17, 72)), 'capstone endpoint')
    need(base['gap_final'] == F(17174705, 42467328), 'capstone gap')
    need(base['rows'][-1]['history'] == F(91627, 2654208), 'capstone history')
    need(base['rows'][-1]['energy'] == F(2071193, 4718592), 'capstone energy')
    tested = 0
    for tau in range(4):
        for reads in itertools.product(*[range(max(0, k-tau), k+1) for k in range(7)]):
            replay({'diagonal': [1, 4], 'x0': [1, 1], 'alpha': str(F(1, 8*(tau+1))),
                    'tau': tau, 'read_versions': list(reads), 'verify_current': False})
            tested += 1
    need(tested == 2087, 'schedule enumeration count')
    periodic = scalar(F(1), 1, 12)
    sequence = [F(1)]+[r['x_next'][0] for r in periodic['rows']]
    need(sequence[:6] == sequence[6:12], 'period six')
    need(sequence[1] == 0 and sequence[2] == -1, 'stale zero stopping failure')
    nonmonotone = scalar(F(1, 6), 2, 14)
    need(nonmonotone['rows'][12]['x_next'] == (F(-7, 7776),), 'x13')
    need(nonmonotone['x_final'] == (F(-1, 432),), 'x14')
    need(nonmonotone['rows'][13]['gap'] > nonmonotone['rows'][12]['gap'], 'nonmonotone objective')
    rates = []
    for tau, expected in [(0, 42), (1, 86), (2, 130), (5, 263)]:
        alpha = F(1, 8*(tau+1))
        b = budget(F(5, 2), F(1, 100), 1-alpha)
        need(b['updates'] == expected, 'exact budget')
        need(b['bound'] <= F(1, 100) < b['previous_bound'], 'minimal geometric budget')
        rates.append({'tau': tau, 'alpha': alpha, 'q': 1-alpha, **b})
    invalid = []
    good = {'diagonal': [1, 4], 'x0': [1, 1], 'alpha': '1/24', 'tau': 2, 'read_versions': [0]}
    mutations = [{'read_versions': [1]}, {'read_versions': [0, 2]}, {'read_versions': [-1]},
                 {'read_versions': [0, 0, 0, 0]}, {'read_versions': [False]}, {'tau': -1},
                 {'tau': True}, {'alpha': '0'}, {'alpha': 'NaN'}, {'alpha': '1/0'},
                 {'alpha': 0.1}, {'diagonal': [0, 4]}, {'x0': [1]}, {'verify_current': 1},
                 {'epsilon': '-1'}, {'unexpected': 1}]
    for changes in mutations:
        try:
            replay({**good, **changes})
        except ValueError as error:
            invalid.append({'input_change': changes, 'rejection': str(error)})
        else:
            raise ValueError('invalid input was accepted: ' + str(changes))
    zero = replay({**good, 'read_versions': [], 'epsilon': '3', 'verify_current': False})
    need(zero['total_full_gradient_queries'] == 0 and zero['x_final'] == (F(1), F(1)), 'zero budget')
    at_minimum = replay({**good, 'x0': [0, 0], 'epsilon': '1/100'})
    need(at_minimum['status'] == 'current_point_residual_certified', 'zero residual')
    unsafe = scalar(F(3, 2), 1, 12)
    need(not unsafe['delay_theorem_parameters_certified'], 'unstable run wrongly certified')
    return {'scope': 'Exact finite replay and implementation checks; the article proves the all-history theorem.',
            'base': base, 'exhaustive_seven_commit_schedules': tested,
            'periodic_boundary': sequence, 'unstable_one_delay': unsafe,
            'nonmonotone_inside_certificate': nonmonotone,
            'fixed_delay_transfer': {'a': '3/4', 'tau1': 'stable: a<1',
                'tau2': 'unstable: a>(sqrt(5)-1)/2',
                'repair_for_variable_tau2': 'alpha=1/6 at lambda=1'},
            'budget_transfer': rates, 'invalid_inputs_rejected': invalid,
            'zero_budget': zero, 'zero_residual': at_minimum,
            'exact_arithmetic_warning': 'Rational bit complexity is additional to vector/query counts.'}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--input', help='JSON input path, or - for stdin')
    args = parser.parse_args()
    try:
        if args.input:
            if args.input == '-':
                spec = json.load(sys.stdin)
            else:
                with open(args.input, encoding='utf-8') as handle:
                    spec = json.load(handle)
            result = replay(spec)
        else:
            result = examples()
        print(json.dumps(encode(result), ensure_ascii=False, indent=2))
    except (ValueError, OSError, TypeError, KeyError, OverflowError) as error:
        print(json.dumps({'status': 'invalid_input_or_arithmetic_failure', 'reason': str(error)}, ensure_ascii=False))
        return 2
    return 0


if __name__ == '__main__':
    sys.exit(main())
