import unittest
from fractions import Fraction as F
import numpy as np
from model import *
from certificates import *


class ScientificTests(unittest.TestCase):
    def test_original_specimen_and_shared_state(self):
        a,b=SpecimenBudget().fractions();self.assertEqual((a,b),(F(9,20),F(9,20)))
        law=SourceLaw(((F(1,2),F(0),F(0)),(F(1,2),F(1),F(1))))
        self.assertEqual(law.negative(2,a,b),F(101,200))
        self.assertNotEqual(law.negative(2,a,b),(1-a/2-b/2)**2)
        for _ in range(10):self.assertEqual(sum(law.sample_specimen(6,a,b,np.random.default_rng(1))['path_counts']),6)
        with self.assertRaises(ValueError):law.negative(2,F(3,4),F(3,4))

    def test_sharp_gate_and_floor(self):
        h=F(11,20)**6;g,z=gate_bound(h,9)
        self.assertLess(g,F(1,20));self.assertGreater(gate_bound(h,8)[0],F(1,20))
        self.assertEqual(SourceLaw.corners(z,0,0).joint_error(6,F(9,20),F(9,20),9),g)
        self.assertEqual(minimum_pairs(F(11,20)**5,F(1,20))['status'],'infeasible')
        self.assertEqual(minimum_pairs(h,F(1,20),2)['status'],'budget-exhausted')
        self.assertEqual(gate_bound(F(1,4),3)[0],F(1,4))

    def test_reporting_and_moments(self):
        p=ExclusionPolicy(F(9,20),F(9,20),6,9)
        self.assertTrue(p.report(['11']*9,0)['issued'])
        for records,count in [(['10']*8,0),(['10']*8+['missing'],0),(['00']*9,0),(['10']*9,1)]:self.assertFalse(p.report(records,count)['issued'])
        self.assertEqual(moment_bound(2,F(9,20),F(9,20),F(4,5),F(4,5),F(16,25)),F(179,1250))
        law=SourceLaw(((F(1),F(2,5),F(4,5)),));p1,p2,q=law.moments()
        self.assertLessEqual(law.negative(6,F(9,20),F(9,20)),moment_bound(6,F(9,20),F(9,20),p1,p2,q))

    def test_batch_exact_decisions(self):
        h=F(11,20)**6
        self.assertTrue(batch_certificate(h,9,1,F(1,20))['certified'])
        self.assertFalse(batch_certificate(h,8,1,F(1,20))['certified'])
        self.assertFalse(batch_certificate(h,100,2,F(1,20))['reachable'])
        for K,T,m in [(8,6,96),(10,20,330),(12,66,1070)]:
            h=F(11,20)**K
            self.assertTrue(batch_certificate(h,m,T,F(1,20))['certified'])
            self.assertFalse(batch_certificate(h,m-1,T,F(1,20))['certified'])

    def test_observation_and_specificity(self):
        h=F(11,20)**6
        self.assertLess(observation_bound(h,9,F(1,10000),F(1,10000)),F(1,20))
        self.assertGreater(observation_bound(h,9,F(1,1000),F(1,1000)),F(1,20))
        self.assertLess(specificity_bound(h,9,F(1,300)),F(1,20))
        self.assertGreater(specificity_bound(h,9,F(1,250)),F(1,20))
        self.assertLessEqual(F(299,300)**898,F(1,20));self.assertGreater(F(299,300)**897,F(1,20))

    def test_poisson_chord_and_counterexample(self):
        ref=PoissonReference(F(5,2),F(9,20),6)
        self.assertEqual(ref.chord_status(),'certified')
        self.assertLess(ref.chord_bound(7)[1],F(1,20));self.assertGreater(ref.chord_bound(6)[0],F(1,20))
        ref=PoissonReference(3,F(9,20),6);self.assertEqual(ref.chord_status(),'invalid');self.assertIsNone(ref.chord_bound(9))
        lo,hi=ref.source_error([(F(1,5),F(23,100)),(F(4,5),F(2))],9)
        invalid=(1-exp_negative(6)[0])**9*gate_bound(F(1,10)**6,9)[0]
        self.assertGreater(lo,invalid)
        env=ref.envelope_numeric(9);self.assertGreaterEqual(env['bound'],float(lo)-1e-12)
        self.assertGreater(PoissonReference(100,F(9,20),6).source_error([(F(1),F(1,10))],9)[0],F(3,4))

    def test_source_constraint_certificate(self):
        h,j=F(11,20)**6,F(1,10)**6
        self.assertEqual(constrained_certificate(h,j,8,F(1,20))['status'],'certified')
        self.assertEqual(constrained_certificate(h,j,7,F(1,20))['status'],'refuted')
        self.assertEqual(constrained_certificate(h,j,8,F(1,20),1)['status'],'budget-exhausted')
        # Nonnegative covariance alone still allows the unrestricted adversary.
        g,z=gate_bound(h,9);law=SourceLaw.corners(z,0,0);p1,p2,q=law.moments()
        self.assertEqual(q,p1*p2);self.assertEqual(law.joint_error(6,F(9,20),F(9,20),9),g)


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