"""Finite-state branching sources and their unconditioned stopped record laws.

Rates are per hour, delays in hours. Rational inputs give exact immediate-read
laws and inverses. Delayed channels and time trajectories are numerical.
"""
from dataclasses import dataclass
import numpy as np
import sympy as sp
from scipy.linalg import expm
from scipy.integrate import solve_ivp


def matrix(values):return sp.Matrix(values).applyfunc(lambda x:sp.Rational(str(x)))


@dataclass(frozen=True)
class MarkerChannel:
    emissions: object

    def __post_init__(self):
        E=matrix(self.emissions)
        if any(v<0 for v in E) or any(sum(E[:,j])!=1 for j in range(E.cols)):raise ValueError('Marker columns must be probability vectors.')
        if E.rank()!=E.cols:raise ValueError('Marker does not distinguish every latent state.')
        object.__setattr__(self,'emissions',E)

    def left_inverse(self):
        E=self.emissions
        return (E.T*E).inv()*E.T


@dataclass(frozen=True)
class BranchingSource:
    preparation: object
    switching: object                  # off-diagonal rates, zero diagonal
    division: object
    death: object
    sisters: object                    # m x m^2, ordered pairs in row-major order

    def __post_init__(self):
        for field in ('preparation','switching','division','death','sisters'):object.__setattr__(self,field,matrix(getattr(self,field)))
        m=len(self.preparation)
        if self.preparation.shape!=(m,1) or self.division.shape!=(m,1) or self.death.shape!=(m,1) or self.switching.shape!=(m,m) or self.sisters.shape!=(m,m*m):raise ValueError('Source dimensions disagree.')
        if any(p<=0 for p in self.preparation) or sum(self.preparation)!=1:raise ValueError('Identification requires positive preparation on every state.')
        if any(v<0 for field in ('switching','division','death','sisters') for v in getattr(self,field)):raise ValueError('Negative rates or probabilities.')
        if any(self.switching[i,i]!=0 for i in range(m)):raise ValueError('Supply off-diagonal switching rates; diagonal must be zero.')
        if any(sum(self.sisters[i,:])!=1 for i in range(m)):raise ValueError('Each sister row must sum to one.')

    @property
    def states(self):return len(self.preparation)

    @property
    def killed(self):
        return self.switching-sp.diag(*[sum(self.switching[i,:])+self.division[i]+self.death[i] for i in range(self.states)])

    @property
    def exits(self):return self.death.row_join(sp.diag(*self.division)*self.sisters)

    def mean_generator(self):
        m=self.states
        offspring=sp.zeros(m,m)
        for i in range(m):
            for j in range(m):offspring[i,j]=sum(self.sisters[i,j*m+k]+self.sisters[i,k*m+j] for k in range(m))
        return self.killed+sp.diag(*self.division)*offspring

    def population(self,times,z=0.):
        """PGF at z, with z=0 giving extinction by time; numerical ODE."""
        t=np.asarray(times,float);m=self.states
        if len(t)<2 or t[0]!=0 or np.any(np.diff(t)<=0) or not 0<=z<=1:raise ValueError('Increasing time grid starting at zero and z in [0,1] required.')
        H=np.array(self.killed,float);b=np.array(self.division,float).ravel();d=np.array(self.death,float).ravel();K=np.array(self.sisters,float)
        def rhs(_,f):return d+H@f+b*(K@np.outer(f,f).ravel())
        sol=solve_ivp(rhs,(0,t[-1]),np.full(m,z),t_eval=t,rtol=1e-10,atol=1e-12)
        if not sol.success:raise RuntimeError(sol.message)
        pi=np.array(self.preparation,float).ravel();G=np.array(self.mean_generator(),float)
        return sol.y,pi@sol.y,np.array([pi@expm(x*G)@np.ones(m) for x in t])

    def concordance_shift(self,epsilon):
        if self.states!=2:raise ValueError('This fixed-marginal perturbation is specific to two states.')
        shift=matrix(epsilon)*sp.Matrix([[1,-1,-1,1]])
        return BranchingSource(self.preparation,self.switching,self.division,self.death,self.sisters+shift)


@dataclass(frozen=True)
class RandomDeadline:
    rate: object = '1'
    daughter_delay: float = 0.

    def __post_init__(self):
        object.__setattr__(self,'rate',sp.Rational(str(self.rate)))
        if self.rate<=0 or self.daughter_delay<0:raise ValueError('Positive deadline rate, nonnegative delay required.')

    def law(self,source,marker):
        E=marker.emissions;m=source.states
        if E.cols!=m:raise ValueError('Marker/source dimensions disagree.')
        R=(self.rate*sp.eye(m)-source.killed).inv();A=self.rate*R
        J=E*sp.diag(*source.preparation)*A*E.T
        if not self.daughter_delay:
            F=E;C=sp.diag(sp.ones(1,1),sp.kronecker_product(F,F))
            V=E*sp.diag(*source.preparation)*R*source.exits*C.T
            return dict(deadline=J,exits=V,channel=F,law=J.row_join(V),occupation=R,
                founder_hours=(source.preparation.T*R*sp.ones(m,1))[0],division_yield=(source.preparation.T*R*source.division)[0],evidence='exact rational')
        H=np.array(source.killed,float);S=expm(self.daughter_delay*H)
        F=np.vstack([np.array(E,float)@S.T,(1-S.sum(axis=1))[None,:]])
        C=np.zeros((1+len(F)**2,1+m*m));C[0,0]=1;C[1:,1:]=np.kron(F,F)
        V=np.array(E*sp.diag(*source.preparation)*R*source.exits,float)@C.T
        return dict(deadline=np.array(J,float),exits=V,channel=F,law=np.column_stack([np.array(J,float),V]),occupation=np.array(R,float),
            founder_hours=float((source.preparation.T*R*sp.ones(m,1))[0]),division_yield=float((source.preparation.T*R*source.division)[0]),evidence='numerical delayed channel')

    def recover(self,observed,marker):
        """Population-law inverse, not a constrained finite-data estimator.

        Immediate inversion is exact. Delayed inversion is numerical and returns
        its reconstruction residual; it does not clip negative fitted rates.
        A zero division rate identifies only zero exit mass, never an unused K.
        """
        if self.daughter_delay:return self._recover_delayed(observed,marker)
        O=matrix(observed);E=marker.emissions;L=marker.left_inverse();p=E.rows;m=E.cols
        if O.shape!=(p,p+1+p*p) or any(v<0 for v in O) or sum(O)!=1:raise ValueError('Invalid unconditioned immediate record law.')
        pi=L*O*sp.ones(O.cols,1)
        if any(v<=0 for v in pi):raise ValueError('Preparation cannot be inverted on every state.')
        J=O[:,:p];V=O[:,p:];A=sp.diag(*[1/x for x in pi])*L*J*L.T
        H=self.rate*(sp.eye(m)-A.inv());LC=sp.diag(sp.ones(1,1),sp.kronecker_product(L,L))
        D=sp.diag(*[1/x for x in pi])*L*V*LC.T;B=self.rate*A.inv()*D
        q=H.copy()
        for i in range(m):q[i,i]=0
        d=B[:,0];b=B[:,1:]*sp.ones(m*m,1)
        if any(v<0 for v in q) or any(v<0 for v in B) or H*sp.ones(m,1)+B*sp.ones(B.cols,1)!=sp.zeros(m,1):raise ValueError('Record law is incompatible with a nonnegative branching source.')
        K=[None if b[i]==0 else [B[i,j+1]/b[i] for j in range(m*m)] for i in range(m)]
        return dict(preparation=pi,switching=q,killed=H,division=b,death=d,sisters=K,exit_mass=B,deadline_kernel=A,evidence='exact population inverse')

    def _recover_delayed(self,observed,marker):
        E=np.array(marker.emissions,float);L=np.array(marker.left_inverse(),float);p,m=E.shape;O=np.array(observed,float)
        if O.shape!=(p,p+1+(p+1)**2) or O.min()<-1e-12 or abs(O.sum()-1)>1e-10:raise ValueError('Invalid delayed record law.')
        pi=L@O.sum(axis=1)
        if pi.min()<=0:raise ValueError('Preparation has missing support.')
        A=np.diag(1/pi)@L@O[:,:p]@L.T;H=float(self.rate)*(np.eye(m)-np.linalg.inv(A))
        S=expm(self.daughter_delay*H);F=np.vstack([E@S.T,(1-S.sum(axis=1))[None,:]])
        if np.linalg.matrix_rank(F)<m:raise ValueError('Numerically unresolved delayed channel rank.')
        LF=np.linalg.pinv(F);LC=np.zeros((1+m*m,1+(p+1)**2));LC[0,0]=1;LC[1:,1:]=np.kron(LF,LF)
        D=np.diag(1/pi)@L@O[:,p:]@LC.T;B=float(self.rate)*np.linalg.solve(A,D)
        b=B[:,1:].sum(axis=1);q=H.copy();np.fill_diagonal(q,0)
        if min(q.min(),B.min())<-1e-8:raise ValueError('Recovered source has negative rates; no clipping applied.')
        return dict(preparation=pi,killed=H,switching=q,division=b,death=B[:,0],sisters=[None if abs(b[i])<1e-12 else B[i,1:]/b[i] for i in range(m)],exit_mass=B,
            channel_min_singular_value=float(np.linalg.svd(F,compute_uv=False)[-1]),row_balance_error=float(np.max(np.abs(H.sum(axis=1)+B.sum(axis=1)))),evidence='numerical population inverse; ill-conditioning is possible')


def known_mixture_calibration(marker_means,compositions):
    """Exact E=Y P^+; composition columns must be independently established."""
    Y=matrix(marker_means);P=matrix(compositions)
    if P.rank()!=P.rows or any(v<0 for v in P) or any(sum(P[:,j])!=1 for j in range(P.cols)):raise ValueError('Full row rank probability compositions required.')
    return MarkerChannel(Y*P.T*(P*P.T).inv())
