import unittest
from fractions import Fraction as F
from decimal import Decimal, localcontext
import numpy as np
from scipy.stats import poisson
from example import (BirthSource, PoissonLoading, FiniteChain, ExactWitness,
                     ErrorLimits, LimitLaw, exp_enclosure, witnesses, large_threshold_bound)


class ProbabilityTests(unittest.TestCase):
    def test_generator_loading_and_zero_time(self):
        chain=FiniteChain(BirthSource(capacity=8))
        np.testing.assert_allclose(chain.generator.sum(axis=1),0,atol=1e-15)
        np.testing.assert_array_equal(chain.generator[-1],np.zeros(6))
        self.assertAlmostEqual(chain.errors(0)[0],0)
        self.assertAlmostEqual(chain.errors(0)[1],poisson.cdf(4,4))
        self.assertGreater(1-chain.errors(0)[1],.37)

    def test_exact_exponential_against_independent_decimal(self):
        with localcontext() as ctx:
            ctx.prec=90
            for x in (F(0),F(-4),F(-1203,125),F(2,3)):
                interval=exp_enclosure(x)
                value=(Decimal(x.numerator)/Decimal(x.denominator)).exp()
                lo=Decimal(interval.lo.numerator)/Decimal(interval.lo.denominator)
                hi=Decimal(interval.hi.numerator)/Decimal(interval.hi.denominator)
                self.assertLessEqual(lo,value);self.assertGreaterEqual(hi,value)
                self.assertLess(float(interval.hi-interval.lo),1e-18)

    def test_exact_witness_decisions_and_grid_exclusion(self):
        rows=witnesses()['rows']
        self.assertTrue(all(row['separating'] for row in rows[:3]))
        self.assertTrue(rows[3]['feasible'])
        self.assertGreater(F(rows[4]['miss']['lower_fraction']),F(1,20))
        self.assertGreater(F(rows[5]['blank']['lower_fraction']),F(1,100))
        self.assertTrue(rows[6]['feasible'])
        self.assertLess(F(rows[7]['miss']['upper_fraction']),F(1,20))
        self.assertLess(F(rows[8]['blank']['upper_fraction']),F(1,100))

    def test_window_clock_and_rate_order(self):
        model=FiniteChain(BirthSource(capacity=8))
        window=ErrorLimits().window(model)
        self.assertEqual(window['status'],'feasible');self.assertEqual(window['grid_frames'],[])
        self.assertAlmostEqual(window['earliest'],3.4237,places=4)
        doubled=FiniteChain(model.source,rate_factors=[2]*5)
        np.testing.assert_allclose(doubled.errors(1.5),model.errors(3),atol=1e-14)
        slow=FiniteChain(model.source,rate_factors=[.98]*5)
        fast=FiniteChain(model.source,rate_factors=[1.02]*5)
        mixed=FiniteChain(model.source,rate_factors=[.98,1.02,1,.99,1.01])
        self.assertLessEqual(fast.errors(3)[1],mixed.errors(3)[1])
        self.assertLessEqual(mixed.errors(3)[1],slow.errors(3)[1])

    def test_single_birth_and_repeated_rates(self):
        chain=FiniteChain(BirthSource(threshold=1,capacity=1))
        self.assertAlmostEqual(chain.errors(3)[0],1-np.exp(-.03))
        self.assertAlmostEqual(chain.errors(3)[1],np.exp(-4-.03))
        # Source with repeated rates remains valid numerically; exact spectral
        # method refuses rather than dividing by zero.
        source=BirthSource(threshold=2,capacity=2,background=F(1),growth=F(1))
        with self.assertRaises(ValueError): ExactWitness(source)
        blank,_=FiniteChain(source).errors(2)
        self.assertAlmostEqual(blank,1-np.exp(-2)*(1+2))

    def test_limit_boundary_and_large_theorem(self):
        two=LimitLaw(headroom=2).boundary();three=LimitLaw(headroom=3).boundary()
        self.assertGreater(two['miss'],.05);self.assertLess(three['miss'],.05)
        self.assertAlmostEqual(two['miss'],.05027,places=5)
        self.assertLess(two['quadrature_error_estimate'],1e-8)
        bounds=large_threshold_bound()['approximate_bounds']
        self.assertLess(bounds['doubled_blank_upper'],.01)
        self.assertLess(bounds['doubled_miss_upper'],.05)
        self.assertGreater(bounds['saturated_blank_lower'],.01)
        self.assertGreater(bounds['saturated_miss_lower'],.05)

    def test_invalid_inputs(self):
        for f in (lambda:BirthSource(capacity=4),lambda:BirthSource(threshold=0),
                  lambda:PoissonLoading(-1),lambda:FiniteChain(BirthSource()).errors(-1),
                  lambda:large_threshold_bound(1000),lambda:LimitLaw(headroom=-1)):
            with self.assertRaises(ValueError): f()


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