"""Literal mass-action reactor and composable single-rate controls.

Chart trajectories conserve all three totals by construction. Numerical results
are not interval orbit or switching certificates.
"""
from dataclasses import dataclass
from typing import Protocol
import numpy as np
from scipy.integrate import solve_ivp
from scipy.optimize import root
from clock import FloatSystem,rate_list,RATE_NAMES

class Signal(Protocol):
    def value(self,t:float)->float:...

@dataclass(frozen=True)
class ZeroSignal:
    def value(self,t):return 0.

@dataclass(frozen=True)
class ResonantPulse:
    amplitude:float
    period:float
    duration:float
    phase:float
    def __post_init__(self):
        if not all(np.isfinite(v) for v in (self.amplitude,self.period,self.duration,self.phase)) or not 0<=self.amplitude<1 or min(self.period,self.duration)<=0:raise ValueError('Finite pulse parameters, amplitude below one and positive times required.')
    def value(self,t):
        if t<=0 or t>=self.duration:return 0.
        return self.amplitude*np.sin(2*np.pi*t/self.period+self.phase)*np.sin(np.pi*t/self.duration)**2

class MassActionReactor(FloatSystem):
    def __init__(self,model,rates=None,input_rate='alpha1',signal=None):
        if model.n!=3:raise ValueError('This paper-local reactor implements the three-site network.')
        rates=np.array(list(map(float,rate_list(model))) if rates is None else rates,float)
        if rates.shape!=(18,) or not np.all(np.isfinite(rates)) or np.min(rates)<=0:raise ValueError('Eighteen positive finite rates required.')
        if input_rate not in RATE_NAMES:raise ValueError('Unknown reaction rate.')
        super().__init__(model,rates=rates);self.rates=rates.copy();self.input_index=RATE_NAMES.index(input_rate);self.signal=signal or ZeroSignal()
        self.conservation=np.zeros((3,12));self.conservation[0,[4,6,7,8]]=1;self.conservation[1,[5,9,10,11]]=1;self.conservation[2,[0,1,2,3,6,7,8,9,10,11]]=1
    def controlled(self,signal,input_rate=None):return MassActionReactor(self.m,self.rates,input_rate or RATE_NAMES[self.input_index],signal)
    def equilibrium(self,guess=None):
        result=root(lambda y:self.f(0,y),np.zeros(9) if guess is None else guess,jac=lambda y:self.jac(0,y),tol=1e-10)
        residual=float(max(abs(self.f(0,result.x))))
        if residual>1e-9 or np.min(self.species(result.x))<=0:raise RuntimeError('No converged positive equilibrium; provide another guess.')
        return result.x,residual
    def modulation(self,t):
        u=float(self.signal.value(t))
        if not np.isfinite(u) or u<=-1:raise ValueError('Control would produce a nonpositive rate.')
        return u
    def flux(self,x,u=0.):
        v=np.empty(18);v[0::3]=self.kon*x[self.lev]*x[self.enz];v[1::3]=self.koff*x[self.cpx];v[2::3]=self.kcat*x[self.cpx];v[self.input_index]*=1+u;return v
    def f(self,t,y,u=None):return self.RN@self.flux(self.species(y),self.modulation(t) if u is None else u)
    def jac(self,t,y,u=None):
        x=self.species(y);D=np.zeros((18,12));j=np.arange(6);D[3*j,self.lev]=self.kon*x[self.enz];D[3*j,self.enz]=self.kon*x[self.lev];D[3*j+1,self.cpx]=self.koff;D[3*j+2,self.cpx]=self.kcat;D[self.input_index]*=1+(self.modulation(t) if u is None else u)
        return self.RN@D@self.P
    def integrate(self,y,duration,method='LSODA',samples=1001,events=None,ledger=False):
        y=np.array(y,float)
        if y.shape!=(9,) or not np.all(np.isfinite(y)) or np.min(self.species(y))<=0 or not np.isfinite(duration) or duration<=0:raise ValueError('Positive duration and physical initial state required.')
        if ledger:
            initial=np.r_[y,np.zeros(18)];rhs=lambda t,w:np.r_[self.f(t,w[:9]),self.flux(self.species(w[:9]),self.modulation(t))];jac=None
        else:initial=y;rhs=self.f;jac=self.jac
        sol=solve_ivp(rhs,(0,duration),initial,method=method,jac=jac,rtol=1e-10,atol=1e-13,max_step=.25,dense_output=True,events=events)
        if not sol.success:raise RuntimeError(sol.message)
        t=np.linspace(0,sol.t[-1],samples);Y=sol.sol(t);X=self.x0[:,None]+self.P@Y[:9]
        if np.min(X)<=0:raise ArithmeticError('Numerical positivity failed; refine the experiment.')
        return dict(time=t,chart=Y[:9],species=X,control=np.array([self.modulation(z) for z in t]),endpoint=sol.y[:9,-1],elapsed=float(sol.t[-1]),events=sol.t_events,ledger=sol.y[9:,-1] if ledger else None,conservation_drift=float(np.max(abs(self.conservation@X-(self.conservation@self.x0)[:,None]))),minimum_concentration=float(X.min()))
    def shoot(self,seed,period,max_iterations=15):
        y=np.array(seed,float);reference=y.copy();normal=self.f(0,y);T=float(period)
        def transport(y,T):
            def rhs(t,z):return np.r_[self.f(t,z[:9]),(self.jac(t,z[:9])@z[9:].reshape(9,9)).ravel()]
            sol=solve_ivp(rhs,(0,T),np.r_[y,np.eye(9).ravel()],method='LSODA',rtol=2e-12,atol=2e-14)
            if not sol.success:raise RuntimeError(sol.message)
            return sol.y[:9,-1],sol.y[9:,-1].reshape(9,9)
        for _ in range(max_iterations):
            end,M=transport(y,T);res=np.r_[end-y,normal@(y-reference)]
            if np.max(abs(res))<2e-10:break
            A=np.zeros((10,10));A[:9,:9]=M-np.eye(9);A[:9,9]=self.f(0,end);A[9,:9]=normal
            delta=np.linalg.solve(A,-res);y+=delta[:9];T+=delta[9]
            if T<=0 or np.min(self.species(y))<=0:raise RuntimeError('Shooting left the physical domain.')
        end,M=transport(y,T);residual=float(np.max(abs(end-y)))
        if residual>1e-8:raise RuntimeError('Periodic shooting failed to converge.')
        multipliers=np.linalg.eigvals(M);trivial=int(np.argmin(abs(multipliers-1)));nontrivial=np.delete(multipliers,trivial)
        return dict(anchor=y,period=T,residual=residual,multipliers=multipliers,largest_nontrivial=float(max(abs(nontrivial))),scope='Numerical shooting and variational equations. Existence, minimal period and attraction require a separate validated proof.')

class SwitchingExperiment:
    def __init__(self,reactor,period,section_x,section_z,amplitude=.1,pulse_periods=3):
        self.reactor=reactor;self.period=period;self.x=np.asarray(section_x);self.z=np.asarray(section_z);self.amplitude=amplitude;self.duration=pulse_periods*period
    def pulse(self,state,phase,method='LSODA'):
        signal=ResonantPulse(self.amplitude,self.period,self.duration,phase)
        return self.reactor.controlled(signal).integrate(state,self.duration,method=method,samples=601)
    def phase_trigger(self,state):
        # Depart from a possible initial zero before looking for another crossing.
        segments=[];departure=self.reactor.integrate(state,.05*self.period,samples=21);segments.append(departure);state=departure['endpoint']
        for _ in range(4):
            def event(t,y):return self.z@y
            event.terminal=True;event.direction=0
            run=self.reactor.integrate(state,2*self.period,samples=301,events=event);segments.append(run);state=run['endpoint']
            if len(run['events'][0]) and self.x@state>0:return state,segments
            run=self.reactor.integrate(state,.05*self.period,samples=21);segments.append(run);state=run['endpoint']
        raise RuntimeError('No positive section crossing found; OFF has not been triggered.')
    def run(self,on_phase=np.pi,off_phase=np.pi/3,wait_periods=14,rest_state=None):
        rest=self.reactor.integrate(np.zeros(9) if rest_state is None else rest_state,2*self.period,samples=101)
        on=self.pulse(rest['endpoint'],on_phase)
        coast=self.reactor.integrate(on['endpoint'],wait_periods*self.period,samples=1501)
        state,triggers=self.phase_trigger(coast['endpoint']);off=self.pulse(state,off_phase)
        final=self.reactor.integrate(off['endpoint'],14*self.period,samples=1501)
        return [rest,on,coast,*triggers,off,final],on,off
    def terminal_window(self,state,periods,rest_state=None):
        if periods<=2:raise ValueError('Settling duration must exceed the two-period readout window.')
        settle=self.reactor.integrate(state,(periods-2)*self.period,samples=101)
        window=self.reactor.integrate(settle['endpoint'],2*self.period,samples=1201)
        return dict(range=float(np.ptp(window['species'][3]+window['species'][11])),distance_to_equilibrium=float(np.linalg.norm(window['endpoint']-(np.zeros(9) if rest_state is None else rest_state))),periods=periods,scope='Finite-time numerical readout; no certified basin or settling time.')
