#!/usr/bin/env python3
"""Quantify reviewed subset-sum guards and conditional amplification costs."""
import argparse
import json
import math
from fractions import Fraction
from pathlib import Path

MAX_N = 10**12
MAX_B = 4096
MAX_REPETITIONS = 2047


def integer(value, low, high, name):
    if type(value) is not int or not low <= value <= high:
        raise ValueError('Invalid ' + name)
    return value


def rational(value):
    if not isinstance(value, str) or len(value) > 128:
        raise ValueError('Require canonical rational confidence error')
    try:
        number = Fraction(value)
    except (ValueError, ZeroDivisionError):
        raise ValueError('Invalid confidence error') from None
    if str(number) != value or not Fraction(1, 10**30) <= number < 1:
        raise ValueError('Confidence error must be canonical, at least 1e-30 and below one')
    return number


def majority_error(repetitions):
    integer(repetitions, 1, MAX_REPETITIONS, 'repetitions')
    if repetitions % 2 == 0:
        raise ValueError('Require odd majority repetitions')
    r = repetitions
    numerator = sum(math.comb(r, k) * (1 << (r - k)) for k in range((r + 1) // 2, r + 1))
    return Fraction(numerator, 3**r)


def majority_plan(call_cap, delta):
    target = delta / call_cap
    low, high = 0, (MAX_REPETITIONS - 1) // 2
    if majority_error(2 * high + 1) > target:
        return dict(status='budget_exceeded', repetitions=None, conditional_all_call_error_bound=None)
    while low < high:
        middle = (low + high) // 2
        if majority_error(2 * middle + 1) <= target:
            high = middle
        else:
            low = middle + 1
    repetitions = 2 * low + 1
    per_call = majority_error(repetitions)
    return dict(status='complete', repetitions=repetitions, per_call_error_bound=str(per_call),
                conditional_all_call_error_bound=str(min(Fraction(1), call_cap * per_call)),
                maximum_base_oracle_calls=call_cap * repetitions)


def calculate(record):
    if not isinstance(record, dict) or set(record) != {'n', 'maximum_input_bit_length', 'witness_error_budget'}:
        raise ValueError('Invalid schedule schema')
    n = integer(record['n'], 2, MAX_N, 'item count')
    b = integer(record['maximum_input_bit_length'], 1, MAX_B, 'bit length')
    delta = rational(record['witness_error_budget'])
    u = n + b + 2
    # At most 100000*12 bits: this guard computation is bounded by MAX_B,
    # and never constructs 2**n for a huge declared n.
    fast_guard_minimum_n = (b**100000 - 1).bit_length()
    fast_guard_passes = n >= fast_guard_minimum_n
    low_guard_passes = (u - 1).bit_length() <= n // 10**9
    q = 0
    while q * 10**9 + b + 2 > 1 << q:
        q += 1
        if q > 128:
            raise ArithmeticError('Bounded low-space guard search exceeded')
    call_cap = n + 1
    two_sided = majority_plan(call_cap, delta)
    one_sided_repetitions = 1
    while Fraction(call_cap, 3**one_sided_repetitions) > delta:
        one_sided_repetitions += 1
    # Exact logarithmic enclosure of the displayed main schedule's cap.
    # Multiplication by 2**ceil(n/2) remains symbolic.
    coefficient = 2 * u**620
    step_floor_log2 = coefficient.bit_length() - 1 + (n + 1) // 2
    rounds = u**120
    word_lower = 4 * (n + b)
    word_upper = 4 * (n + b + (n + 1).bit_length())
    ideal_time_exponent = Fraction(n, 100)
    ideal_space_exponent = Fraction(n, 20)
    return dict(status='complete', n=n, maximum_input_bit_length=b, source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb',
                fast_decision=dict(known_guard='b^100000 <= 2^n', known_guard_passes=fast_guard_passes,
                                   minimum_n_for_this_known_guard=fast_guard_minimum_n,
                                   fixed_source_cutoff_n_star='effective but no evaluated numerical value recorded',
                                   main_branch_eligibility='unknown' if fast_guard_passes else 'known_guard_failed',
                                   fallback='exact meet-in-the-middle', error_model='two_sided_at_most_one_third_per_fixed_input',
                                   main_time_formula='K0*2^(0.489995*n)*(n+b+2)^K; K0,K unselected here',
                                   ordinary_time_claim='O_c(2^(0.49*n)) for every fixed polynomial-bit class, eventually'),
                low_space_decision=dict(known_guard='n+b+2 <= 2^floor(n/1000000000)', known_guard_passes=low_guard_passes,
                                       minimum_n_for_this_known_guard=q * 10**9,
                                       fixed_source_cutoff_n_zero='no evaluated numerical value recorded',
                                       main_branch_eligibility='unknown' if low_guard_passes else 'known_guard_failed',
                                       fallback='exact indexed subset enumeration with O(n) writable words',
                                       error_model='no_false_YES; each_YES_found_with_probability_at_least_two_thirds',
                                       rounds_formula='(n+b+2)^120', rounds_floor_log2=rounds.bit_length() - 1,
                                       total_trial_cap_formula='2*(n+b+2)^620*2^ceil(n/2)',
                                       total_trial_cap_log2_interval=[step_floor_log2, step_floor_log2 + 1],
                                       space_formula='A_c*(n+b+2)^15*2^(0.199*n) writable words',
                                       ordinary_space_claim='O_c(2^(n/5)); polynomial factor remains at exponent 0.199'),
                ideal_exponent_only_comparison=dict(time_vs_half_exponent_log2_factor=str(ideal_time_exponent),
                                                    time_factor_if_at_most_2_to_60=2.0**float(ideal_time_exponent) if ideal_time_exponent <= 60 else None,
                                                    space_vs_quarter_exponent_log2_factor=str(ideal_space_exponent),
                                                    ignored_factors='All constants, polynomial overhead, thresholds, arithmetic and actual implementation; neither comparison is a measured speedup.'),
                word_width_enclosure_bits=[word_lower, word_upper],
                conditional_witness_plan=dict(maximum_adaptive_decision_questions=call_cap, error_budget=str(delta),
                                               fast_two_sided_majority=two_sided,
                                               low_space_one_sided_OR=dict(repetitions=one_sided_repetitions,
                                                                         conditional_all_call_error_bound=str(Fraction(call_cap, 3**one_sided_repetitions)),
                                                                         maximum_base_oracle_calls=call_cap * one_sided_repetitions),
                                               assumptions='Fresh independent runs conditional on each adaptive history; correct source/input/machine correspondence and proved base error <=1/3 are prerequisites, not checked by this planner.',
                                               final_witness_check='Check distinct original item IDs and exact integer sum. Invalid recovered output is unknown, never an exact negative certificate.',
                                               negative_answer='Amplified randomized NO remains probabilistic unless separately established by a complete exact method.'),
                source_algorithms_implemented=False, source_proof_verification='not_run', formal_program_verification='not_run',
                interpretation='Exact bounded guard and error-budget algebra for supplied n,b. Displayed caps are permitted upper schedules, not runtime lower bounds. Unknown source cutoffs prevent an implementation-ready main-branch badge. Separate faster-time and lower-space results cannot be combined into one backend guarantee.')


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--output', type=Path)
    args = parser.parse_args()
    body = json.dumps(calculate(json.loads(args.input.read_text())), 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()
