import unittest
from dataclasses import replace
from fractions import Fraction as F
from example import (TwoStageType,CoverageContract,TypePopulation,Allocation,
    NegativePlanner,NegativeHistoryLaw,ExactBudgetExceeded,independent_minimum)


class ScientificChecks(unittest.TestCase):
    def test_exact_floor_profiles_and_all_attainment_regimes(self):
        for c in [F(0),F(2,5),F(19,20),F(1),F(3,2),F(2)]:
            for d in [F(0),F(1,5),F(1,2),F(1)]:
                con=CoverageContract(c,d)
                for w in [F(0),F(1,10),F(1,3),F(1,2),F(4,5),F(1)]:
                    q=con.witness(w);self.assertGreaterEqual(q.x+q.a,c);self.assertGreaterEqual(q.y+q.b,c)
                    self.assertLessEqual(abs(q.x-q.y),d)
                    self.assertEqual(w*q.x*q.y+(1-w)*q.a*q.b,con.profile(w))
                    self.assertLessEqual(con.profile(w),con.floor)
                self.assertEqual(con.profile(F(1,2)),con.floor)
        plateau=CoverageContract('3/2','1/2')
        self.assertEqual(plateau.profile(F(1,3)),F(1,2));self.assertEqual(plateau.profile(F(2,3)),F(1,2))
        self.assertLess(plateau.profile(F(1,4)),F(1,2))
        # At c=2 the interval collapses: every allocation succeeds with response one.
        self.assertEqual(CoverageContract(2,1).profile(0),1)
        for c,d,w in [('19/20','1/2','1/8'),('3/2','1/10','3/8'),('3/2','1/2','7/8')]:
            con=CoverageContract(c,d);self.assertAlmostEqual(independent_minimum(con,F(w)),float(con.profile(F(w))),places=10)

    def test_pointwise_lower_bound_independent_rational_grid(self):
        # Direct 4D sources, including slack coverage, are not constructed from the minimizer.
        grid=[F(i,4) for i in range(5)]
        import itertools
        checked=0
        for x,y,a,b in itertools.product(grid,repeat=4):
            c=min(x+a,y+b);d=abs(x-y);con=CoverageContract(c,d)
            self.assertGreaterEqual((x*y+a*b)/2,con.floor)
            for w in [F(1,4),F(3,4)]:self.assertGreaterEqual(w*x*y+(1-w)*a*b,con.profile(w))
            checked+=1
        self.assertEqual(checked,625)

    def test_rms_coverage_recording_exception_and_pooled_trap(self):
        # One quarter of the covered mass violates uniform d=.5, but RMS <=.5 is valid.
        pop=TypePopulation(((F(1,4),TwoStageType(1,0,0,1)),(F(3,4),TwoStageType(F(1,2),F(1,2),F(1,2),F(1,2)))))
        contract=CoverageContract(1,F(1,2),1,0)
        self.assertTrue(pop.verify(contract,'rms')['admitted']);self.assertFalse(pop.verify(contract,'uniform')['admitted'])
        self.assertEqual(pop.verify(contract)['mean_square'],F(1,4))
        con=CoverageContract();w=con.worst_population()
        self.assertTrue(w.verify(con)['admitted']);self.assertEqual(w.responses,(con.recorded_floor,)*2)
        trap=TypePopulation(((F(1,2),TwoStageType(1,0,0,1)),(F(1,2),TwoStageType(0,1,1,0))))
        pooled=trap.pooled_condition1();self.assertEqual(pooled['mean_activation'],pooled['mean_second_stage'])
        self.assertEqual(pooled['conditional_progeny'],0);self.assertEqual(trap.responses,(0,0))
        invalid=TypePopulation(((F(1),TwoStageType(0,0,0,0,covered=True)),))
        self.assertFalse(invalid.verify(con)['admitted'])

    def test_source_terminal_normalization_and_actual_sampling_laws(self):
        source=TwoStageType('4/5','3/4','1/2','2/3','9/10','4/5')
        p=F(1,5);pop=TypePopulation(((F(1),source),))
        for condition in [0,1]:
            terminal=source.terminals(p,condition);self.assertEqual(sum(terminal),1)
            self.assertEqual(sum(terminal[:4]),1-p*pop.responses[condition])
        balanced=Allocation(2,2)
        self.assertLess(balanced.negative_mass(p,pop),balanced.random_assignment_mass(p,pop))
        uneven=Allocation(3,1)
        self.assertNotEqual(uneven.negative_mass(p,pop),uneven.random_assignment_mass(p,pop,F(3,4)))
        # Conditional shared bad environments create a persistent floor.
        g=CoverageContract().recorded_floor;planner=NegativePlanner(g)
        self.assertGreaterEqual(planner.total_error(10000),planner.delta)

    def test_exact_planning_neighbors_and_perfect_edge(self):
        con=CoverageContract();P=NegativePlanner(con.recorded_floor);plan=P.plan()
        self.assertEqual(plan['balanced_minimum'],2300);self.assertEqual(plan['random_minimum'],2300)
        self.assertLessEqual(P.total_error(2300),F(1,20));self.assertGreater(P.total_error(2299),F(1,20))
        lo,hi=P.endpoint(2300);self.assertLess(hi,.01);self.assertTrue(lo<=.009996058<=hi or abs((lo+hi)/2-.009996058)<1e-8)
        full=NegativePlanner(CoverageContract('3/2','1/2').recorded_floor).plan()
        self.assertEqual((full['random_minimum'],full['balanced_minimum']),(749,750))
        perfect=NegativePlanner(1,theta=1,delta='1/20',alpha='1/20')
        self.assertEqual(perfect.plan()['balanced_minimum'],2);self.assertIsNone(perfect.endpoint(2))
        self.assertFalse(NegativePlanner('9/10',theta=1,delta='1/20',alpha='1/20').feasible)
        self.assertFalse(NegativePlanner(0).feasible)
        with self.assertRaises(ExactBudgetExceeded):replace(P,bit_budget=10).plan()

    def test_adaptive_full_negative_history_tree(self):
        class Policy:
            def weight(self,h):return F(1,7) if h and h[-1][1]==2 else F(5,6)
        con=CoverageContract();pop=con.worst_population();p=F(1,100)
        self.assertEqual(NegativeHistoryLaw(pop,p,Policy()).mass(4),(1-p*con.recorded_floor)**4)
        with self.assertRaises(ValueError):NegativeHistoryLaw(pop,p,Policy()).mass(20)

    def test_record_denominators_loss_and_invalid_inputs(self):
        p=NegativePlanner(CoverageContract().recorded_floor);a=Allocation(1150,1150)
        self.assertEqual(p.record(a)['status'],'conditional calculation only')
        self.assertEqual(p.record(a,contract_validated=True)['status'],'excludes target fraction')
        for pos,loss in [(1,0),(0,1)]:self.assertEqual(p.record(a,pos,loss,True)['confidence_endpoint'],1)
        self.assertEqual(p.record(Allocation(1200,1100),contract_validated=True)['status'],'unavailable')
        with self.assertRaises(ValueError):p.record(a,positives=2301)
        with self.assertRaises(ValueError):CoverageContract(c=F(21,10))
        with self.assertRaises(ValueError):TwoStageType(-1,0,0,0)
        with self.assertRaises(ValueError):Allocation(1.5,2)


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