"""Composable literal mass-action models; concentration units are micromolar/minute."""
from dataclasses import dataclass
import numpy as np
import sympy as s
from scipy.integrate import solve_ivp
from certificates import Interval
R=s.Rational

@dataclass(frozen=True)
class Reaction:
    reactant:tuple
    product:tuple
    rate:object
    def __post_init__(self):
        object.__setattr__(self,'rate',R(self.rate))
        if len(self.reactant)!=len(self.product) or any(type(v)is not int or v<0 for v in self.reactant+self.product) or self.rate<=0:raise ValueError('Positive rate and equal-sized integer complexes required.')

class MassActionModel:
    def __init__(self,species,reactions):
        self.species=tuple(species);self.reactions=tuple(reactions);n=len(species)
        if not n or len(set(species))!=n or any(len(r.reactant)!=n for r in reactions):raise ValueError('Species and reaction dimensions do not match.')
        self.variables=s.symbols(' '.join(species),seq=True);self.N=s.Matrix([[r.product[i]-r.reactant[i] for r in reactions] for i in range(n)]);self.fluxes=s.Matrix([r.rate*s.prod(x**p for x,p in zip(self.variables,r.reactant)) for r in reactions]);self.field=tuple(map(s.expand,self.N*self.fluxes));self.J=s.Matrix(self.field).jacobian(self.variables)
        self._field=s.lambdify(self.variables,self.field,'numpy');self._jac=s.lambdify(self.variables,self.J,'numpy')
    def integrate(self,initial,duration,samples=501,method='Radau'):
        initial=np.asarray(initial,float)
        if initial.shape!=(len(self.species),) or min(initial)<=0 or duration<=0:raise ValueError('Positive initial concentrations and duration required.')
        sol=solve_ivp(lambda t,x:self._field(*x),(0,duration),initial,jac=lambda t,x:self._jac(*x),method=method,rtol=2e-10,atol=1e-12,t_eval=np.linspace(0,duration,samples))
        if not sol.success or np.min(sol.y)<=0:raise ArithmeticError('Trajectory solver or numerical positivity failed.')
        return sol.t,sol.y
    def regularity(self,state):
        if len(state)!=len(self.species) or any(s.sympify(x)<=0 for x in state):raise ValueError('Positive state required.')
        sub=dict(zip(self.variables,map(s.sympify,state)));res=[s.simplify(f.subs(sub)) for f in self.field]
        if any(v!=0 for v in res):return dict(status='not_a_steady_state')
        rank=self.J.subs(sub).rank();stoich=self.N.rank();return dict(status='regular' if rank==stoich else 'rank_deficient',jacobian_rank=rank,stoichiometric_rank=stoich)

@dataclass(frozen=True)
class LoadedReactor:
    rates:tuple=('1/20','1/20','1/20','1/20','3/20')
    dilution:object='1/100'
    feed:object='10'
    load:object='1/50'
    def parameters(self):
        k=tuple(map(R,self.rates));D,ci,ell=map(R,[self.dilution,self.feed,self.load])
        if len(k)!=5 or min(k)<=0 or D<=0 or ci<=0 or ell<0:raise ValueError('Positive rate constants, dilution and feed; nonnegative load required.')
        return k,D,ci,ell
    def model(self):
        k,D,ci,ell=self.parameters();r=[Reaction(a,b,c) for a,b,c in [((1,1,0),(0,2,0),k[0]),((0,1,0),(1,0,0),k[1]),((0,1,0),(0,0,1),k[2]),((1,0,1),(0,0,2),k[3]),((0,0,1),(1,0,0),k[4]),((0,0,0),(1,0,0),D*ci)]]
        for i in range(3):r.append(Reaction(tuple(int(j==i) for j in range(3)),(0,0,0),D))
        if ell>0:r.append(Reaction((1,0,0),(0,0,0),ell))
        return MassActionModel(('a','b','c'),r)
    def equilibrium(self,floor=1):
        k,D,ci,ell=self.parameters();alpha=(k[1]+k[2]+D)/k[0];delta=k[4]+D-k[3]*alpha;N=ci-alpha*(1+ell/D);floor=R(floor)
        if floor<=0:raise ValueError('Positive floor required.')
        result=dict(alpha=alpha,delta=delta,available_pool=N,critical_load=D*(ci/alpha-1))
        if delta<=0 or N<=0:return dict(status='no_positive_equilibrium',**result)
        r=delta/k[2];z=(alpha,r*N/(1+r),N/(1+r));sigma=k[0]*z[1];tau=k[3]*z[2];A1=D+ell+delta+sigma+tau;A2=delta*(D+ell)+sigma*(delta+k[2]+D)+D*tau;A3=D*sigma*(delta+k[2]);margin=A1*A2-A3
        return dict(status='positive_equilibrium',state=z,floor_satisfied=z[1]>=floor,floor_load_limit=D/alpha*(ci-alpha-(1+r)*floor/r),routh_coefficients=(A1,A2,A3),routh_margin=margin,locally_stable=bool(min(A1,A2,A3,margin)>0),**result)

@dataclass(frozen=True)
class EnvZOmpR:
    rates:tuple=('1',)*9
    def constants(self):
        k=tuple(map(R,self.rates))
        if len(k)!=9 or min(k)<=0:raise ValueError('Nine positive rates required.')
        alpha=k[2]*(k[7]+k[8])/(k[6]*k[8]);C=k[2]*(k[4]+k[5])/(k[3]*k[5]);K=k[2]/k[5]+k[2]/k[8];L=(k[1]+k[2])/k[0]+1+K
        return k,alpha,C,K,L
    def model(self):
        k,*_=self.constants();E=lambda *ids:tuple(int(i in ids) for i in range(7));pairs=[(E(0),E(1)),(E(1),E(0)),(E(1),E(2)),(E(2,3),E(5)),(E(5),E(2,3)),(E(5),E(0,4)),(E(1,4),E(6)),(E(6),E(1,4)),(E(6),E(1,3))]
        return MassActionModel(tuple(f'x{i}' for i in range(1,8)),[Reaction(a,b,ki) for (a,b),ki in zip(pairs,k)])
    def equilibrium(self,Xtotal,Ytotal):
        Xt,Yt=map(R,[Xtotal,Ytotal]);k,alpha,C,K,L=self.constants();U=Yt-alpha
        if Xt<=0:raise ValueError('Positive X total required.')
        if U<=0:return dict(status='no_positive_equilibrium',alpha=alpha)
        B=C+K*Xt-L*U;y=(-B+s.sqrt(B*B+4*L*C*U))/(2*L);t=Xt*y/(L*y+C);state=((k[1]+k[2])*t/k[0],t,C*t/y,y,alpha,k[2]*t/k[5],k[2]*t/k[8]);return dict(status='positive_equilibrium',state=tuple(map(s.simplify,state)),alpha=alpha)

def envz_box_floors(rate_boxes,Xbox,Ybox):
    k=tuple(rate_boxes);Xm,Xp=Xbox.lo,Xbox.hi
    if len(k)!=9 or min(x.lo for x in k)<=0 or Xm<=0:raise ValueError('Positive parameter and X-total intervals required.')
    alpha=k[2]*(k[7]+k[8])/(k[6]*k[8]);C=k[2]*(k[4]+k[5])/(k[3]*k[5]);K=k[2]/k[5]+k[2]/k[8];L=(k[1]+k[2])/k[0]+1+K;U=Ybox.lo-alpha.hi
    if U<=0:return dict(status='not_certified_uniform_positive_region',alpha=alpha.strings())
    y=U/(1+K.hi*Xp/C.lo);t=Xm*y/(L.hi*y+C.hi);z=k[2].lo/k[8].hi*t;g1=(k[8].hi+2*(k[7].hi+k[8].hi))/(k[6].lo*k[8].lo*t);g2=(k[2].hi+2*k[6].hi*alpha.hi)/(k[6].lo*k[8].lo*z)
    return dict(status='certified_equilibrium_floors',alpha=alpha.strings(),y_floor=y,x2_floor=t,x7_floor=z,residual_gain_certificate1=g1,residual_gain_certificate2=g2,scope='Floors cover exact equilibria of the full box, not arbitrary transient states with these totals.')

@dataclass(frozen=True)
class PrivateRelease:
    reaction_index:int
    release_rates:tuple
    leakage_rates:tuple=()
    def apply(self,model):
        old=model.reactions[self.reaction_index];m=sum(old.product);p=m-1;n=len(model.species);lam=tuple(map(R,self.release_rates));delta=tuple(map(R,self.leakage_rates or (0,)*p))
        if sum(old.reactant)>2 or m<=2 or len(lam)!=p or len(delta)!=p or min(lam)<=0 or min(delta)<0:raise ValueError('At-most-bimolecular reactant, >2 products, m-1 positive release rates and nonnegative leaks required.')
        labels=tuple(f'private_{self.reaction_index}_{i+1}' for i in range(p))
        if set(labels)&set(model.species):raise ValueError('Private intermediate label collision.')
        extend=lambda x:x+(0,)*p;species=model.species+labels;reactions=[Reaction(extend(r.reactant),extend(r.product),r.rate) for i,r in enumerate(model.reactions) if i!=self.reaction_index];unit=lambda i:tuple(int(j==i) for j in range(n+p))
        reactions.append(Reaction(extend(old.reactant),unit(n),old.rate));products=[i for i,count in enumerate(old.product) for _ in range(count)];W=s.zeros(n,p)
        for j in range(p):
            dst=[0]*(n+p)
            for i in ([products[j]] if j<p-1 else products[j:]):dst[i]+=1;W[i,j]+=1
            if j<p-1:dst[n+j+1]=1
            reactions.append(Reaction(unit(n+j),tuple(dst),lam[j]))
            if delta[j]>0:reactions.append(Reaction(unit(n+j),(0,)*(n+p),delta[j]))
        return MassActionModel(species,reactions),W
    def steady_costs(self,flux,weights=None):
        F=R(flux);lam=tuple(map(R,self.release_rates));delta=tuple(map(R,self.leakage_rates or (0,)*len(lam)));weights=tuple(map(R,weights or (1,)*len(lam)))
        if not lam or len(weights)!=len(lam) or len(delta)!=len(lam) or min(lam)<=0 or min(delta)<0 or min(weights)<0 or F<0:raise ValueError('Invalid release-chain inputs.')
        pi=R(1);z=[];yields=[]
        for l,d in zip(lam,delta):z.append(F*pi/(l+d));pi*=l/(l+d);yields.append(pi)
        return dict(intermediate_concentrations=z,product_yields=yields,weighted_storage=sum(w*v for w,v in zip(weights,z)),lossless_storage_per_flux=sum(w/l for w,l in zip(weights,lam)),scope='Steady intermediate balances only; existence of a full leaking-network equilibrium is not implied.')
