"""Retained-state redox dynamics with composable maintained and finite supplies."""
from dataclasses import dataclass,replace
from typing import Protocol
from math import log
import numpy as np
from scipy.integrate import solve_ivp
from branches import nominal,MaintainedKinetics


class SupplyLaw(Protocol):
    initial_stock:float
    finite:bool
    def scale(self,time,stock):...


@dataclass(frozen=True)
class MaintainedSupply:
    strength:float=.12
    amplitude:float=0.
    frequency:float=1.
    initial_stock:float=0.
    finite:bool=False
    def __post_init__(self):
        if not self.strength>abs(self.amplitude) or self.frequency<0:raise ValueError('Maintained supply must stay positive.')
    def scale(self,time,stock):return self.strength+self.amplitude*np.sin(self.frequency*time)


@dataclass(frozen=True)
class LinearDonor:
    initial_stock:float=120.
    conversion:float=1000.
    finite:bool=True
    def __post_init__(self):
        if min(self.initial_stock,self.conversion)<=0:raise ValueError('Positive donor and concentration conversion required.')
    def scale(self,time,stock):return stock/self.conversion


@dataclass(frozen=True)
class SaturatingDonor:
    initial_stock:float=200.
    affinity:float=.01
    initial_strength:float=.12
    finite:bool=True
    def __post_init__(self):
        if min(self.initial_stock,self.affinity,self.initial_strength)<=0:raise ValueError('Positive saturating supply parameters required.')
    def scale(self,time,stock):return self.initial_strength*stock/(self.affinity+stock)*(self.affinity+self.initial_stock)/self.initial_stock


class OperatingModel:
    def __init__(self,kinetics=None,repair_scale=1.,quotas=(10.,4.)):
        base=kinetics or nominal()
        if repair_scale<=0 or len(quotas)!=2 or min(quotas)<=0:raise ValueError('Positive repair scale and two quotas required.')
        self.kinetics=MaintainedKinetics(base.gpx,replace(base.trx,b=base.trx.b*repair_scale,c=base.trx.c*repair_scale),base.source)
        self.repair_scale=repair_scale;self.quotas=np.array(quotas,float)
    def service(self,u):
        s=self.kinetics.full(u);return np.array([self.kinetics.gpx.a*s['e0'],self.kinetics.trx.e*s['y']*s['v']])
    def storage(self,u):
        x,z,e1,e2,zt,h,w,v=u
        G,T=self.kinetics.gpx,self.kinetics.trx
        normalizer=G.G/2-G.z0+G.E+T.T-T.z0
        return normalizer+x-(z-G.z0+e1+e2)-(zt-T.z0)
    def necessary_repair_delay(self,initial_w):
        T=self.kinetics.trx;wq=T.E-self.quotas[1]/(T.e*(T.T-T.z0))
        if wq<=0:return np.inf
        return max(0.,log(initial_w/wq)/T.c) if initial_w>0 else 0.
    def solve(self,initial,supply,times):
        times=np.asarray(times,float);u0=np.asarray(initial,float)
        if u0.shape!=(8,) or self.kinetics.margins(u0).min()<-1e-12 or len(times)<2 or times[0]!=0 or not np.all(np.diff(times)>0):raise ValueError('Physical eight-state preparation and increasing times from zero required.')
        # state, donor, integrated G/T services, regeneration, G/T shortfalls.
        y0=np.r_[u0,supply.initial_stock,np.zeros(5)]
        def rhs(t,y):
            u=y[:8];scale=supply.scale(t,y[8]);current=scale*self.kinetics.source.unit_current(u[0]);services=self.service(u)
            return np.r_[self.kinetics.rhs(t,u,scale),-current if supply.finite else 0.,services,current,np.maximum(self.quotas-services,0)]
        def crossing(t,y):return self.service(y[:8])[1]-self.quotas[1]
        crossing.direction=0
        sol=solve_ivp(rhs,(0,times[-1]),y0,t_eval=times,events=crossing,method='Radau',rtol=2e-10,atol=2e-12)
        if not sol.success:raise RuntimeError(sol.message)
        states=sol.y.T;physical=min(self.kinetics.margins(y[:8]).min() for y in states)
        if physical<-1e-7 or (supply.finite and states[:,8].min()<-1e-7):raise RuntimeError('Numerical physical-domain violation.')
        W=np.array([self.storage(y[:8]) for y in states]);total_service=states[:,9:11].sum(axis=1)
        ledger=W-W[0]-states[:,11]+total_service
        donor_residual=states[:,8]+states[:,11]-supply.initial_stock if supply.finite else np.zeros(len(times))
        return dict(times=times,states=states[:,:8],donor=states[:,8],services=np.array([self.service(y[:8]) for y in states]),
            integrated_services=states[:,9:11],regeneration=states[:,11],shortfalls=states[:,12:14],storage=W,
            scales=np.array([supply.scale(t,y[8]) for t,y in zip(times,states)]),crossings=sol.t_events[0],
            storage_identity_residual=float(max(abs(ledger))),donor_identity_residual=float(max(abs(donor_residual))),minimum_physical_margin=float(physical),
            evidence='Numerical full retained-state trajectory. Sampled margins and quota crossings are diagnostics, not a uniform preparation certificate.')
