"""A finite-type inherited-mark branching source with correlated daughters."""
from dataclasses import dataclass
from fractions import Fraction as F
from math import comb,log,log1p,expm1
import numpy as np
from scipy.integrate import solve_ivp


@dataclass(frozen=True)
class MolecularRates:
    sites:int=2
    basal_writing:F=F(1,100)
    recruited_writing:F=F(1)
    recruited_erasure:F=F(1,2)
    intrinsic_erasure:F=F(1,100)
    division:F=F(1,10)
    protected_division:F=F(1,10)
    protected_death:F=F(1,100)
    unprotected_death:F=F(3,10)

    def __post_init__(self):
        if self.sites<1 or not isinstance(self.sites,int):raise ValueError('Positive integer site count required.')
        for name in self.__dataclass_fields__:
            if name!='sites':
                value=F(getattr(self,name));object.__setattr__(self,name,value)
                if value<0:raise ValueError('Nonnegative rates required.')


@dataclass(frozen=True)
class Action:
    eraser:F=F(0)
    protected_death:F=F(0)
    def __post_init__(self):
        for name in ['eraser','protected_death']:
            value=F(getattr(self,name));object.__setattr__(self,name,value)
            if value<0:raise ValueError('Actions are additional nonnegative rates.')


@dataclass(frozen=True)
class Phase:
    duration:F
    action:Action
    def __post_init__(self):
        object.__setattr__(self,'duration',F(self.duration))
        if self.duration<=0:raise ValueError('Positive phase duration required.')


class InheritedPopulation:
    def __init__(self,rates=MolecularRates()):
        self.rates=rates;N=rates.sites
        self.states=tuple((a,r) for a in range(N+1) for r in range(N-a+1));self.index={s:i for i,s in enumerate(self.states)};self.size=len(self.states)
        self.protected=tuple(a>r for a,r in self.states)
        self.birth=tuple(rates.protected_division if p else rates.division for p in self.protected)
        self.death=tuple(rates.protected_death if p else rates.unprotected_death for p in self.protected)
        self.pairs=[]
        for a,r in self.states:
            self.pairs.append(tuple((self.index[(j,k)],self.index[(a-j,r-k)],F(comb(a,j)*comb(r,k),2**(a+r))) for j in range(a+1) for k in range(r+1)))

    def generator(self,action=Action()):
        N=self.rates.sites;Q=[[F(0)]*self.size for _ in self.states];e=self.rates.intrinsic_erasure+action.eraser
        for i,(a,r) in enumerate(self.states):
            s=N-a-r
            channels=[((a+1,r),s*(self.rates.basal_writing+self.rates.recruited_writing*F(a,N))),
                ((a,r+1),s*(self.rates.basal_writing+self.rates.recruited_writing*F(r,N))),
                ((a-1,r),a*(e+self.rates.recruited_erasure*F(r,N))),
                ((a,r-1),r*(e+self.rates.recruited_erasure*F(a,N)))]
            for target,rate in channels:
                if rate:Q[i][self.index[target]]+=rate;Q[i][i]-=rate
        return Q

    def daughter(self,x,y):return [sum((p*x[j]*y[k] for j,k,p in row),F(0)) for row in self.pairs]

    def phi(self,x,action=Action()):
        if len(x)!=self.size:raise ValueError('State-vector length mismatch.')
        Q=self.generator(action);D=self.daughter(x,x)
        return [sum((a*b for a,b in zip(Q[i],x)),F(0))+(self.death[i]+action.protected_death*self.protected[i])*(1-x[i])+self.birth[i]*(D[i]-x[i]) for i in range(self.size)]

    def mean_matrix(self,action=Action()):
        A=self.generator(action)
        for i,row in enumerate(self.pairs):
            for j,k,p in row:A[i][j]+=self.birth[i]*p;A[i][k]+=self.birth[i]*p
            A[i][i]-=self.birth[i]+self.death[i]+action.protected_death*self.protected[i]
        return A

    def eraser_direction(self,q):
        Q0=self.generator();Q1=self.generator(Action(F(1)))
        return [sum(((b-a)*x for a,b,x in zip(Q0[i],Q1[i],q)),F(0)) for i in range(self.size)]

    def numerical_flow(self,terminal,phase):
        Q=np.array(self.generator(phase.action),dtype=float);b=np.array(self.birth,dtype=float)
        d=np.array([self.death[i]+phase.action.protected_death*self.protected[i] for i in range(self.size)],dtype=float)
        terms=[(i,j,k,float(p)) for i,row in enumerate(self.pairs) for j,k,p in row]
        def fun(t,x):
            offspring=np.zeros(self.size)
            for i,j,k,p in terms:offspring[i]+=p*x[j]*x[k]
            return Q@x+d*(1-x)+b*(offspring-x)
        sol=solve_ivp(fun,(0,float(phase.duration)),np.asarray(terminal,dtype=float),method='DOP853',rtol=2e-11,atol=2e-13)
        if not sol.success:raise RuntimeError(sol.message)
        if sol.y[:,-1].min()<-1e-10 or sol.y[:,-1].max()>1+1e-10:raise RuntimeError('PGF solve left the probability cube.')
        return np.clip(sol.y[:,-1],0,1)

    def schedule(self,phases,terminal=None):
        """Deterministic chronological phases composed backwards from the endpoint."""
        value=np.zeros(self.size) if terminal is None else np.asarray(terminal,dtype=float)
        if value.shape!=(self.size,) or value.min()<0 or value.max()>1:raise ValueError('Terminal extinction vector must lie in the cube.')
        for phase in reversed(phases):value=self.numerical_flow(value,phase)
        return value

    def taylor_coefficients(self,action,order=4):
        Q=self.generator(action);d=[self.death[i]+action.protected_death*self.protected[i] for i in range(self.size)]
        J=[row[:] for row in Q]
        for i in range(self.size):J[i][i]-=d[i]+self.birth[i]
        cs=[[F(0)]*self.size]
        for k in range(order):
            products=[self.daughter(cs[j],cs[k-j]) for j in range(k+1)]
            cs.append([(sum((a*b for a,b in zip(J[i],cs[k])),F(0))+(d[i] if k==0 else 0)+self.birth[i]*sum(p[i] for p in products))/(k+1) for i in range(self.size)])
        return cs


def validate_configuration(z,size):
    if len(z)!=size or any(not isinstance(v,(int,np.integer)) or v<0 for v in z):raise ValueError('Founder configuration must contain nonnegative integer counts.')


def population_survival(single_extinction,z):
    validate_configuration(z,len(single_extinction))
    if any(not 0<=q<=1 for q in single_extinction):raise ValueError('Extinction probabilities must lie in [0,1].')
    if any(q==0 and n for q,n in zip(single_extinction,z)):return 1.
    return -expm1(sum(n*log(float(q)) for q,n in zip(single_extinction,z) if n))


class FeedbackSimulator:
    """Optional seeded SSA with common piecewise-constant feedback and pathwise cap.

    A callback chooses eraser from the current history after every jump or budget
    boundary. It must be constant between calls. Budget is common rate*time,
    never multiplied by cell count. Event/population caps return unresolved.
    """
    def __init__(self,source,amplitude,budget,seed=123):
        self.source=source;self.amplitude=float(amplitude);self.budget=float(budget);self.rng=np.random.default_rng(seed)
        if self.amplitude<0 or self.budget<0:raise ValueError('Nonnegative limits required.')

    def run(self,z,horizon,policy,max_events=100000,max_population=10000):
        validate_configuration(z,self.source.size);z=np.array(z,dtype=int);t=0.;spent=0.;events=0;history=[]
        while t<horizon and z.sum():
            if events>=max_events or z.sum()>max_population:return dict(status='unresolved_resource_cap',time=t,spent=spent,population=z.tolist(),events=events)
            v=float(policy(t,z.copy(),spent,tuple(history))) if spent<self.budget else 0.
            if not 0<=v<=self.amplitude:raise ValueError('Policy exceeds the admitted amplitude.')
            action=Action(F(str(v)));Q=self.source.generator(action);channels=[]
            for i,count in enumerate(z):
                if not count:continue
                for j,q in enumerate(Q[i]):
                    if j!=i and q:channels.append((float(count*q),('switch',i,j)))
                if self.source.death[i]:channels.append((float(count*self.source.death[i]),('death',i,None)))
                for j,k,p in self.source.pairs[i]:channels.append((float(count*self.source.birth[i]*p),('divide',i,(j,k))))
            rate=sum(x[0] for x in channels);wait=self.rng.exponential(1/rate) if rate else np.inf
            budget_time=(self.budget-spent)/v if v else np.inf;dt=min(wait,horizon-t,budget_time)
            t+=dt;spent=min(self.budget,spent+v*dt)
            if dt==budget_time:spent=self.budget;history.append((t,'budget_exhausted'));continue
            if dt<wait:break
            selector=self.rng.uniform(0,rate)
            for propensity,event in channels:
                selector-=propensity
                if selector<=0:break
            kind,i,target=event;z[i]-=1
            if kind=='switch':z[target]+=1
            elif kind=='divide':z[target[0]]+=1;z[target[1]]+=1
            events+=1;history.append((t,kind,int(i)))
        return dict(status='extinct' if z.sum()==0 else 'alive_at_horizon',time=t,spent=spent,population=z.tolist(),events=events,
            scope='One stochastic path; not a risk bound or policy optimality certificate.')
