#!/usr/bin/env python3
"""Exact fixed-delay history, continuous Euler and polynomial-defect certificates.

Python standard library only. All certificate checks are explicit under -O too.
The model is scalar x'=-a*x+b*x(t-tau)+p(t), with rational coefficients,
quadratic initial/candidate histories and quadratic p. No general DDE solver,
interval library or sampled estimate of a continuum supremum is claimed.
"""
from __future__ import annotations
import argparse
from collections import Counter
from copy import deepcopy
from fractions import Fraction as F
from math import factorial, comb
from pathlib import Path
import json

COUNTS=Counter()

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


def rat(s):
    if isinstance(s,bool) or not isinstance(s,(str,int,F)):
        raise ValueError('rational input required')
    return F(s)


def poly(p):
    if not isinstance(p,(list,tuple)) or len(p)!=3:
        raise ValueError('exactly three quadratic coefficients required')
    return [rat(x) for x in p]


def add(*ps):
    return [sum((p[k] for p in ps),F(0)) for k in range(3)]


def scale(p,c):
    return [c*x for x in p]


def value(p,t):
    return (p[2]*t+p[1])*t+p[0]


def shift(p,d):
    # p(t+d), by an exact binomial expansion.
    return [value(p,d),p[1]+2*p[2]*d,p[2]]


def maximum_abs(p,lo,hi):
    if not lo<hi:
        raise ValueError('positive interval required')
    points=[lo,hi]
    if p[2]:
        critical=-p[1]/(2*p[2])
        if lo<critical<hi:
            points.append(critical)
    values=[abs(value(p,t)) for t in points]
    bound=max(values)
    return bound,points,points[values.index(bound)]


def strings(xs):
    return [str(x) for x in xs]


def exp_bounds(x,terms=40):
    """Rational Taylor lower/upper enclosure of exp(x), x>=0."""
    x=rat(x)
    if x<0:
        raise ValueError('nonnegative exponential argument required')
    terms=max(terms,int(x)+2)
    term=F(1);lower=term
    for k in range(1,terms+1):
        term*=x/k;lower+=term
    following=term*x/(terms+1)
    ratio=x/(terms+2)
    if ratio>=1:
        raise ValueError('invalid exponential tail ratio')
    return lower,lower+following/(1-ratio)


def rate_certificate(a,b,tau,lam):
    a,b,tau,lam=map(rat,(a,b,tau,lam))
    if not (a>abs(b) and tau>0 and lam>0):
        raise ValueError('dissipativity, positive delay and rate required')
    lo,hi=exp_bounds(lam*tau)
    check(lam+abs(b)*hi<=a,'certified Halanay rate')
    return {'rate':str(lam),'exponential_lower':str(lo),
            'exponential_upper':str(hi),'positive_margin':str(a-lam-abs(b)*hi)}


def model(a=3,b=1,tau=F(3,5),h=F(1,4),T=2,
          true_history=(0,0,1),candidate_history=None,p=None):
    a,b,tau,h,T=map(rat,(a,b,tau,h,T))
    q=poly(true_history)
    if p is None:
        # Manufacture exact x=q(t), while verification below need not use it.
        p=add([q[1],2*q[2],F(0)],scale(q,a),scale(shift(q,-tau),-b))
    return {'a':str(a),'b':str(b),'tau':str(tau),'h':str(h),'T':str(T),
            'true_history':strings(q),'candidate_history':strings(poly(q if candidate_history is None else candidate_history)),
            'p':strings(poly(p))}


def parse_model(data):
    required={'a','b','tau','h','T','true_history','candidate_history','p'}
    if set(data)!=required:
        raise ValueError('exact model fields required')
    a,b,tau,h,T=(rat(data[k]) for k in ['a','b','tau','h','T'])
    if not (a>abs(b) and 0<h<=tau and T>0 and (T/h).denominator==1):
        raise ValueError('invalid dissipative uniform-grid model')
    N=int(T/h)
    return a,b,tau,h,T,N,poly(data['true_history']),poly(data['candidate_history']),poly(data['p'])


def query(s,h,values,history,available=None):
    available=len(values) if available is None else available
    if s<=0:
        return {'argument':str(s),'kind':'initial_history','left_index':None,
                'weight':None,'value':str(value(history,s))}
    j=int(s/h)
    if j+1>=available:
        raise ValueError('delay query uses unavailable history')
    alpha=(s-j*h)/h
    if not 0<=alpha<=1:
        raise ValueError('invalid interpolation interval')
    z=(1-alpha)*values[j]+alpha*values[j+1]
    return {'argument':str(s),'kind':'interpolation','left_index':j,
            'weight':str(alpha),'value':str(z)}


def line(j,h,ys):
    slope=(ys[j+1]-ys[j])/h
    return [ys[j]-j*h*slope,slope,F(0)]


def merge_breakpoints(nodes,shifted):
    # Two already sorted lists; no set or comparison sort and no repeated slices.
    i=j=0;result=[];iterations=0
    while i<len(nodes) or j<len(shifted):
        iterations+=1
        if j==len(shifted) or (i<len(nodes) and nodes[i]<shifted[j]):
            point=nodes[i];i+=1
        elif i==len(nodes) or shifted[j]<nodes[i]:
            point=shifted[j];j+=1
        else:
            point=nodes[i];i+=1;j+=1
        if not result or point!=result[-1]:result.append(point)
    check(iterations<=len(nodes)+len(shifted),'linear original and delayed breakpoint merge')
    return result


def residual_records(data,ys):
    a,b,tau,h,T,N,truth,history,p=parse_model(data)
    cuts=merge_breakpoints([j*h for j in range(N+1)],
        [j*h+tau for j in range(N+1) if 0<j*h+tau<T])
    records=[]
    for lo,hi in zip(cuts,cuts[1:]):
        mid=(lo+hi)/2;j=int(mid/h);current=line(j,h,ys)
        delayed=shift(history,-tau) if mid<tau else shift(line(int((mid-tau)/h),h,ys),-tau)
        r=add([current[1],F(0),F(0)],scale(current,a),scale(delayed,-b),scale(p,-1))
        bound,points,argmax=maximum_abs(r,lo,hi)
        records.append({'lo':str(lo),'hi':str(hi),'coefficients':strings(r),
                        'test_points':strings(points),'argmax':str(argmax),'bound':str(bound)})
    return cuts,records


def build_certificate(data):
    a,b,tau,h,T,N,truth,history,p=parse_model(data)
    ys=[value(history,0)];slopes=[];queries=[]
    for n in range(N):
        t=n*h;record=query(t-tau,h,ys,history);z=rat(record['value'])
        k=-a*ys[n]+b*z+value(p,t)
        queries.append(record);slopes.append(k);ys.append(ys[n]+h*k)
    cuts,records=residual_records(data,ys)
    eta=max(rat(r['bound']) for r in records)
    M,_,_=maximum_abs(add(truth,scale(history,-1)),-tau,F(0))
    floor=eta/(a-abs(b))
    return {'nodes':strings([j*h for j in range(N+1)]),'values':strings(ys),
            'slopes':strings(slopes),'queries':queries,'cuts':strings(cuts),
            'residuals':records,'initial_error':str(M),'residual_bound':str(eta),
            'forcing_floor':str(floor),'uniform_error_bound':str(max(M,floor))}


def verify_certificate(data,cert):
    a,b,tau,h,T,N,truth,history,p=parse_model(data)
    required={'nodes','values','slopes','queries','cuts','residuals',
              'initial_error','residual_bound','forcing_floor','uniform_error_bound'}
    check(set(cert)==required,'certificate fields')
    check(cert['nodes']==strings([j*h for j in range(N+1)]),'complete uniform time grid')
    check(len(cert['values'])==N+1 and len(cert['slopes'])==N and len(cert['queries'])==N,
          'state slope and query lengths')
    ys=[rat(x) for x in cert['values']];slopes=[rat(x) for x in cert['slopes']]
    check(ys[0]==value(history,0),'initial candidate-history join')
    for n in range(N):
        t=n*h;expected=query(t-tau,h,ys,history,available=n+1)
        check(cert['queries'][n]==expected,'causal exact historical query')
        k=-a*ys[n]+b*rat(expected['value'])+value(p,t)
        check(slopes[n]==k,'Euler slope from the declared model')
        check(ys[n+1]==ys[n]+h*k,'continuous Euler update')
    cuts,records=residual_records(data,ys)
    check(cert['cuts']==strings(cuts),'original and shifted breakpoints complete')
    check(len(cert['residuals'])==len(records),'residual interval coverage')
    for actual,expected in zip(cert['residuals'],records):
        check(actual==expected,'quadratic extrema and exact residual coefficients')
    eta=max(rat(r['bound']) for r in records)
    M,_,_=maximum_abs(add(truth,scale(history,-1)),-tau,F(0))
    check(rat(cert['initial_error'])==M,'complete initial-history error')
    check(rat(cert['residual_bound'])==eta,'global residual supremum')
    floor=eta/(a-abs(b))
    check(rat(cert['forcing_floor'])==floor,'Halanay persistent-defect floor')
    check(rat(cert['uniform_error_bound'])==max(M,floor),'uniform continuous error certificate')
    return {'status':'CERTIFIED','interval':['0',str(T)],'error_bound':str(max(M,floor)),
            'history_match':M==0,'residual_segments':len(records)}


def actual_polynomial_error(data,cert):
    a,b,tau,h,T,N,truth,history,p=parse_model(data)
    exact_rhs=add([truth[1],2*truth[2],F(0)],scale(truth,a),scale(shift(truth,-tau),-b))
    check(exact_rhs==p,'manufactured polynomial solves full delay equation')
    ys=list(map(rat,cert['values']));best=F(0);where=F(0)
    for j in range(N):
        err=add(line(j,h,ys),scale(truth,-1));v,_,t=maximum_abs(err,j*h,(j+1)*h)
        if v>best:best,where=v,t
    check(best<=rat(cert['uniform_error_bound']),'true all-time error lies within certificate')
    return {'exact_maximum_error':str(best),'attained_at':str(where)}


def delayed_sum(t,tau):
    if t<0:return F(1)
    return F(1)+sum(((t-k*tau)**(k+1)/factorial(k+1)
                     for k in range(int(t/tau)+1)),F(0))


def polynomial_steps(history,tau,windows):
    # Independent explicit integration route, for x'=x(t-tau).
    def evaluate(p,t):
        return sum((c*t**j for j,c in enumerate(p)),F(0))
    previous=list(map(rat,history));endpoint=evaluate(previous,F(0));pieces=[]
    for k in range(windows):
        shifted=[sum((previous[j]*comb(j,i)*(-tau)**(j-i)
                      for j in range(i,len(previous))),F(0))
                 for i in range(len(previous))]
        current=[F(0)]+[c/(j+1) for j,c in enumerate(shifted)]
        left=k*tau;current[0]=endpoint-evaluate(current,left)
        check(evaluate(current,left)==endpoint,'method-of-steps continuous polynomial join')
        for j in range(11):
            t=left+tau*j/10
            derivative=sum((i*current[i]*t**(i-1) for i in range(1,len(current))),F(0))
            check(derivative==evaluate(previous,t-tau),'method-of-steps differentiated shifted history')
        endpoint=evaluate(current,(k+1)*tau);pieces.append(current);previous=current
    return pieces,endpoint


def step_checks():
    for tau in [F(1,3),F(2,3),F(3,2)]:
        for j in range(1,71):
            t=tau*j/10
            derivative=sum(((t-k*tau)**k/factorial(k)
                            for k in range(int(t/tau)+1)),F(0))
            check(derivative==delayed_sum(t-tau,tau),'delayed polynomial satisfies differential equation')
        for m in range(1,8):
            t=m*tau
            # The new degree m+1 term vanishes at its start, as do its first m derivatives.
            for order in range(m+1):
                check((t-m*tau)**(m+1-order)/factorial(m+1-order)==0,
                      'delayed polynomial derivative joins')
    tau=F(2,3)
    check(delayed_sum(F(2),tau)==F(319,81),'three-window exact endpoint')
    check(F(1)+tau!=F(1)+tau/2,'same current value does not determine future')
    constant,ce=polynomial_steps([F(1)],tau,3)
    affine,ae=polynomial_steps([F(1),1/tau],tau,3)
    check(ce==F(319,81) and ae==F(259,81),'two histories independently integrated to terminal')
    return {'tau':'2/3','constant_history_endpoint':str(ce),'affine_history_endpoint':str(ae),
            'same_initial_value_first_endpoints':['5/3','4/3'],
            'constant_history_pieces':[strings(p) for p in constant],
            'affine_history_pieces':[strings(p) for p in affine]}


def reject_call(fn,label):
    try:fn()
    except ValueError:
        COUNTS['rejected '+label]+=1
        return label
    raise ValueError('invalid certificate accepted: '+label)


def main():
    COUNTS.clear()
    steps=step_checks()
    rates=[rate_certificate(3,1,F(3,5),1),rate_certificate(3,1,1,F(1,2))]
    exp2=exp_bounds(F(2));exp14=exp_bounds(F(7,5))
    check(exp2[0]>7 and exp14[0]>4,'point and history exponential lower bounds')
    check(F(1,500)+F(7,250)/exp2[0]<F(3,500),'current endpoint rational error budget')
    check(F(1,500)+F(7,250)/exp14[0]<F(9,1000),'entire delayed-window rational error budget')
    reject_rate=reject_call(lambda:rate_certificate(3,1,1,1),'old rate after increased delay')
    principal=[]
    for h in [F(1,4),F(1,8)]:
        data=model(h=h);cert=build_certificate(data);result=verify_certificate(data,cert)
        principal.append({'model':data,'certificate':cert,'verification':result,
                          'independent_exact_solution_comparison':actual_polynomial_error(data,cert)})
    check(principal[0]['certificate']['residual_bound']=='11/16','coarse residual value')
    check(principal[1]['certificate']['residual_bound']=='19/64','refined residual value')
    check(rat(principal[0]['certificate']['uniform_error_bound'])>F(3,20) and
          F(3,20)-rat(principal[1]['certificate']['uniform_error_bound'])==F(1,640),
          'refinement crosses the requested certified tolerance')
    check(principal[0]['certificate']['queries'][4]['value']=='33/320','off-grid delay example')
    check(principal[0]['certificate']['values'][-1]=='127500911/32768000','coarse endpoint value')
    # A one-step example where every mesh-point defect is zero but its interior maximum is one.
    spike=model(tau=1,h=1,T=1,true_history=(0,0,0),p=(0,4,-4))
    sc=build_certificate(spike);sv=verify_certificate(spike,sc)
    check(sc['residuals'][0]['test_points']==['0','1','1/2'],'interior stationary point included')
    check(sc['residual_bound']=='1' and sc['uniform_error_bound']=='1/2','zero endpoints are not zero defect')
    rejected=[reject_rate]
    source=principal[0];data=source['model'];cert=source['certificate']
    mutations=[]
    bad=deepcopy(cert);bad['cuts'].remove('3/5');mutations.append(('missing delayed breakpoint',bad))
    bad=deepcopy(cert);bad['queries'][4]['value']='0';mutations.append(('nearest-left history substitution',bad))
    bad=deepcopy(cert);bad['values'][3]='0';mutations.append(('altered Euler state',bad))
    bad=deepcopy(cert);bad['residuals'][2]['coefficients'][2]='-2';mutations.append(('wrong history-side polynomial',bad))
    bad=deepcopy(cert);bad['residual_bound']='1/4';mutations.append(('understated residual supremum',bad))
    bad=deepcopy(cert);bad['nodes'].pop();mutations.append(('missing terminal time',bad))
    bad=deepcopy(cert);bad['initial_error']='1/10';mutations.append(('wrong initial-history budget',bad))
    for label,bad in mutations:
        rejected.append(reject_call(lambda bad=bad:verify_certificate(data,bad),label))
    bad=deepcopy(sc);bad['residuals'][0]['test_points']=['0','1'];bad['residuals'][0]['bound']='0'
    rejected.append(reject_call(lambda:verify_certificate(spike,bad),'omitted interior residual extremum'))
    rejected.append(reject_call(lambda:build_certificate(model(a=1,b=1)),'missing strict dissipativity'))
    rejected.append(reject_call(lambda:build_certificate(model(h=1)),'step beyond selected causal-window contract'))
    rejected.append(reject_call(lambda:poly([0,0,1,4]),'extra polynomial coordinate'))
    rejected.append(reject_call(lambda:model(candidate_history=[]),'empty candidate history is not a default'))
    # Signed delayed feedback, nonmatching histories and varied meshes.
    model_count=0
    for tau in [F(1,3),F(3,5),F(1)]:
        for h in [F(1,8),F(1,12)]:
            for b in [F(-1),F(0),F(1)]:
                for q in [[F(0),F(0),F(1)],[F(1),F(-2),F(1,2)]]:
                    for delta in [F(0),F(1,20)]:
                        history=add(q,[delta,-delta,F(0)])
                        d=model(a=3,b=b,tau=tau,h=h,T=1,true_history=q,candidate_history=history)
                        c=build_certificate(d);verify_certificate(d,c);actual_polynomial_error(d,c)
                        model_count+=1
    # Check the extremum procedure on a rational grid as a diagnostic, not its proof.
    for A in range(-2,3):
        for B in range(-3,4):
            for C in range(-2,3):
                p=list(map(F,(A,B,C)));bound,_,_=maximum_abs(p,F(-2,3),F(5,4))
                for j in range(25):
                    t=F(-2,3)+F(j,24)*(F(5,4)+F(2,3))
                    check(abs(value(p,t))<=bound,'sampled rational values inside analytic quadratic bound')
    unstable_model=model(tau=1,h=1,T=11,true_history=(1,0,0),p=(0,0,0))
    unstable=build_certificate(unstable_model);verify_certificate(unstable_model,unstable)
    seq=list(map(rat,unstable['values']))
    check(unstable['queries'][0]['kind']=='initial_history' and
          unstable['queries'][0]['value']=='1' and unstable['slopes'][0]=='-2',
          'unstable Euler first step actually reads constant history')
    check(seq[:6]==list(map(F,[1,-1,3,-7,17,-41])),
          'continuous stable equation has unstable Euler example')
    return {'status':'PASS','checks':sum(COUNTS.values()),'categories':dict(COUNTS),
            'method_of_steps':steps,'rate_certificates':rates,
            'continuous_euler_models':principal,'interior_extremum_example':{
                'model':spike,'certificate':sc,'verification':sv},
            'additional_manufactured_models':model_count,'rejected_certificates':rejected,
            'unstable_euler_sequence':strings(seq),
            'scope':'Exact finite rational certificates and polynomial extrema; general existence, comparison and convergence are proved in the formal pages.'}


if __name__=='__main__':
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output',type=Path,required=True)
    args=parser.parse_args();result=main()
    args.output.parent.mkdir(parents=True,exist_ok=True)
    args.output.write_text(json.dumps(result,ensure_ascii=False,indent=2)+'\n')
    print(json.dumps({'status':result['status'],'checks':result['checks'],
                      'additional_models':result['additional_manufactured_models']}))
