"""Exact rational replays of the paper's two conventional certificates.

No Lean invocation. A local certificate is never promoted to a global result.
"""
from dataclasses import dataclass
from fractions import Fraction as F
import json
from pathlib import Path
import sympy as sp
from example import Reactor, Resident, Consumers, Reservoir


@dataclass(frozen=True)
class Interval:
    lo:F
    hi:F

    def __post_init__(self):
        object.__setattr__(self,'lo',F(self.lo));object.__setattr__(self,'hi',F(self.hi))
        if self.lo>self.hi:raise ValueError('Reversed interval.')
    @staticmethod
    def point(v):return v if isinstance(v,Interval) else Interval(v,v)
    def __add__(self,v):
        v=self.point(v);return Interval(self.lo+v.lo,self.hi+v.hi)
    __radd__=__add__
    def __neg__(self):return Interval(-self.hi,-self.lo)
    def __sub__(self,v):return self+-self.point(v)
    def __mul__(self,v):
        v=self.point(v);p=[self.lo*v.lo,self.lo*v.hi,self.hi*v.lo,self.hi*v.hi];return Interval(min(p),max(p))
    __rmul__=__mul__
    def __pow__(self,n):
        if not isinstance(n,int) or n<0:raise ValueError('Nonnegative integer powers only.')
        value=self.point(1)
        for _ in range(n):value=value*self
        return value
    def magnitude(self):return max(abs(self.lo),abs(self.hi))


def polynomial_interval(poly,values):
    total=Interval.point(0)
    for powers,c in poly.terms():
        term=Interval.point(F(c))
        for value,power in zip(values,powers):
            if power:term=term*value**power
        total=total+term
    return total


def norm_inf(matrix):return max(sum(abs(F(matrix[i,j])) for j in range(matrix.cols)) for i in range(matrix.rows))


def local_certificate(radius=None,tolerance=None,inputs=None):
    data=inputs if inputs is not None else json.loads(Path(__file__).with_name('local_certificate_inputs.json').read_text())
    center=list(map(F,data['center']));P=sp.Matrix([[sp.Rational(v) for v in row] for row in data['P']])
    radius=F(data['radius']) if radius is None else F(radius);delta=F(data['rate_tolerance']) if tolerance is None else F(tolerance)
    if radius<=0 or delta<0:raise ValueError('Positive radius and nonnegative tolerance required.')
    model=Reactor(Resident(),Consumers.equal(2),Reservoir.reference(F(1,20)))
    rates=[F(r.rate) for r in model.reactions]
    if delta>=min(rates):raise ValueError('Rate cube must remain positive.')
    states=sp.symbols('A B z H R X1 X2');rs=sp.symbols('r0:21')
    # Derive the exact polynomial field from the literal directed reaction list.
    field=sp.zeros(7,1)
    for r,k in zip(model.reactions,rs):
        monomial=k
        for x,p in zip(states,r.inputs):monomial*=x**p
        for i,(a,b) in enumerate(zip(r.inputs,r.outputs)):field[i]+=(b-a)*monomial
    J=field.jacobian(states);sub=dict(zip(states+rs,center+rates));J0=J.subs(sub)
    minors=[F(P[:k,:k].det()) for k in range(1,8)]
    if P!=P.T or min(minors)<=0:raise ValueError('P is not symmetric positive definite.')
    pmin=1/norm_inf(P.inv());pmax=norm_inf(P);Q=-(P*J0+J0.T*P)
    qmin=min(F(Q[i,i])-sum(abs(F(Q[i,j])) for j in range(7) if j!=i) for i in range(7))
    vals=[Interval(c-radius,c+radius) for c in center]+[Interval(k-delta,k+delta) for k in rates]
    dj=max(sum((polynomial_interval(sp.Poly(J[i,j],states+rs),vals)-F(J0[i,j])).magnitude() for j in range(7)) for i in range(7))
    fixed=[Interval.point(c) for c in center]+vals[7:]
    fc=max(polynomial_interval(sp.Poly(f,states+rs),fixed).magnitude() for f in field)
    ell=F(data['sqrt_ratio_lower']);invnorm=norm_inf(J0.inv());qeff=qmin-6*pmax*dj;k=invnorm*dj
    margin=qeff*ell*radius-6*pmax*fc;maps=invnorm*fc+k*radius<radius
    root_radius=invnorm*fc/(1-k) if k<1 else None
    accepted=qmin>0 and qeff>0 and margin>0 and k<1 and maps and ell>0 and ell**2<=pmin/pmax and root_radius is not None and 3*root_radius<ell*radius and min(center)>radius
    values={'radius':radius,'rate_tolerance':delta,'p_min':pmin,'p_max':pmax,'q_min':qmin,'jacobian_inf_error':dj,'field_inf_error':fc,'q_eff':qeff,'contraction_bound':k,'boundary_margin':margin,'equilibrium_inf_radius':root_radius,'consumer_floor':min(center[5:])-radius,'ellipsoid_level':pmin*radius**2}
    return {'accepted':bool(accepted),'scope':'Exact rational interval checks for the conventional invariant-ellipsoid and contraction argument. Applies only inside E, not to every positive state. Lean not rerun.',
        'center':list(map(str,center)),'P':[[str(v) for v in row] for row in data['P']],
        'sylvester_minors':list(map(str,minors)),'newton_maps_box_into_itself':bool(maps),**{key:str(value) if value is not None else None for key,value in values.items()}}


def donor_certificate():
    z=sp.Symbol('z');e=sp.Rational(1,100000);h=sp.Rational(20001,10000)
    B=60/(z+2);H=(16*z+2*z*z)/h;A=B*z+16*z+4*z*z-3*H
    polynomial=sp.Poly(sp.cancel(A+B+e*A*A-e*B-33).as_numer_denom()[0],z)
    lo=sp.Rational(9957940122,10**10);hi=sp.Rational(9957940124,10**10)
    shifted=sp.Poly(polynomial.diff().as_expr().subs(z,lo+z),z)
    lower=shifted.eval(0)+sum(min(c,0)*(hi-lo)**power[0] for power,c in shifted.terms() if power[0]>0)
    # Explicit interval expressions also certify positivity of eliminated A,B,H.
    # A = B*z + (6668*z^2-53328*z)/6667; use separate conservative bounds.
    A_lower=60*lo/(hi+2)+(6668*lo**2-53328*hi)/6667
    valid=polynomial.eval(lo)<0<polynomial.eval(hi) and lower>0 and A_lower>0 and lo>sp.Rational(1,2)
    if not valid:raise ArithmeticError('Donor enclosure failed.')
    return {'accepted':True,'scope':'Rational IVT enclosure and nonsingularity through scalar elimination; local small-supply branch, not a global continuation certificate.',
        'polynomial_coefficients':list(map(str,polynomial.all_coeffs())),'z_interval':[str(lo),str(hi)],'p_left':str(polynomial.eval(lo)),'p_right':str(polynomial.eval(hi)),
        'derivative_lower':str(lower),'A_lower':str(A_lower),'small_supply_slope_interval':[str(2-1/lo),str(2-1/hi)]}
