"""Full distributive phosphorylation network with retained complexes and shared enzymes."""
from dataclasses import dataclass
from fractions import Fraction as F
import math
import numpy as np
from scipy.integrate import solve_ivp


@dataclass(frozen=True)
class EdgeKinetics:
    association: object = 1.
    dissociation: object = 1.
    catalysis: object = 1.
    reverse_association: object = 1.
    reverse_dissociation: object = 1.
    reverse_catalysis: object = 1.

    def __post_init__(self):
        if any(not math.isfinite(float(v)) or v<=0 for v in self.values()):
            raise ValueError('all six elementary rate constants must be finite and positive')

    def values(self):return (self.association,self.dissociation,self.catalysis,self.reverse_association,self.reverse_dissociation,self.reverse_catalysis)
    @property
    def p(self):return self.association/(self.dissociation+self.catalysis)
    @property
    def q(self):return self.reverse_association/(self.reverse_dissociation+self.reverse_catalysis)
    @property
    def up(self):return self.p*self.catalysis
    @property
    def down(self):return self.q*self.reverse_catalysis
    @property
    def ratio(self):return self.up/self.down

    def scaled(self,clock):
        if clock<=0:raise ValueError('clock multiplier must be positive')
        return EdgeKinetics(*(clock*v for v in self.values()))


@dataclass(frozen=True)
class Reaction:
    name: str
    rate: object
    reactants: tuple
    products: tuple


class PhosphorylationNetwork:
    """A graph whose edges each have a distinct kinase and phosphatase complex."""
    def __init__(self,vertices,edges,kinetics,reverse_activity=0.):
        self.vertices=tuple(vertices);self.edges=tuple(edges);self.kinetics=tuple(kinetics)
        self.q=len(vertices);self.m=len(edges);self.E=self.q;self.F=self.q+1;self.size=self.q+2+2*self.m
        if len(kinetics)!=self.m or len(set(vertices))!=self.q or len(set(edges))!=self.m:
            raise ValueError('one rate set per distinct edge and distinct vertex labels required')
        if not math.isfinite(float(reverse_activity)) or reverse_activity<0:raise ValueError('reverse activity must be nonnegative')
        self.reverse_activity=reverse_activity
        if any(a==b or not 0<=a<self.q or not 0<=b<self.q for a,b in edges):raise ValueError('invalid edge indices')
        self.names=[f'S{v}' for v in vertices]+['E','F']+[f'C{i}' for i in range(self.m)]+[f'Y{i}' for i in range(self.m)]
        self.reactions=[]
        for i,((a,b),r) in enumerate(zip(edges,kinetics)):
            C=self.q+2+i;Y=self.q+2+self.m+i
            definitions=[('E_bind',r.association,(a,self.E),(C,)),('E_release',r.dissociation,(C,),(a,self.E)),('E_cat',r.catalysis,(C,),(b,self.E)),
                ('F_bind',r.reverse_association,(b,self.F),(Y,)),('F_release',r.reverse_dissociation,(Y,),(b,self.F)),('F_cat',r.reverse_catalysis,(Y,),(a,self.F))]
            if reverse_activity:
                definitions += [('E_cat_reverse',reverse_activity*r.association*r.catalysis/r.dissociation,(b,self.E),(C,)),('F_cat_reverse',reverse_activity*r.reverse_association*r.reverse_catalysis/r.reverse_dissociation,(a,self.F),(Y,))]
            for name,rate,rea,pro in definitions:self.reactions.append(Reaction(f'{i}:{name}',rate,rea,pro))
        self.stoichiometry=np.zeros((self.size,len(self.reactions)),dtype=int)
        for j,r in enumerate(self.reactions):
            for k in r.reactants:self.stoichiometry[k,j]-=1
            for k in r.products:self.stoichiometry[k,j]+=1
        self.inventory=np.zeros((3,self.size),dtype=int)
        self.inventory[0,self.E]=1;self.inventory[0,self.q+2:self.q+2+self.m]=1
        self.inventory[1,self.F]=1;self.inventory[1,self.q+2+self.m:]=1
        self.inventory[2,:self.q]=1;self.inventory[2,self.q+2:]=1
        assert not np.any(self.inventory@self.stoichiometry)
        self.chart_indices=np.r_[np.arange(1,self.q),np.arange(self.q+2,self.size)]
        self.chart=np.zeros((self.size,self.size-3),dtype=int)
        self.chart[self.chart_indices,np.arange(self.size-3)]=1;self.chart[0,:]=-1
        self.chart[self.E,self.q-1:self.q-1+self.m]=-1;self.chart[self.F,self.q-1+self.m:]=-1

    @classmethod
    def cube(cls,sites,kinetics,reverse_activity=0.):
        if type(sites) is not int or not 1<=sites<=4:raise ValueError('this explicit builder supports 1..4 sites')
        edges=[(v,v|1<<i) for v in range(2**sites) for i in range(sites) if not v>>i&1]
        return cls(range(2**sites),edges,kinetics,reverse_activity)

    @classmethod
    def symmetric_lift(cls,chain_kinetics,reverse_activity=0.):
        n=len(chain_kinetics);kinetics=[]
        for v in range(2**n):
            k=v.bit_count()
            for i in range(n):
                if not v>>i&1:
                    r=chain_kinetics[k]
                    kinetics.append(EdgeKinetics(r.association/(n-k),r.dissociation,r.catalysis,r.reverse_association/(k+1),r.reverse_dissociation,r.reverse_catalysis))
        return cls.cube(n,kinetics,reverse_activity)

    @classmethod
    def chain(cls,kinetics,reverse_activity=0.):
        n=len(kinetics);return cls(range(n+1),[(i,i+1) for i in range(n)],kinetics,reverse_activity)

    def concentrations(self,chart_state,totals):
        if len(chart_state)!=self.size-3 or len(totals)!=3:raise ValueError('invalid state or total dimensions')
        result=self.chart@np.asarray(chart_state,float)
        result[0]+=totals[2];result[self.E]+=totals[0];result[self.F]+=totals[1]
        return result

    def rates(self,x):
        return np.array([float(r.rate)*math.prod(x[k] for k in r.reactants) for r in self.reactions])

    def rhs(self,t,x):return self.stoichiometry@self.rates(x)

    def jacobian(self,x):
        gradient=np.zeros((len(self.reactions),self.size))
        for i,r in enumerate(self.reactions):
            for k in r.reactants:gradient[i,k]=float(r.rate)*math.prod(x[h] for h in r.reactants if h!=k)
        return self.stoichiometry@gradient

    def simulate(self,initial,times):
        initial=np.asarray(initial,float);times=np.asarray(times,float)
        if initial.shape!=(self.size,) or np.any(~np.isfinite(initial)) or np.any(initial<0):raise ValueError('invalid nonnegative initial state')
        if times.ndim!=1 or len(times)<2 or times[0]!=0 or np.any(np.diff(times)<=0) or np.any(~np.isfinite(times)):raise ValueError('finite increasing grid must start at zero')
        totals=self.inventory@initial
        def rhs(t,z):return self.rhs(t,self.concentrations(z,totals))[self.chart_indices]
        def jac(t,z):return self.jacobian(self.concentrations(z,totals))[self.chart_indices]@self.chart
        sol=solve_ivp(rhs,(0,times[-1]),initial[self.chart_indices],t_eval=times,method='Radau',jac=jac,rtol=2e-9,atol=2e-11)
        if not sol.success:raise RuntimeError(sol.message)
        states=np.array([self.concentrations(z,totals) for z in sol.y.T])
        if states.min() < -1e-7:raise RuntimeError('trajectory left positive class beyond numerical tolerance')
        return states

    def ssa(self,initial,volume,horizon,seed=72,event_budget=200000):
        counts=np.array(initial,dtype=object)
        if counts.shape!=(self.size,) or any(not isinstance(x,(int,np.integer)) or x<0 for x in counts):raise ValueError('counts must be nonnegative integers')
        if not math.isfinite(volume) or volume<=0 or not math.isfinite(horizon) or horizon<0:raise ValueError('invalid volume or horizon')
        rng=np.random.default_rng(seed);t=0.;times=[t];history=[counts.copy()];events=[]
        for _ in range(event_budget):
            a=np.array([float(r.rate)*math.prod(counts[k] for k in r.reactants)/volume**(len(r.reactants)-1) for r in self.reactions],float)
            total=a.sum()
            if not np.isfinite(total):raise RuntimeError('propensity overflow; use an appropriate scale')
            if total==0:break
            dt=rng.exponential(1/total)
            if t+dt>horizon:break
            t+=dt;j=int(rng.choice(len(a),p=a/total));counts=counts+self.stoichiometry[:,j]
            if min(counts)<0:raise AssertionError('negative count')
            times.append(t);history.append(counts.copy());events.append(j)
        else:raise RuntimeError('SSA event budget reached before horizon; no completed path returned')
        times.append(horizon);history.append(counts.copy())
        return np.array(times),np.array(history,dtype=object),events

    def readout(self,x):
        return x[self.q-1]+sum(x[self.q+2+self.m+i] for i,(a,b) in enumerate(self.edges) if b==self.q-1)

    def cube_aggregation(self):
        n=(self.q-1).bit_length()
        if self.q!=2**n or self.m!=n*2**(n-1):raise ValueError('aggregation requires the full cube')
        A=np.zeros((3*n+3,self.size),dtype=int)
        for v in range(self.q):A[v.bit_count(),v]=1
        A[n+1,self.E]=1;A[n+2,self.F]=1
        for i,(a,b) in enumerate(self.edges):
            k=a.bit_count();A[n+3+k,self.q+2+i]=1;A[2*n+3+k,self.q+2+self.m+i]=1
        return A

    def square_ratios(self):
        n=(self.q-1).bit_length();lookup={e:r.ratio for e,r in zip(self.edges,self.kinetics)};result=[]
        for v in range(self.q):
            for i in range(n):
                for j in range(i+1,n):
                    if not(v>>i&1 or v>>j&1):
                        a=v|1<<i;b=v|1<<j;c=a|1<<j
                        if all(e in lookup for e in [(v,a),(a,c),(v,b),(b,c)]):result.append((v,i,j,lookup[v,a]*lookup[a,c]/(lookup[v,b]*lookup[b,c])))
        return result


@dataclass(frozen=True)
class QuadraticMemoryContract:
    pmax: F
    exit_level: F
    decay: F
    noise: F
    dimension: int

    def __post_init__(self):
        if min(self.pmax,self.exit_level,self.decay,self.noise)<=0 or self.dimension<1:raise ValueError('positive verified constants required')

    def sufficient_size(self,recovery,storage):
        return max(F(1),1000*self.noise*storage/self.exit_level,10**6*self.noise/(self.decay*self.exit_level),1024*self.dimension*self.pmax/self.exit_level)

    def failures(self,volume,recovery,storage):
        if min(volume,recovery,storage)<=0 or recovery<8/self.decay:raise ValueError('invalid system size or insufficient recovery duration')
        retention=F(1,4096)+self.noise*storage/(volume*self.exit_level)
        recovery_failure=F(1,1024)+self.noise*recovery/(volume*self.exit_level)+F(1,256)+4096*self.noise/(self.decay*volume*self.exit_level)
        return retention,recovery_failure,retention+recovery_failure

    def rescale_clock(self,factor):
        if factor<=0:raise ValueError('clock factor must be positive')
        return QuadraticMemoryContract(self.pmax,self.exit_level,self.decay*factor,self.noise*factor,self.dimension)
