import math
import unittest
from fractions import Fraction as F
import numpy as np
import sympy as sp
from models import FounderPreparation,TwoTypeBranching,BirthDeathEnvelope
from prediction import MenuRule,DetectorFeasibleRule,ObservationBudget,InheritedRateClass,dependence_decision,scalar_finite_certificate,binomial_cdf,hoeffding_width,unresolved_coverage


class ScientificChecks(unittest.TestCase):
    def test_preparation_source_convolutions_and_boundary(self):
        I,W=FounderPreparation(F(0)),FounderPreparation(F(1))
        for k in range(2,25):
            for source in [I,W]:self.assertEqual(sum(source.detected_coefficients(k)),source.cdf_two(k))
        for z in [F(0),F(1,3),F(1)]:self.assertEqual(I.pgf(z,founders=1),W.pgf(z,founders=1))
        self.assertLess(I.cdf_two(5),F(19,20));self.assertGreater(I.cdf_two(6),F(19,20))
        self.assertLess(W.cdf_two(6),F(19,20));self.assertGreater(W.cdf_two(7),F(19,20))
        boundary=FounderPreparation(F(3,5))
        self.assertEqual(boundary.cdf_two(6),F(19,20))
        coeff=boundary.detected_coefficients(4);self.assertEqual(coeff[3]+coeff[4],F(1,4))
        for k in [2,6,10]:self.assertAlmostEqual(I.many_cdf_numeric(k,2),float(I.cdf_two(k)),places=14)

    def test_integer_calibration_and_ordered_error_account(self):
        for n,k,p in [(20,7,F(5,8)),(33,10,F(7,32))]:
            direct=sum((F(math.comb(n,j))*p**j*(1-p)**(n-j) for j in range(k+1)),F(0))
            self.assertEqual(binomial_cdf(n,k,p),direct)
        for event,n,threshold in [('at_most_two',1024,607),('at_most_two',1280,759),('at_most_two',4000,2374),('three_or_four',576,149),('three_or_four',640,165)]:
            rule=MenuRule(event);self.assertEqual(rule.threshold(n),threshold)
            r=rule.exact_errors(n);target=F(19,20) if n in [576,1024] else F(245,256)
            self.assertGreaterEqual(r['uniform_coverage'],target)
            # Adversarial overlap can place all dangerous decisions on future count seven.
            worstW=F(31,32)-min(r['W_error'],F(3,128))
            self.assertEqual(r['uniform_coverage'],min(F(245,256),worstW))
        self.assertEqual(MenuRule().choose(164,640,True)['cutoff'],7)
        self.assertEqual(MenuRule().choose(165,640,True)['cutoff'],6)
        self.assertEqual(MenuRule().choose(640,640)['cutoff'],7)
        self.assertEqual(dependence_decision(640,640)['cutoff'],6)
        self.assertEqual(dependence_decision(0,640)['cutoff'],7)

    def test_detection_gaps_factorial_moments_and_erasure(self):
        I,W=FounderPreparation(),FounderPreparation(F(1))
        for d in [F(0),F(1,4),F(1,2),F(3,4),F(1)]:
            self.assertEqual(sum(W.detected_coefficients(2,d))-sum(I.detected_coefficients(2,d)),d**3*(2*d-1)/(1+d)**4)
            for z in [F(1,3),F(1,2)]:
                u=d*(1-z);self.assertEqual(W.pgf(z,detection=d)-I.pgf(z,detection=d),u*u*(1-u)**2/(4*(1+u)**2))
        z,d=sp.symbols('z d');w=1-d+d*z;g0=w;g1=w/(2-w)
        for g,want in [((g0+g1)**2/4,F(17,2)),((g0*g0+g1*g1)/2,F(9))]:
            self.assertEqual(sp.simplify(sp.diff(g,z).subs(z,1)-3*d),0)
            self.assertEqual(sp.simplify(sp.diff(g,z,2).subs(z,1)-want*d*d),0)
        self.assertEqual(I.detected_coefficients(5,F(0)),W.detected_coefficients(5,F(0)))
        self.assertEqual(W.pgf(F(1,2),detection=F(1,2))-I.pgf(F(1,2),detection=F(1,2)),F(9,1600))

    def test_detector_feasible_rule_and_outward_width(self):
        width=hoeffding_width(2000,F(3,512))
        self.assertGreaterEqual(float(width)+1e-17,math.sqrt(math.log(1024/3)/4000))
        rule=DetectorFeasibleRule()
        no_information=rule.evaluate([0]*10)
        self.assertEqual(no_information['cutoff'],7);self.assertEqual(no_information['retained_models'],['I','W'])
        impossible=rule.evaluate([100]*2000)
        self.assertIsNone(impossible['detector_interval']);self.assertEqual(impossible['cutoff'],7)
        self.assertEqual(no_information['uniform_coverage'],F(245,256))
        with self.assertRaises(ValueError):rule.evaluate([None,0])

    def test_scalar_forward_comparison_and_rational_budget(self):
        finite=scalar_finite_certificate()
        self.assertEqual(finite['rows'][-1],840711)
        self.assertEqual(finite['lower'],F(3134312565151423153,3252877056000000000))
        self.assertGreater(finite['lower'],F(963,1000))
        budget=ObservationBudget();self.assertEqual(budget.recorded(F(963,1000)),F(952037,1000000))
        for birth,death,k,want in [(.105,.05,4,F(95034241414,10**11)),(.1,0,5,F(95670013878,10**11))]:
            envelope=BirthDeathEnvelope(birth,death);lower=envelope.exact_supercritical_lower(k)
            self.assertGreater(budget.recorded(lower),want)
            p,overflow=TwoTypeBranching(birth,death,.3,birth,death,.2).killed_distribution(7,50)
            self.assertLess(abs(p[:k+1].sum()-envelope.cdf(k)),2e-9+overflow)
            self.assertLessEqual(float(lower),envelope.cdf(k))
        for birth,death in [(0,.1),(.1,.1),(.05,.1)]:
            p,overflow=TwoTypeBranching(birth,death,0,birth,death,0).killed_distribution(7,40)
            self.assertAlmostEqual(p[:5].sum(),BirthDeathEnvelope(birth,death).cdf(4),places=8)

    def test_inherited_success_histories_and_minimality(self):
        cls=InheritedRateClass();budget=ObservationBudget();a=cls.source_lower(3);b=cls.source_lower(4)
        self.assertEqual(a['lower'],F(176701212731773153,182749885209000000))
        B3=budget.recorded(a['lower']);B4=budget.recorded(b['lower'])
        self.assertEqual(B3,F(6470259728405606661,6768514267000000000))
        self.assertEqual(B4,F(16473049504746783406959,17007584224404250000000))
        self.assertLess(F(99,100)*B3,F(19,20));self.assertGreater(F(99,100)*B4,F(19,20))
        self.assertGreater(cls.reference_endpoint_two_failure_lower(),F(1,20))
        self.assertTrue(cls.contains((.0001,.2999,.0011,.1001,.0001,.00002)))
        self.assertFalse(cls.contains((.0005,.3005,.0015,.1,0,.00001)))
        p,overflow=TwoTypeBranching().killed_distribution(math.log(2)/.10001)
        self.assertGreater(p[:4].sum(),float(a['lower']));self.assertLess(p[:3].sum()+overflow,.95)

    def test_sampling_and_missing_outcome_accounting(self):
        source=FounderPreparation(F(1));latent,observed=source.sample(100,2,.4,12)
        self.assertTrue(np.all(latent>=2));self.assertTrue(np.all(observed<=latent))
        np.testing.assert_array_equal(observed,source.sample(100,2,.4,12)[1])
        self.assertEqual(unresolved_coverage(51,24,75),(F(17,25),F(1)))
        # False objects can concentrate on successful outcomes; subtract the whole budget.
        self.assertEqual(ObservationBudget(F(0),F(1,100)).recorded(F(95,100)),F(94,100))
        for call in [lambda:FounderPreparation(F(2)),lambda:TwoTypeBranching(-1),lambda:MenuRule('posthoc'),lambda:ObservationBudget(F(-1))]:
            with self.assertRaises(ValueError):call()


if __name__=='__main__':unittest.main()
