#!/usr/bin/env python3
"""Finite exact checks for the regular Sturm–Liouville capstone.

Python standard library only. Fractions and Taylor/Lagrange bounds certify
signs; no floating-point value is used in any proof decision. This does NOT
prove infinite-mode completeness or validate a general ODE solver.
"""
import argparse
import json
from fractions import Fraction as Q
from math import factorial, isqrt


def rational_argument(text):
    """Reject malformed rationals with an ordinary argparse diagnostic."""
    try:
        return Q(text)
    except (ValueError, ZeroDivisionError) as exc:
        raise argparse.ArgumentTypeError('expected a rational number with nonzero denominator') from exc


def ceilq(x):
    return -((-x.numerator) // x.denominator)


def floorq(x):
    return x.numerator // x.denominator


def atan_box(x, terms=32):
    """Alternating series, 0 < x < 1; sum of terms 0..terms-1."""
    if not 0 < x < 1 or terms < 1:
        raise ValueError('atan needs 0 < x < 1 and a positive term count')
    s = sum(((-1) ** j * x ** (2*j+1) / (2*j+1)
             for j in range(terms)), Q(0))
    nxt = (-1) ** terms * x ** (2*terms+1) / (2*terms+1)
    return min(s, s+nxt), max(s, s+nxt)


def pi_box():
    # Machin: pi = 16 atan(1/5) - 4 atan(1/239).
    # The capstone verifies the tangent identity and the angle branch.
    a, b = atan_box(Q(1, 5))
    c, d = atan_box(Q(1, 239))
    return 16*a-4*d, 16*b-4*c


def trig_boxes(x, degree_half=26):
    """Taylor sin through degree 2m+1; cos through 2m.

    Lagrange bounds |sin remainder| <= |x|^(2m+2)/(2m+2)!,
    |cos remainder| <= |x|^(2m+1)/(2m+1)! are global.
    """
    if degree_half < 1:
        raise ValueError('Taylor half-degree must be positive')
    m = degree_half
    s = sum(((-1)**j * x**(2*j+1) / factorial(2*j+1)
             for j in range(m+1)), Q(0))
    c = sum(((-1)**j * x**(2*j) / factorial(2*j)
             for j in range(m+1)), Q(0))
    rs = abs(x)**(2*m+2) / factorial(2*m+2)
    rc = abs(x)**(2*m+1) / factorial(2*m+1)
    return (s-rs, s+rs), (c-rc, c+rc)


def residual_box(k, terms=26):
    if k <= 0:
        raise ValueError('frequency must be positive')
    (sl, su), (cl, cu) = trig_boxes(k, terms)
    return k*sl-cu, k*su-cl


def signed_box(box, n):
    return box if n % 2 == 0 else (-box[1], -box[0])


def sign(box):
    return -1 if box[1] < 0 else (1 if box[0] > 0 else 0)


def rational_interval(box):
    return [str(x) for x in box]


def readable_interval(box, digits=12):
    """Outward decimal formatting using integer floor/ceiling, not float."""
    scale = 10**digits
    def fmt(n):
        sg = '-' if n < 0 else ''
        n = abs(n)
        return f'{sg}{n//scale}.{n%scale:0{digits}d}'
    return [fmt(floorq(box[0]*scale)), fmt(ceilq(box[1]*scale))]


INITIAL = [(Q(43,50), Q(861,1000)),
           (Q(137,40), Q(1713,500)),
           (Q(6437,1000), Q(3219,500))]


def robin_certificate(n, refinements=0, terms=26):
    lo, hi = INITIAL[n]
    pl, pu = pi_box()
    # Prove the entire bracket is in the nth positive-tangent branch.
    branch = n*pu < lo and hi < (Q(n)+Q(1,2))*pl
    left, right = signed_box(residual_box(lo, terms), n), signed_box(residual_box(hi, terms), n)
    if not branch or sign(left) != -1 or sign(right) != 1:
        return {'index': n, 'status': 'undecided',
                'reason': 'Taylor bounds do not certify both endpoint signs',
                'k_interval': rational_interval((lo,hi)),
                'left_signed_residual': rational_interval(left),
                'right_signed_residual': rational_interval(right)}
    done = 0
    for _ in range(refinements):
        mid = (lo+hi)/2
        s = sign(signed_box(residual_box(mid, terms), n))
        if s == 0:
            break  # Keep the last certified outer interval; never guess a sign.
        if s < 0:
            lo = mid
        else:
            hi = mid
        done += 1
    lam = (lo*lo-pu*pu, hi*hi-pl*pl)
    return {'index': n, 'status': 'certified',
            'requested_refinements': refinements, 'completed_refinements': done,
            'refinement_stopped_unresolved': done < refinements,
            'k_interval': rational_interval((lo,hi)),
            'k_decimal_outward': readable_interval((lo,hi)),
            'lambda_interval': rational_interval(lam),
            'lambda_decimal_outward': readable_interval(lam),
            'k_width': str(hi-lo), 'branch_membership': branch,
            'left_signed_residual': rational_interval(signed_box(residual_box(lo,terms),n)),
            'right_signed_residual': rational_interval(signed_box(residual_box(hi,terms),n)),
            'uniqueness_basis': 'analytic strict monotonicity of k tan k on (n*pi,n*pi+pi/2)'}


def nn_count(cut):
    """cut is lambda/pi^2, exactly rational; count n^2-1 < or <= cut."""
    if abs(cut) > 10**6:
        raise ValueError('absolute normalized cut must not exceed 1000000')
    if cut < -1:
        return {'strict': 0, 'closed': 0, 'equal_index': None}
    # Largest integer n with n^2 <= cut+1, using exact integer arithmetic.
    top = isqrt(floorq(cut+1))
    equal = Q(top*top-1) == cut
    return {'strict': top if equal else top+1,
            'closed': top+1, 'equal_index': top if equal else None}


def phase_count_box(low, high, beta=Q(1,2)):
    """Input interval is already a VALIDATED enclosure of theta(b)/pi.

    The caller, not this routine, must supply the ODE defect certificate.
    beta denotes the right boundary angle divided by pi.
    """
    if low > high or not 0 < beta <= 1 or low < 0:
        raise ValueError('require 0 <= low <= high and 0 < beta/pi <= 1')
    # Any allowed eigen-threshold beta+n touching the CLOSED enclosure blocks.
    first = max(0, ceilq(low-beta))
    if Q(first)+beta <= high:
        return {'status': 'undecided', 'reason': 'enclosure touches a spectral threshold',
                'threshold': str(Q(first)+beta)}
    count = max(0, ceilq(low-beta))
    if count != max(0, ceilq(high-beta)):
        raise ArithmeticError("internal threshold-free interval count mismatch")
    return {'status': 'certified', 'strict': count, 'closed': count,
            'assumption': 'input is a previously validated normalized phase enclosure'}


def self_test():
    checks = []
    def check(name, ok):
        checks.append({'name': name, 'passed': bool(ok)})
        if not ok:
            raise AssertionError(name)
    pl,pu = pi_box()
    check('pi bracket positive and narrower than 10^-40', 3 < pl < pu < Q(22,7) and pu-pl < Q(1,10**40))
    tan2 = 2*Q(1,5)/(1-Q(1,5)**2)
    tan4 = 2*tan2/(1-tan2**2)
    check('tan(4 atan(1/5)-atan(1/239))=1', (tan4-Q(1,239))/(1+tan4*Q(1,239)) == 1)
    for n in range(3):
        cert = robin_certificate(n, 12)
        check(f'Robin root {n} exact bracket and refinement', cert['status']=='certified' and cert['completed_refinements']==12)
        l,r = map(Q,cert['k_interval'])
        check(f'Robin root {n} refined width', r-l == (INITIAL[n][1]-INITIAL[n][0])/2**12)
        ll,rr = map(Q,cert['lambda_interval'])
        check(f'Robin root {n} sign of eigenvalue', rr<0 if n==0 else ll>0)
    check('insufficient Taylor degree reports undecided', robin_certificate(2,terms=1)['status']=='undecided')
    check('NN below ground', nn_count(Q(-2))=={'strict':0,'closed':0,'equal_index':None})
    for n in range(11):
        c=nn_count(Q(n*n-1))
        check(f'NN strict/closed at mode {n}', c=={'strict':n,'closed':n+1,'equal_index':n})
    for c,e in [(Q(-1,2),1),(Q(1,10),2),(Q(7,2),3)]:
        check('NN non-spectral cut '+str(c), nn_count(c)['strict']==e and nn_count(c)['closed']==e)
    a,b=nn_count(Q(-1)),nn_count(Q(0))
    check('open interval (-pi^2,0) empty', b['strict']-a['closed']==0)
    check('closed interval [-pi^2,0] has two', b['closed']-a['strict']==2)
    check('phase interval away from thresholds', phase_count_box(Q(7,4),Q(9,4))['strict']==2)
    for low,high in [(Q(149,100),Q(151,100)),(Q(3,2),Q(8,5)),(Q(7,5),Q(3,2)),(Q(3,2),Q(3,2))]:
        check('phase touching threshold is undecided '+str((low,high)), phase_count_box(low,high)['status']=='undecided')
    check('phase lowest range zero', phase_count_box(Q(0),Q(1,4))['strict']==0)
    eps=Q(1,10**12)
    check('naive floating-like strict count fails above equality', ceilq(1+eps)==2)
    check('naive floating-like closed count fails below equality', floorq(1-eps)+1==1)
    for c in [Q(-1),Q(0),Q(3)]:
        check('exact NN spectral equality remains distinguished '+str(c), nn_count(c)['closed']-nn_count(c)['strict']==1)
    try:
        phase_count_box(Q(2),Q(1))
        check('invalid phase interval rejected',False)
    except ValueError:
        check('invalid phase interval rejected',True)
    return {'status':'passed','assertions':len(checks),'checks':checks,
            'scope':'finite exact identities, certified signs and counting conventions only; no infinite completeness claim'}


def main():
    p=argparse.ArgumentParser(description=__doc__)
    p.add_argument('--refine',type=int,default=0,help='certified bisections, 0..40')
    p.add_argument('--terms',type=int,default=26,help='Taylor half-degree, 1..60')
    p.add_argument('--nn-cut',type=rational_argument,default=Q(0),help='exact lambda/pi^2, e.g. 3 or 1/2')
    p.add_argument('--phase-box',nargs=2,type=rational_argument,metavar=('LOW','HIGH'),help='previously validated enclosure of theta(b)/pi')
    p.add_argument('--beta',type=rational_argument,default=Q(1,2),help='right boundary angle/pi for --phase-box')
    p.add_argument('--self-test',action='store_true')
    args=p.parse_args()
    if not 0<=args.refine<=40 or not 1<=args.terms<=60:
        p.error('require 0 <= refine <= 40 and 1 <= terms <= 60')
    try:
        if args.self_test:
            result=self_test()
        else:
            result={'scope':'finite exact capstone certificate; analytic proofs establish completeness, branch uniqueness and general phase theorem',
                    'pi_interval':rational_interval(pi_box()),
                    'nn':{'cut_lambda_over_pi_squared':str(args.nn_cut),**nn_count(args.nn_cut)},
                    'robin':[robin_certificate(n,args.refine,args.terms) for n in range(3)]}
            if args.phase_box:
                result['phase']=phase_count_box(*args.phase_box,beta=args.beta)
        print(json.dumps(result,ensure_ascii=False,indent=2))
    except ValueError as exc:
        p.error(str(exc))


if __name__=='__main__':
    main()
