"""Exact equilibrium elimination for small cubes; stability is a separate question."""
from fractions import Fraction as F
import sympy as s
import mpmath as mp
from reactor import EdgeKinetics,PhosphorylationNetwork


def rational(x):return s.Rational(str(x))


class EquilibriumGeometry:
    def __init__(self,network,tree_weights=None):
        if network.q>8:raise ValueError('exact cofactor example is limited to eight phosphoforms')
        if network.reverse_activity:raise ValueError('irreversible eliminant does not cover reverse catalysis')
        self.network=network;self.u=s.Symbol('u');u=self.u;q=network.q
        G=s.zeros(q)
        for (a,b),r in zip(network.edges,network.kinetics):
            up=rational(r.up)*u;down=rational(r.down)
            G[b,a]+=up;G[a,a]-=up;G[a,b]+=down;G[b,b]-=down
        self.G=G
        if tree_weights is None:
            self.tau=[s.expand((-1)**(q-1)*G.minor_submatrix(i,i).det(method='berkowitz')) for i in range(q)]
        else:
            self.tau=[s.sympify(t) for t in tree_weights]
        if len(self.tau)!=q or any(s.expand(v)!=0 for v in G*s.Matrix(self.tau)):
            raise ValueError('weights must lie in the effective-generator kernel')
        if any(s.Poly(t,u).is_zero or any(c<0 for c in s.Poly(t,u).all_coeffs()) for t in self.tau):raise ValueError('positive tree polynomials required')
        self.A=s.expand(sum(self.tau))
        self.B=s.expand(sum(rational(r.p)*self.tau[a] for (a,b),r in zip(network.edges,network.kinetics)))
        self.D=s.cancel(sum(rational(r.q)*self.tau[b] for (a,b),r in zip(network.edges,network.kinetics))/u)

    def eliminant(self,totals):
        E,Ft,S=map(rational,totals)
        if min(E,Ft,S)<=0:raise ValueError('all totals must be positive')
        u=self.u;r=E/Ft;L=self.B-r*self.D;M=self.B-u*self.D
        P=s.Poly(s.expand((S-E-Ft)*u*L*M-(r-u)*self.A*M+(u+1)*Ft*u*L**2),u)
        return P.clear_denoms()[1].primitive()[1]

    def census(self,totals):
        E,Ft,S=map(rational,totals);r=E/Ft;u=self.u;L=s.Poly(self.B-r*self.D,u);M=s.Poly(self.B-u*self.D,u)
        P=self.eliminant(totals);rows=[]
        for (lo,hi),multiplicity in P.intervals(eps=s.Rational(1,10**40)):
            if hi<=0:continue
            if lo<=r<=hi and P.eval(r)==0:
                # Reconstruction at u=r must use the remaining strictly increasing total equation.
                admissible=L.eval(r)==0
                rows.append(dict(interval=[F(r),F(r)],multiplicity=multiplicity,admissible=bool(admissible),exceptional_ratio=True))
                continue
            if L.is_zero or M.is_zero or L.count_roots(lo,hi) or M.count_roots(lo,hi):
                raise RuntimeError('sign enclosure unresolved; refine or handle a shared root explicitly')
            mid=(lo+hi)/2
            admissible=(r-mid)*L.eval(mid)>0 and L.eval(mid)*M.eval(mid)>0
            rows.append(dict(interval=[F(lo),F(hi)],multiplicity=multiplicity,admissible=bool(admissible),exceptional_ratio=False))
        return rows

    def state_on_curve(self,ratio,enzyme_totals,substrate_total=None,digits=70):
        """High-precision diagnostic reconstruction; not an interval enclosure."""
        with mp.workdps(digits):
            conv=lambda x:mp.mpf(str(s.N(rational(x),digits)))
            z=conv(ratio);ET,FT=map(conv,enzyme_totals);r=ET/FT
            polynomial=lambda t:mp.polyval([conv(c) for c in s.Poly(t,self.u).all_coeffs()],z)
            A,B,D=map(polynomial,[self.A,self.B,self.D]);L=B-r*D
            if z==r:
                if s.simplify((self.B-rational(enzyme_totals[0])/rational(enzyme_totals[1])*self.D).subs(self.u,rational(enzyme_totals[0])/rational(enzyme_totals[1])))!=0 or substrate_total is None:
                    raise ValueError('exceptional ratio needs L(r)=0 and a substrate total')
                ST=conv(substrate_total);a=A*r*D;b=A+r*FT*(B+D)-ST*r*D
                scale=2*ST/(b+mp.sqrt(b*b+4*a*ST))
            else:scale=(r-z)/(z*L)
            freeF=FT/(1+scale*z*D);freeE=z*freeF
            if min(scale,freeF,freeE)<=0:raise ValueError('ratio is outside the positive equilibrium branch')
            substrate=[scale*polynomial(t) for t in self.tau]
            c=[conv(k.p)*freeE*substrate[a] for (a,b),k in zip(self.network.edges,self.network.kinetics)]
            y=[conv(k.q)*freeF*substrate[b] for (a,b),k in zip(self.network.edges,self.network.kinetics)]
            return substrate+[freeE,freeF]+c+y


def capacity_bounds(sites):
    if type(sites) is not int or sites<2:raise ValueError('the displayed theorem is for integer sites >=2')
    lower=max(sites,(2**(sites-2)+sites)//2,4 if sites==3 else 0)
    return dict(sites=sites,constructive_lower=lower,universal_upper=2**sites,square_balanced_capacity=sites,exact_unrestricted_capacity_known=False)
