#!/usr/bin/env python3
"""Exact finite wave certificates. Standard library only; --output is explicit.

Polynomial identities certify the displayed finite models. Grid checks diagnose
implementation errors; the general propagation theorems are proved in the pages.
All acceptance comparisons use Fraction, including outward radical intervals.
"""
from argparse import ArgumentParser
from collections import Counter
from dataclasses import dataclass
from fractions import Fraction as Q
from math import comb, factorial, isqrt
from pathlib import Path
import json

CHECKS = Counter()

def check(condition, label):
    if not condition:
        raise ValueError(label)
    CHECKS[label] += 1

# Univariate polynomials, ascending powers. Zero is represented by (0,).
def trim(p):
    p = list(map(Q, p))
    while len(p) > 1 and p[-1] == 0:
        p.pop()
    return tuple(p or [Q(0)])

def add(p, q):
    return trim([(p[i] if i < len(p) else 0) +
                 (q[i] if i < len(q) else 0)
                 for i in range(max(len(p), len(q)))])

def scale(p, a):
    return trim([Q(a) * x for x in p])

def mul(p, q):
    r = [Q(0)] * (len(p) + len(q) - 1)
    for i, a in enumerate(p):
        for j, b in enumerate(q):
            r[i+j] += a*b
    return trim(r)

def power(p, n):
    r = (Q(1),)
    for _ in range(n):
        r = mul(r, p)
    return r

def derivative(p, n=1):
    for _ in range(n):
        p = trim([i*p[i] for i in range(1, len(p))])
    return p

def evaluate(p, x):
    v = Q(0)
    for a in reversed(p):
        v = v*x + a
    return v

def compose(p, q):
    r = (Q(0),)
    for a in reversed(p):
        r = add(mul(r, q), (a,))
    return r

def integrate(p, a, b):
    antiderivative = (Q(0),) + tuple(v/Q(i+1) for i, v in enumerate(p))
    return evaluate(antiderivative, b)-evaluate(antiderivative, a)

PULSE = power((1, 0, -1), 3)
PRIMITIVE = (Q(0), Q(1), Q(0), -Q(1), Q(0), Q(3,5), Q(0), -Q(1,7))
SATURATION = Q(16,35)

def pulse(x):
    return evaluate(PULSE, x) if abs(x) < 1 else Q(0)

def primitive(x):
    return -SATURATION if x <= -1 else SATURATION if x >= 1 else evaluate(PRIMITIVE, x)

def reflected_pulse(t, x, a, h, amplitude, c=Q(2)):
    t, x, a, h, amplitude, c = map(Q, (t, x, a, h, amplitude, c))
    if t < 0 or x < 0 or h <= 0 or c <= 0 or a < h:
        raise ValueError('half-line input contract')
    return amplitude*h/(2*c)*(primitive((x+c*t-a)/h)-primitive((x-c*t-a)/h)
                               -primitive((x+c*t+a)/h)+primitive((x-c*t+a)/h))

def detector(t):
    return reflected_pulse(t, 9, 2, 1, 1) + reflected_pulse(t, 9, 5, 2, Q(-1,2))

def detector_polynomial(midpoint):
    result = (Q(0),)
    for a, h, amplitude in [(Q(2),Q(1),Q(1)), (Q(5),Q(2),Q(-1,2))]:
        for offset, slope, sign in [(9-a, 2, 1), (9-a, -2, -1),
                                    (9+a, 2, -1), (9+a, -2, 1)]:
            affine = (offset/h, Q(slope)/h)
            at_mid = evaluate(affine, midpoint)
            value = (-SATURATION,) if at_mid < -1 else (SATURATION,) if at_mid > 1 else compose(PRIMITIVE, affine)
            result = add(result, scale(value, amplitude*h*sign/4))
    return result

@dataclass(frozen=True)
class Interval:
    lo: Q
    hi: Q
    def __post_init__(self):
        if self.lo > self.hi:
            raise ValueError('reversed interval')
    def __add__(self, other):
        b = as_interval(other)
        return Interval(self.lo+b.lo, self.hi+b.hi)
    __radd__ = __add__
    def __neg__(self):
        return Interval(-self.hi, -self.lo)
    def __sub__(self, other):
        return self + -as_interval(other)
    def __mul__(self, other):
        b = as_interval(other)
        values = [self.lo*b.lo, self.lo*b.hi, self.hi*b.lo, self.hi*b.hi]
        return Interval(min(values), max(values))
    __rmul__ = __mul__
    def reciprocal(self):
        if self.lo <= 0 <= self.hi:
            raise ValueError('zero in denominator')
        return Interval(1/self.hi, 1/self.lo)
    def __truediv__(self, other):
        return self * as_interval(other).reciprocal()

def as_interval(x):
    return x if isinstance(x, Interval) else Interval(Q(x), Q(x))

def sqrt_interval(value, digits=70):
    value = Q(value)
    if value < 0:
        raise ValueError('negative radicand')
    denominator = 10**digits
    k = isqrt(value.numerator*denominator**2 // value.denominator)
    lo = Q(k, denominator)
    hi = lo if lo*lo == value else Q(k+1, denominator)
    check(lo*lo <= value <= hi*hi, 'integer_square_root_enclosure')
    return Interval(lo, hi)

def interval_evaluate(p, x):
    v = as_interval(0)
    for a in reversed(p):
        v = v*x + a
    return v

# Bivariate polynomials in t,x (or t,d for radial models).
def clean(p):
    return {k: Q(v) for k, v in p.items() if v}

def diff2(p, coordinate, order=1):
    for _ in range(order):
        r = {}
        for key, value in p.items():
            k = list(key)
            if k[coordinate]:
                n = k[coordinate]
                k[coordinate] -= 1
                r[tuple(k)] = r.get(tuple(k), 0) + n*value
        p = clean(r)
    return p

def sum2(p, q, factor=Q(1)):
    r = dict(p)
    for k, v in q.items():
        r[k] = r.get(k, 0) + factor*v
    return clean(r)

def evaluate2(p, t, x):
    return sum((v*t**i*x**j for (i,j),v in p.items()), Q(0))

def forced_monomial(m, n, c):
    # Exact symmetric inner integral, then beta integral in s.
    return {(m+k+2,n-k):Q(comb(n,k)*factorial(m)*factorial(k), factorial(m+k+2))*c**k
            for k in range(0,n+1,2)}

def pulse_and_reflection_checks():
    check(derivative(PRIMITIVE) == PULSE, 'primitive_identity')
    check(integrate(PULSE,-1,1) == 2*SATURATION, 'pulse_mass')
    for degree in [3,4]:
        polynomial = power((1,0,-1),degree)
        for order in range(degree):
            for endpoint in [-1,1]:
                check(evaluate(derivative(polynomial,order),endpoint)==0,'endpoint_smoothness')
    times = list(map(Q,[2,3]))+[Q(7,2),Q(9,2),Q(11,2),Q(6),Q(7),Q(8)]
    expected = [Q(-4,35),Q(-8,35),Q(-4,35),Q(0),Q(-4,35),Q(-8,35),Q(-4,35),Q(0)]
    for t, answer in zip(times,expected):
        check(detector(t)==answer,'double_pulse_table')
    breaks = list(map(Q,[0,1,3,4,5,6,8,10]))
    pieces = []
    for left,right in zip(breaks,breaks[1:]):
        polynomial = detector_polynomial((left+right)/2)
        pieces.append(polynomial)
        for k in range(41):
            t = left+(right-left)*Q(k,40)
            check(evaluate(polynomial,t)==detector(t),'piecewise_polynomial_grid')
    # Exact derivative identities give signs on the four active windows;
    # each is a scalar multiple of a nonnegative pulse on its full interval.
    derivative_specs = [(1,Q(-1,4),(Q(2),Q(-1))),
                        (2,Q(1,2),(Q(7),Q(-2))),
                        (4,Q(-1,2),(Q(11),Q(-2))),
                        (5,Q(1,4),(Q(7),Q(-1)))]
    for index,factor,affine in derivative_specs:
        check(derivative(pieces[index])==scale(compose(PULSE,affine),factor),'window_derivative_identity')
    for index in [0,3,6]:
        check(pieces[index]==(Q(0),),'quiet_window_identity')
    for k in range(1,len(pieces)):
        for order in range(3):
            check(evaluate(derivative(pieces[k-1],order),breaks[k])==
                  evaluate(derivative(pieces[k],order),breaks[k]),'window_C2_join')
    for c in [Q(1,3),Q(1),Q(2),Q(5)]:
        for k in range(101):
            t = Q(k,10)
            check(reflected_pulse(t,0,2,1,1,c)==0,'wall_boundary')
            if c*t>=8:
                check(reflected_pulse(t,5,2,1,1,c)==0,'complete_reflection_cancellation')
    # Omitting the images leaves a nonzero late plateau for one positive pulse.
    wrong = (primitive(4+2*4-2)-primitive(4-2*4-2))/4
    check(wrong==Q(8,35) and reflected_pulse(4,4,2,1,1)==0,'missing_image_rejected')
    check(reflected_pulse(1,4,2,1,1,2)!=reflected_pulse(1,4,2,1,1,1),'wrong_wave_speed_rejected')
    return {'times':list(map(str,times)),'values':list(map(str,expected)),
            'windows':[list(map(str,[a,b])) for a,b in zip(breaks,breaks[1:])],
            'piece_coefficients':[[str(x) for x in p] for p in pieces]}

def forced_checks():
    for c in [Q(1,3),Q(1),Q(2),Q(5)]:
        for m in range(7):
            for n in range(9):
                p = forced_monomial(m,n,c)
                residual = sum2(diff2(p,0,2),diff2(p,1,2),-c*c)
                check(residual=={(m,n):Q(1)},'forced_polynomial_PDE_identity')
                check(all(i>=2 for i,j in p),'forced_zero_initial_data')
    p = forced_monomial(1,2,Q(2))
    check(p=={(3,2):Q(1,6),(5,0):Q(1,15)},'source_model_coefficients')
    check(evaluate2(p,Q(1),Q(2))==Q(11,15),'source_model_value')
    wrong = {(3,2):Q(1,6)}
    check(sum2(diff2(wrong,0,2),diff2(wrong,1,2),-4)!={(1,2):Q(1)},'temporal_only_source_rejected')
    return {'solution':{'t^3*x^2':'1/6','t^5':'1/15'},'u(1,2)':'11/15'}

def energy_checks():
    phi = (Q(0),Q(1),Q(-1))
    a = integrate(mul(phi,phi),0,1)
    b = integrate(power(derivative(phi),2),0,1)
    check(a==Q(1,30) and b==Q(1,3),'energy_integrals')
    eta = Q(1,100)
    root30 = sqrt_interval(30)
    budget = (as_interval(2)/root30 + Q(8,3))*eta
    actual = sqrt_interval(Q(22,15))*eta
    displacement = as_interval(eta)/root30
    displacement_budget = (as_interval(1)/root30+Q(2,3))*eta
    check(actual.hi < budget.lo < budget.hi < Q(31,1000),'energy_budget_outward')
    check(displacement.hi < displacement_budget.lo,'displacement_budget_outward')
    check(Q(23,750)<Q(31,1000),'coarse_rational_energy_bound')
    # Polynomial integration checks actual norms and a uniform-in-t majorant.
    for j in range(101):
        t = Q(j,100)
        q_squared = eta**2*(4*t*t*a+4*t**4*b)
        check(q_squared==eta**2*(Q(2,15)*t*t+Q(4,3)*t**4),'exact_energy_norm')
        residual = add(scale(phi,2*eta),(8*eta*t*t,))
        residual_squared = integrate(power(residual,2),0,1)
        check(residual_squared==4*eta**2*(Q(1,30)+Q(4,3)*t*t+16*t**4),'residual_norm_identity')
        majorant = (as_interval(2)/root30+8*t*t)*eta
        check(residual_squared <= majorant.hi**2,'residual_triangle_majorant_grid')
        # The universal triangle estimate is proved in the page; this additionally
        # proves this particular polynomial inequality with sqrt(30) < 6.
        lower_majorant_squared = 4*eta**2*a+64*eta**2*t**4+Q(32,6)*eta**2*t*t
        check(residual_squared <= lower_majorant_squared,'residual_bound_rational_comparison')
    def zero_data_and_wall_contract(p):
        initial = {j:v for (i,j),v in p.items() if i==0 and v}
        velocity = {j:v for (i,j),v in diff2(p,0).items() if i==0 and v}
        at_zero = {i:v for (i,j),v in p.items() if j==0 and v}
        at_one = {}
        for (i,j),v in p.items():
            at_one[i] = at_one.get(i,Q(0))+v
        return not initial and not velocity and not at_zero and not any(at_one.values())
    check(zero_data_and_wall_contract({(2,1):eta,(2,2):-eta}), 'candidate_data_and_wall_contract')
    check(not zero_data_and_wall_contract({(0,0):Q(7)}),'constant_offset_initial_data_rejected')
    check(not zero_data_and_wall_contract({(1,1):eta,(1,2):-eta}),'initial_velocity_mismatch_rejected')
    return {'Q(1)^2':str(eta**2*Q(22,15)),
            'Q(1)_interval':interval_strings(actual),
            'energy_budget_interval':interval_strings(budget),
            'displacement_interval':interval_strings(displacement),
            'displacement_budget_interval':interval_strings(displacement_budget)}

def radial_polynomial(j,c):
    n = 2*j+2
    return {(k,n-k-1):Q(comb(n,k),n)*c**(k-1) for k in range(1,n+1,2)}

def spherical_pulse(t,d,c=Q(1)):
    t,d,c=map(Q,(t,d,c))
    if t<0 or d<0 or c<=0:
        raise ValueError('spherical input contract')
    if d==0:
        return t*pulse(c*t)
    lower,upper=abs(d-c*t),min(d+c*t,Q(1))
    if lower>=upper:
        return Q(0)
    return ((1-lower**2)**4-(1-upper**2)**4)/(16*c*d)

def spherical_checks():
    for c in [Q(1,3),Q(1),Q(2),Q(5)]:
        for j in range(8):
            p=radial_polynomial(j,c)
            radial_gradient=diff2(p,1)
            extra={(i,k-1):2*v for (i,k),v in radial_gradient.items()}
            laplacian=sum2(diff2(p,1,2),extra)
            check(sum2(diff2(p,0,2),laplacian,-c*c)=={},'spherical_polynomial_PDE_identity')
            initial={key:value for key,value in diff2(p,0).items() if key[0]==0}
            check(initial=={(0,2*j):Q(1)},'spherical_polynomial_initial_velocity')
    check(spherical_pulse(Q(3,2),2)==Q(81,8192),'off_center_first_value')
    check(spherical_pulse(2,2)==Q(1,32),'off_center_peak_value')
    for k in range(201):
        t=Q(k,40)
        expected=(1-(2-t)**2)**4/32 if 1<t<3 else Q(0)
        check(spherical_pulse(t,2)==expected,'off_center_window_grid')
    for c in [Q(1,3),Q(1),Q(2),Q(5)]:
        for d in [Q(1,4),Q(1),Q(2),Q(7,2)]:
            for k in range(41):
                t=Q(k,10)
                value=spherical_pulse(t,d,c)
                check(value>=0,'positive_velocity_contribution')
                if c*t>d+1 or d>c*t+1:
                    check(value==0,'sphere_misses_support')
    f=power((1,0,-1),4)
    displacement=add(f,mul((0,1),derivative(f)))
    check(evaluate(displacement,Q(1,2))==Q(-135,256),'positive_displacement_negative_solution')
    # f=|x|^2-1 has M_r f=r^2-1 at x=0; d_t[t(t^2-1)] at t=1 is 2.
    check(evaluate(derivative((0,-1,0,1)),1)==2,'normal_derivative_counterexample')
    return {'off_center_3/2':'81/8192','off_center_2':'1/32',
            'positive_displacement_at_1/2':'-135/256','sphere_value_only_rejected':True}

def two_dimensional_center(t):
    t=Q(t)
    if t<0:
        raise ValueError('negative time')
    a=1-t*t
    antiderivative=(Q(0),a**3,Q(0),a*a,Q(0),Q(3,5)*a,Q(0),Q(1,7))
    lower=sqrt_interval(max(Q(0),t*t-1))
    return as_interval(evaluate(antiderivative,t))-interval_evaluate(antiderivative,lower)

def tail_checks():
    root3=sqrt_interval(3)
    explicit=(root3*432-746)/35
    general=two_dimensional_center(2)
    check(general.lo<=explicit.hi and explicit.lo<=general.hi,'dimension_two_radical_crosscheck')
    check(Q(1,16)<explicit.lo<explicit.hi<Q(3,40),'dimension_two_outward_interval')
    check(primitive(2)==SATURATION and spherical_pulse(2,0)==0,'dimension_one_three_at_two')
    for k in range(1,81):
        t=Q(k,20)
        value=two_dimensional_center(t)
        check(value.lo>0,'two_dimensional_positive_grid')
        if t<=1:
            exact=evaluate((0,1,0,-2,0,Q(8,5),0,Q(-16,35)),t)
            check(value.lo==exact==value.hi,'early_two_dimensional_polynomial')
        else:
            upper=as_interval(Q(1,8))/sqrt_interval(t*t-1)
            check(Q(1,8)/t<value.lo and value.hi<upper.lo,'two_dimensional_tail_outward_grid')
    # Integral mass gives the general 1/(8t) bound in the prose.
    check(integrate(mul((0,1),PULSE),0,1)==Q(1,8),'tail_mass_constant')
    return {'u1(2)':'16/35','u3(2)':'0','u2(2)_interval':interval_strings(explicit),
            'strict_outer_interval':['1/16','3/40'],'limit_t_times_u2':'1/8'}

def interval_strings(value):
    return [str(value.lo),str(value.hi)]

def main():
    CHECKS.clear()
    parser=ArgumentParser(description=__doc__)
    parser.add_argument('--output',type=Path)
    args=parser.parse_args()
    results={'reflection':pulse_and_reflection_checks(),'source':forced_checks(),
             'energy':energy_checks(),'spherical':spherical_checks(),'tails':tail_checks()}
    results={'status':'PASS','checks':sum(CHECKS.values()),'groups':dict(CHECKS),
             'results':results,'scope':'Finite exact polynomial/rational certificates; grid diagnostics are explicitly labelled. General theorems are proved in the companion pages.'}
    output=json.dumps(results,ensure_ascii=False,indent=2)+'\n'
    if args.output:
        args.output.write_text(output)
    print(output,end='')

if __name__=='__main__':
    main()
