"""Equilibrium geometry, independent kinetics and full mass-action dynamics."""
from dataclasses import dataclass
from fractions import Fraction as Q
import numpy as np
from scipy.integrate import solve_ivp
import phos_sharp as ps

class Reactor:
    """n-site physical kinetics in class coordinates q=(S1..Sn,C1..Cn,Y1..Yn).
    Independent totals and rates can be changed; no equilibrium is hard-coded.
    """
    def __init__(self,rates,totals):
        self.rates=np.array(rates,float);self.n=len(rates);n=self.n;self.totals=np.asarray(totals,float)
        if n<1 or self.rates.shape!=(n,6) or self.totals.shape!=(3,) or min(self.rates.ravel())<=0 or min(self.totals)<=0 or not np.all(np.isfinite(self.rates)) or not np.all(np.isfinite(self.totals)):raise ValueError('Positive finite six-rate site rows and three totals required.')
        self.indices=np.array(list(range(1,n+1))+list(range(n+3,3*n+3)));self.P=np.zeros((3*n+3,3*n));self.P[0]=-1;self.P[self.indices,np.arange(3*n)]=1;self.P[n+1,n:2*n]=-1;self.P[n+2,2*n:]=-1
        self.offset=np.zeros(3*n+3);self.offset[0]=self.totals[2];self.offset[n+1:n+3]=self.totals[:2]
        self.N=np.zeros((3*n+3,6*n));self.reactants=[]
        for i in range(n):
            E,F,C,Y=n+1,n+2,n+3+i,2*n+3+i
            for j,(reac,prod) in enumerate([([i,E],[C]),([C],[i,E]),([C],[i+1,E]),([i+1,F],[Y]),([Y],[i+1,F]),([Y],[i,F])]):
                for k in reac:self.N[k,6*i+j]-=1
                for k in prod:self.N[k,6*i+j]+=1
                self.reactants.append(reac)
        self.labels=[f'S{i}' for i in range(n+1)]+['E','F']+[f'C{i+1}' for i in range(n)]+[f'Y{i+1}' for i in range(n)]
    def species(self,q):return self.offset+self.P@np.asarray(q)
    def flux(self,z):return np.array([k*np.prod(z[reac]) for k,reac in zip(self.rates.ravel(),self.reactants)])
    def field(self,t,q):return (self.N@self.flux(self.species(q)))[self.indices]
    def jacobian(self,t,q):
        z=self.species(q);D=np.zeros((6*self.n,3*self.n+3))
        for j,(k,react) in enumerate(zip(self.rates.ravel(),self.reactants)):
            for l,s in enumerate(react):D[j,s]+=k*np.prod([z[v] for m,v in enumerate(react) if m!=l])
        return (self.N@D)[self.indices]@self.P
    def integrate(self,state,duration,samples=1001,method='LSODA'):
        state=np.asarray(state,float)
        if state.shape!=(3*self.n+3,) or np.min(state)<=0 or duration<=0:raise ValueError('Physical full initial state and positive duration required.')
        q=state[self.indices]
        if max(abs(self.species(q)-state))>1e-9:raise ValueError('Initial state does not have the selected conserved totals.')
        sol=solve_ivp(self.field,(0,duration),q,method=method,jac=self.jacobian,rtol=2e-10,atol=1e-13,dense_output=True)
        if not sol.success:raise RuntimeError(sol.message)
        times=np.r_[0,np.geomspace(max(duration*1e-7,1e-7),duration,samples-1)];Y=sol.sol(times);X=self.offset[:,None]+self.P@Y
        if np.min(X)<=0:raise ArithmeticError('Numerical positivity failed.')
        ledger=max(max(abs(np.array(ps.totals(list(z)),float)-self.totals)) for z in X.T)
        return dict(time=times,species=X,endpoint=X[:,-1],conservation_drift=float(ledger),minimum_concentration=float(np.min(X)))
