"""Finite-capacity reserve: exact path-risk certificates and numerical chains."""
from dataclasses import dataclass
from fractions import Fraction as F
from math import ceil,floor,factorial,log,exp,comb
import numpy as np
from scipy.integrate import solve_ivp
from scipy.sparse import diags,lil_matrix
import mpmath as mp


@dataclass(frozen=True)
class Reserve:
    capacity:int
    threshold:int
    renewal:F
    mortality_ceiling:F
    def __post_init__(self):
        if not isinstance(self.capacity,int) or not isinstance(self.threshold,int) or not 1<=self.threshold<=self.capacity:raise ValueError('Integer threshold in [1,capacity] required.')
        for name in ['renewal','mortality_ceiling']:
            object.__setattr__(self,name,F(getattr(self,name)))
            if getattr(self,name)<=0:raise ValueError('Positive comparison rates required.')

    def anchored(self,anchor,horizon):return AnchoredCertificate(self,anchor,F(horizon))

    def capacity_envelope(self,theta,horizon):
        theta=F(theta);xstar=1-self.mortality_ceiling/self.renewal
        if not 0<theta<1 or xstar<=theta:return dict(status='outside_capacity_law_regime',equilibrium_fraction=xstar)
        h=ceil(self.capacity*theta);M=floor(self.capacity*xstar)
        if h>M:return dict(status='rounding_leaves_no_admissible_anchor')
        f=log(float(self.renewal*(1-theta)/self.mortality_ceiling));I=float(1-theta)*f-float(1-theta)+float(self.mortality_ceiling/self.renewal)
        bound=min(1.,self.capacity*float(self.mortality_ceiling)*float(horizon)*exp(-self.capacity*I+2*f))
        return dict(status='numerical_display_of_proved_formula',threshold=h,anchor=M,equilibrium_fraction=xstar,action=I,rounding_correction=2*f,bound=bound)

    def high_renewal_coefficient(self,horizon):
        d=self.capacity-self.threshold
        if d<1:raise ValueError('Zero buffer is a different first-death problem.')
        return (self.capacity*self.mortality_ceiling)**(d+1)*F(horizon)/factorial(d)


class AnchoredCertificate:
    def __init__(self,reserve,anchor,horizon):
        self.reserve=reserve;self.anchor=anchor;self.horizon=F(horizon)
        if not isinstance(anchor,int) or not reserve.threshold<=anchor<=reserve.capacity or horizon<=0:raise ValueError('Anchor must be an integer between threshold and capacity; horizon positive.')
        self.products={reserve.threshold-1:F(1)}
        for h in range(reserve.threshold,anchor):
            self.products[h]=self.products[h-1]*reserve.capacity*reserve.mortality_ceiling/(reserve.renewal*(reserve.capacity-h))
        self.denominator=sum(self.products.values(),F(0));self.excursion=self.products[anchor-1]/self.denominator
        self.departures=anchor*reserve.mortality_ceiling*self.horizon

    def initial_failure_before_return(self,initial):
        if not isinstance(initial,int) or not 0<=initial<=self.reserve.capacity:raise ValueError('Initial count outside capacity.')
        if initial<self.reserve.threshold:return F(1)
        if initial>=self.anchor:return F(0)
        return sum((p for j,p in self.products.items() if j>=initial),F(0))/self.denominator

    def evaluate(self,initial):
        first=self.initial_failure_before_return(initial);repeated=self.departures*self.excursion
        return dict(initial=initial,anchor=self.anchor,initial_term=first,repeated_excursion_term=repeated,
            sharp_bound=min(F(1),first+repeated),simple_product_bound=min(F(1),first+self.departures*self.products[self.anchor-1]),
            denominator=self.denominator,scope='All-time failure H<threshold, including crossings followed by recovery; sufficient comparison bound, not the exact risk.')

    def random_preparation(self,probabilities):
        if any(p<0 for p in probabilities.values()) or sum(probabilities.values())!=1:raise ValueError('Exact probability distribution required.')
        # Averaging capped conditional bounds is sound and at least as strong as
        # adding a common excursion bound after averaging the initial term.
        return sum((F(p)*self.evaluate(h)['sharp_bound'] for h,p in probabilities.items()),F(0))


class ReserveChain:
    """Forward equations with an explicit absorbing failure counter.

    The full chain gives terminal shortfall, the killed chain gives any-time
    threshold loss. Returning above threshold cannot remove absorbed mass.
    """
    def __init__(self,reserve,mortality=None):
        self.reserve=reserve;self.mortality=mortality or (lambda t:float(reserve.mortality_ceiling))

    def solve(self,initial,horizon,samples=241,killed=True):
        R=self.reserve;K=R.capacity;low=R.threshold if killed else 0
        if not low<=initial<=K:raise ValueError('Initial count must be in the modeled safe state set.')
        states=np.arange(low,K+1);n=len(states);birth=float(R.renewal)*states*(1-states/K)
        def matrix(t):
            death=float(self.mortality(t))*states
            if min(death)<0:raise ValueError('Negative death rate.')
            Q=diags([birth[:-1],-(birth+death),death[1:]],[ -1,0,1],shape=(n,n),format='lil')
            if killed:
                A=lil_matrix((n+1,n+1));A[:n,:n]=Q;A[n,0]=death[0];return A.tocsc()
            return Q.tocsc()
        dim=n+int(killed);y=np.zeros(dim);y[initial-low]=1
        atol=np.full(dim,1e-12)
        if killed:atol[-1]=1e-24
        sol=solve_ivp(lambda t,p:matrix(t)@p,(0,float(horizon)),y,method='BDF',jac=lambda t,p:matrix(t),rtol=2e-9,atol=atol,t_eval=np.linspace(0,float(horizon),samples))
        if not sol.success:raise RuntimeError(sol.message)
        risk=sol.y[-1] if killed else sol.y[:R.threshold].sum(axis=0)
        return sol.t,risk,dict(evidence='Numerical forward equation, not an interval enclosure.',mass_residual=float(np.max(abs(sol.y.sum(axis=0)-1))),minimum_numerical_mass=float(sol.y.min()))

    def high_precision_constant_risk(self,initial,horizon,digits=60):
        R=self.reserve;states=list(range(R.threshold,R.capacity+1));n=len(states)
        if initial not in states:raise ValueError('Safe initial state required.')
        # Direct failed-state mass avoids cancellation of 1 - total survivor mass.
        with mp.workdps(digits):
            Q=mp.zeros(n+1);r=mp.mpf(R.renewal.numerator)/R.renewal.denominator;m=mp.mpf(R.mortality_ceiling.numerator)/R.mortality_ceiling.denominator
            for i,h in enumerate(states):
                birth=r*h*(1-mp.mpf(h)/R.capacity);death=m*h;Q[i,i]=-birth-death
                if i+1<n:Q[i,i+1]=birth
                if i:Q[i,i-1]=death
                else:Q[i,n]=death
            T=F(horizon);probability=mp.expm(Q*(mp.mpf(T.numerator)/T.denominator))[initial-R.threshold,n]
            return mp.nstr(probability,digits-5)


def latency_floor(capacity,threshold,single_survival):
    p=F(single_survival)
    if not 0<=p<=1 or not 1<=threshold<=capacity:raise ValueError('Invalid latency comparison.')
    return sum((F(comb(capacity,k))*p**k*(1-p)**(capacity-k) for k in range(threshold)),F(0))


def fluid_limit_diagnostic(renewal,mortality,theta,horizon,capacity):
    r,m,th,T=map(float,(renewal,mortality,theta,horizon))
    if not r>m>0 or not 1-m/r<th<1:return dict(status='outside_subcritical_reserve_regime')
    xstar=1-m/r;cross=log(m*th/(r*th-r+m))/(r-m);xT=xstar/(1-m/r*exp(-(r-m)*T))
    if T<=cross:return dict(status='deadline_before_fluid_crossing',crossing_time=cross)
    eps=(th-xT)/2;lower=max(0.,1-4*(r/4+m)*T*exp(2*(r+m)*T)/(capacity*eps**2))
    return dict(status='fluid_failure_regime',crossing_time=cross,equilibrium_fraction=xstar,terminal_fraction=xT,finite_K_lower_display=lower,
        scope='Constant source started full; a numerical display of the proved conservative bound. Not an exclusion of every intervention policy.')
