"""Finite weighted populations, sharp support bounds and measurement contracts."""
from dataclasses import dataclass
from fractions import Fraction as F
from math import lcm
import sympy as sp
from scipy.optimize import linprog
import numpy as np
from enzyme import rational, CapacityFamily


@dataclass(frozen=True)
class Population:
    capacities: tuple
    weights: tuple

    def __post_init__(self):
        object.__setattr__(self,'capacities',tuple(map(rational,self.capacities)))
        object.__setattr__(self,'weights',tuple(map(rational,self.weights)))
        if not self.weights or len(self.weights)!=len(self.capacities) or min(self.weights)<0 or min(self.capacities)<=0 or sum(self.weights)!=1:
            raise ValueError('Positive capacities and normalized nonnegative finite weights required.')

    @property
    def mean(self):return sum(w*v for w,v in zip(self.weights,self.capacities))

    @property
    def variance(self):return sum(w*(v-self.mean)**2 for w,v in zip(self.weights,self.capacities))

    def count_realization(self):
        denominator=lcm(*(w.denominator for w in self.weights))
        return [int(w*denominator) for w in self.weights]

    def recovery_enclosure(self,family,n=128):
        decisions=[family.at(v).classify(n) for v in self.capacities]
        lower=sum(w for w,d in zip(self.weights,decisions) if d['status']=='success')
        upper=sum(w for w,d in zip(self.weights,decisions) if d['status']!='failure')
        return dict(lower=lower,upper=upper,decisions=decisions)

    def pooled_rate(self,family,inputs):
        # Literal common-clamped rate at each capacity; only V may vary.
        return sum(w*family.at(v).parameters.rate(*inputs) for w,v in zip(self.weights,self.capacities))


@dataclass(frozen=True)
class TwoBandClass:
    low_min: F = F(5,8)
    low_max: F = F(9,14)
    high_min: F = F(1)
    high_max: F = F(11,8)
    mean: F = F(1)

    def __post_init__(self):
        for n in self.__dataclass_fields__:object.__setattr__(self,n,rational(getattr(self,n)))
        if not 0<self.low_min<=self.low_max<self.high_min<=self.mean<self.high_max:raise ValueError('Require separated bands and mean in the high band below its upper endpoint.')

    def classify_bands(self,family):
        low=family.at(self.low_max);high=family.at(self.high_min)
        inward=family.at(self.low_min).drift(0)
        margin=high.drift(high.target)
        time=(high.target-high.initial)/margin if margin>0 else None
        valid=inward>=0 and low.drift(low.target)<=0 and time is not None and time<=high.deadline
        return dict(status='certified' if valid else 'not certified',inward_lower=inward,
                    low_target_drift_upper=low.drift(low.target),high_drift_lower=margin,high_arrival_upper=time)

    def bounds(self,family):
        contract=self.classify_bands(family)
        if contract['status']!='certified':return dict(status='not certified',interval=None,contract=contract)
        return dict(status='exact identified set',interval=((self.mean-self.low_max)/(self.high_max-self.low_max),F(1)),
                    attained=(True,True),contract=contract)

    def witness(self,p):
        p=rational(p);b,d,m=self.low_max,self.high_max,self.mean
        lower=(m-b)/(d-b)
        if not lower<=p<=1:raise ValueError('Requested recovery fraction is outside this support/mean construction.')
        fail=1-p;compensation=(m-b)*fail/(d-m)
        return Population((b,d,m),(fail,compensation,1-fail-compensation))


@dataclass(frozen=True)
class ReadoutCalibration:
    sensitivity: tuple = (F(9,10),F(1))
    false_positive: tuple = (F(0),F(1,50))
    observed: tuple = (F(17,25),F(18,25))

    def __post_init__(self):
        for name in self.__dataclass_fields__:
            v=tuple(map(rational,getattr(self,name)));object.__setattr__(self,name,v)
            if len(v)!=2 or not 0<=v[0]<=v[1]<=1:raise ValueError('Probability intervals must be ordered within [0,1].')
        if self.sensitivity[0]<=self.false_positive[1]:raise ValueError('This interval formula requires separated sensitivity and false-positive bands.')

    def interval(self,prior=(F(0),F(1))):
        sL,sU=self.sensitivity;fL,fU=self.false_positive;rL,rU=self.observed
        L=max(rational(prior[0]),F(0),(rL-fU)/(sU-fU));U=min(rational(prior[1]),F(1),(rU-fL)/(sL-fL))
        return None if L>U else (L,U)

    def witness(self,p):
        p=rational(p);interval=self.interval()
        if interval is None or not interval[0]<=p<=interval[1]:raise ValueError('Recovery fraction cannot produce the readout band.')
        sL,sU=self.sensitivity;fL,fU=self.false_positive;rL,rU=self.observed
        low=sL*p+fL*(1-p);high=sU*p+fU*(1-p);r=max(rL,low)
        t=(r-low)/(high-low) if high>low else F(0)
        s=sL+t*(sU-sL);f=fL+t*(fU-fL)
        assert s*p+f*(1-p)==r and r<=rU
        return dict(sensitivity=s,false_positive=f,observed=r,
                    successful_positive=s*p,successful_negative=(1-s)*p,
                    failing_positive=f*(1-p),failing_negative=(1-f)*(1-p))


def minimax(interval):
    if interval is None:return None
    L,U=interval;return dict(midpoint=(L+U)/2,radius=(U-L)/2)


def convert_weights(interval,ratio):
    ratio=rational(ratio)
    if ratio<1:raise ValueError('Positive cell-weight ratio must be at least one.')
    if interval is None:return None
    L,U=interval
    return L/(ratio*(1-L)+L),ratio*U/(1-U+ratio*U)


def error_normalization_frontier(ratio):
    ratio=rational(ratio)
    if ratio<1:raise ValueError('Weight ratio must be at least one.')
    epsilon=(33-32*ratio)/(50*(1+2*ratio))
    return dict(status='feasible' if epsilon>=0 else 'no nonnegative error budget',maximum_additional_error=epsilon if epsilon>=0 else None)


def connected_limits(cutoff,a=F(5,8),d=F(11,8),mean=F(1)):
    """Sharp limits when success means V>=cutoff, including boundary cases.

    Accepts exact cutoff inputs; a numerical root gives numerical endpoints.
    """
    a,d,mean=map(rational,[a,d,mean]);c=cutoff
    if not a<mean<d:raise ValueError('Interior mean required.')
    if c>d:return dict(lower=0,upper=0,lower_attained=True,upper_attained=True)
    if c<=a:return dict(lower=1,upper=1,lower_attained=True,upper_attained=True)
    if c==d:return dict(lower=0,upper=(mean-a)/(d-a),lower_attained=True,upper_attained=True)
    return dict(lower=max(0,(mean-c)/(d-c)),upper=min(1,(mean-a)/(c-a)),lower_attained=c>mean,upper_attained=True)


def offband_infimum(eta,cutoff,b=F(9,14),d=F(11,8)):
    eta,b,d=map(rational,[eta,b,d]);c=cutoff
    if not 0<=eta<=1 or not b<c<1<d:raise ValueError('Require eta in [0,1] and the reference cutoff ordering.')
    return dict(infimum=max((1-c)/(d-c),(1-b-(c-b)*eta)/(d-b)),attained=eta==0,
                transition=(d-1)/(d-c))


def offband_witness(eta,failing_capacity):
    """Mean-one construction; caller must certify failing_capacity misses.

    The code limits off-band mass so all weights stay nonnegative, including
    exactly at the transition where the manuscript's first construction needs
    to be replaced by its two-atom construction.
    """
    eta,x=map(rational,[eta,failing_capacity]);b,d=F(9,14),F(11,8)
    if not 0<=eta<=1 or not b<x<1:raise ValueError('Invalid off-band fraction or capacity.')
    off=min(eta,(d-1)/(d-x));success=(1-b-off*(x-b))/(d-b)
    return Population((b,x,d),(1-off-success,off,success))


def resolved_capacity_bounds(population,measurements,error,cutoff_interval):
    z=tuple(map(rational,measurements));error=rational(error);cL,cU=map(rational,cutoff_interval)
    if len(z)!=len(population.weights) or error<0 or cL>cU:raise ValueError('Invalid measurement contract.')
    lo=sum(w for w,x in zip(population.weights,z) if x>=cU+error)
    hi=sum(w for w,x in zip(population.weights,z) if x>=cL-error)
    return dict(lower=lo,upper=hi,unresolved_mass=hi-lo,ambiguity_band=(cL-error,cU+error))


def row_span_certificate(rows,target):
    A=sp.Matrix([[sp.Rational(x) for x in r] for r in rows]);h=sp.Matrix([sp.Rational(x) for x in target])
    if A.cols!=h.rows or not any(all(A[i,j]==1 for j in range(A.cols)) for i in range(A.rows)):
        raise ValueError('Matching feature rows, including the constant row, required.')
    if A.rank()==A.col_join(h.T).rank():
        solution=A.T.gauss_jordan_solve(h)[0]
        solution=solution.subs({s:0 for s in solution.free_symbols})
        assert A.T*solution==h
        return dict(identified_for_all_weights=True,multipliers=list(solution))
    z=next(z for z in A.nullspace() if (h.T*z)[0]!=0)
    positive=[max(v,0) for v in z];negative=[max(-v,0) for v in z];mass=sum(positive)
    w1=[v/mass for v in positive];w2=[v/mass for v in negative]
    assert A*sp.Matrix(w1)==A*sp.Matrix(w2)
    return dict(identified_for_all_weights=False,null_direction=list(z),witness_weights=[w1,w2],target_values=[(h.T*sp.Matrix(w))[0] for w in [w1,w2]])


def finite_support_bounds(rows,values,target):
    """Numerical LP for an edited finite support; never labeled an exact bound."""
    A=np.array(rows,float);b=np.array(values,float);h=np.array(target,float)
    lo=linprog(h,A_eq=A,b_eq=b,bounds=(0,None),method='highs')
    hi=linprog(-h,A_eq=A,b_eq=b,bounds=(0,None),method='highs')
    if not lo.success or not hi.success:return dict(status='not certified',reason='Numerical feasible-set solve failed or found incompatibility.')
    return dict(status='numerical bounds',lower=float(lo.fun),upper=float(-hi.fun),lower_weights=lo.x.tolist(),upper_weights=hi.x.tolist())
