"""Distribution, preparation, overflow, exact arithmetic and simulation checks."""
import unittest
from fractions import Fraction as Q
import numpy as np
from example import (Rates,FounderPreparation,MomentClosure,ScalarTerminalLaw,BackwardCountLaw,
    ForwardOverflowLaw,PredictionAssessment,EventSimulator,exact_witness,scenario,HORIZON_DAYS)


class ScientificChecks(unittest.TestCase):
    def test_exact_witness_and_polynomial(self):
        c=exact_witness()
        self.assertLess(c['actual_coverage_at_most_two_upper'],Q(19,20))
        self.assertGreater(c['actual_coverage_at_most_three_lower'],Q(19,20))
        self.assertGreater(c['scalar_coverage_at_most_two_lower'],Q(19,20))
        self.assertGreater(c['actual_variance_lower'],Q(4,5))
        self.assertEqual(c['minimal_one_sided_95_endpoint'],3)
        self.assertEqual(c['markov_only_endpoint'],10)

    def test_full_witness_and_founder_preparation(self):
        m,actual,scalar,rows=scenario(Rates(),5/24,HORIZON_DAYS,(1,2,100),cap=192)
        self.assertAlmostEqual(m['mean'],.5186169458,places=9)
        self.assertAlmostEqual(m['variance_actual'],1.0875939888,places=8)
        self.assertAlmostEqual(sum(actual[:3]),.9475977178,places=9)
        self.assertAlmostEqual(sum(actual[:4]),.9738271218,places=9)
        self.assertLess(m['mean_identity_error'],1e-10);self.assertLess(m['defect_identity_error'],1e-10)
        self.assertEqual(rows[0]['scalar_interval'],[0,2]);self.assertEqual(rows[0]['actual_one_sided_endpoint'],3)
        self.assertGreater(rows[1]['actual_coverage'],.95)
        self.assertEqual(rows[2]['scalar_interval'],[38,67]);self.assertAlmostEqual(rows[2]['actual_coverage'],.851145704,places=7)
        self.assertLess(abs(rows[2]['actual_mass_deficit']),1e-10)
        # Random founder type produces nonzero type covariance but zero initial count variance.
        self.assertAlmostEqual(m['trace'][0][5],0)

    def test_low_coefficients_include_grow_then_shrink(self):
        r=Rates((1.,1.),(1.,1.),(.2,.3))
        low=BackwardCountLaw(r).solve(1.,3);high=BackwardCountLaw(r).solve(1.,32)
        np.testing.assert_allclose(low,high[:,:4],atol=1e-11)
        forward=ForwardOverflowLaw(r,3).solve(1.,FounderPreparation())
        actual=(19/24)*low[0]+(5/24)*low[1]
        self.assertGreater(actual[:2].sum()-forward['pmf'][:2].sum(),1e-3)
        self.assertLess(actual[:2].sum()-forward['pmf'][:2].sum(),forward['overflow'])
        self.assertLess(forward['mass_residual'],1e-12)
        with self.assertRaises(ArithmeticError):PredictionAssessment().quantile(np.array([.1,.2]),.975)

    def test_scalar_special_cases_and_equal_growth(self):
        # Type-independent birth/death gives the same count law, despite switching.
        r=Rates((.2,.2),(.1,.1),(.3,.4));m,actual,scalar,_=scenario(r,.6,2.,cap=40)
        np.testing.assert_allclose(actual,scalar,atol=1e-11)
        self.assertAlmostEqual(m['variance_actual'],m['variance_scalar'],places=10)
        law=ScalarTerminalLaw(m['mean'],m['A']);self.assertAlmostEqual(sum(law.coefficients(40))+law.tail_above(40),1.,places=12)
        # Equal net growth with unequal turnover guarantees equal variances, not equal laws.
        r=Rates((.2,.4),(.1,.3),(.2,.1));m=MomentClosure(r,.4).solve(3.)
        self.assertAlmostEqual(m['variance_actual'],m['variance_scalar'],places=9)
        self.assertGreater(m['trace'][0][1],0)

    def test_plating_polynomials_and_source_control(self):
        one=np.array([.2,.3,.5]);p=FounderPreparation(3,.4,.2)
        base=.4*one;base[0]+=.6
        np.testing.assert_allclose(p.compose(one,6),np.convolve(np.convolve(base,base),base),atol=1e-16)
        np.testing.assert_allclose(FounderPreparation(100,0,.2).compose(one,5),[1,0,0,0,0,0])
        mean,var=p.moments(1.3,.61);self.assertAlmostEqual(mean,1.56);self.assertAlmostEqual(var,3*(.4*.61+.4*.6*1.3**2))
        theta=(.57-np.sqrt(.57**2-4*.05*.02))/.1;f=817*614/(np.pi*4500**2)
        _,_,_,rows=scenario(Rates((0,.1),(.3,0),(1.,.01)),theta,7.,(1000,),cap=128,plating=f)
        self.assertEqual(rows[0]['scalar_interval'],[3,24]);self.assertAlmostEqual(rows[0]['actual_coverage'],.9532165,places=6)
        self.assertGreater(rows[0]['variance_actual']/rows[0]['variance_scalar'],1.)

    def test_forward_crosscheck_and_zero_time(self):
        q=ForwardOverflowLaw(Rates(),32).solve(HORIZON_DAYS,FounderPreparation())
        actual=BackwardCountLaw(Rates()).solve(HORIZON_DAYS,3);mixture=19/24*actual[0]+5/24*actual[1]
        self.assertLess(abs(q['pmf'][:3].sum()-mixture[:3].sum()),1e-11)
        self.assertGreater(q['overflow'],0);self.assertLess(q['overflow'],5e-11)
        prep=FounderPreparation(5,.7,.2);q=ForwardOverflowLaw(Rates(),3).solve(0.,prep)
        from scipy.stats import binom
        np.testing.assert_allclose(q['pmf'],binom.pmf(np.arange(4),5,.7),atol=1e-14)
        self.assertAlmostEqual(q['overflow'],binom.sf(3,5,.7),places=14)

    def test_event_semantics_and_invalid_inputs(self):
        sim=EventSimulator(Rates((10,10),(0,0),(0,0)))
        result=sim.run(FounderPreparation(1,1,1),10,seed=1,event_budget=0)
        self.assertFalse(result['finished']);self.assertEqual((result['S'],result['R']),(0,1))
        result=EventSimulator(Rates((0,0),(0,0),(0,0))).run(FounderPreparation(8,1,0),2)
        self.assertTrue(result['finished']);self.assertEqual((result['S'],result['R']),(8,0))
        with self.assertRaises(ValueError):Rates((-1,0),(0,0),(0,0))
        with self.assertRaises(ValueError):FounderPreparation(1,1,2)
        with self.assertRaises(ValueError):sim.run(FounderPreparation(),-1)


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