import unittest
from example import *


class ScientificChecks(unittest.TestCase):
    def test_literal_generator_all_words_and_support_invariants(self):
        p=Parameters();chem=FiniteFoodChemistry(p);chain=CountChain(p);self.assertEqual(len(chem.reactions),24)
        for word in product((0,1),repeat=2):
            for n in chain.states:
                state=MolecularState.pure(9,word,n);aggregated={}
                for r in chem.reactions:
                    rate=r.propensity(state)
                    if rate:
                        target=r.fire(state);self.assertEqual(target.totals,(9,9));self.assertEqual(target.word,word)
                        counts=tuple(sum(target.counts[2*i:2*i+2]) for i in range(2));aggregated[counts]=aggregated.get(counts,0)+rate
                self.assertEqual(aggregated,dict(chain.transitions(n)))
        # A contaminant replacing one food molecule survives every enabled step.
        state=MolecularState((1,1,2,0,7,7))
        for r in chem.reactions:
            if r.propensity(state):self.assertGreater(r.fire(state).counts[1],0)
        # Reverse 2S -> S+F consumes no last copy; no factorial divisor.
        state=MolecularState.pure(9,(0,0),(2,1));r=chem.reactions[1];self.assertEqual(r.propensity(state),2)

    def test_exact_partition_refill_and_failed_mass(self):
        division=FairDivision(9)
        for n in range(1,10):self.assertEqual(sum(F(comb(n,k),2**n) for k in range(1,n)),1-division.empty_probability(n))
        parent=MolecularState((7,0,0,6,2,3));d=division.allocate(parent,(3,0,0,2,1,2))
        self.assertTrue(d['success']);self.assertEqual(d['total_food_added'],18);self.assertEqual(d['left'].counts,(3,0,0,2,6,7));self.assertEqual(d['right'].counts,(4,0,0,4,5,5))
        failed=division.allocate(parent,(0,0,0,2,1,2));self.assertFalse(failed['success']);self.assertEqual(failed['total_food_added'],18)
        # Enumerate a contaminated parent allocation: both pure original daughters impossible.
        contaminated=MolecularState((2,1,2,0,0,1));small=FairDivision(3)
        for left in product(*(range(n+1) for n in contaminated.counts)):
            d=small.allocate(contaminated,left)
            self.assertFalse(d['left'].admitted(3,(0,0)) and d['right'].admitted(3,(0,0)))

    def test_exact_generator_certificates(self):
        for K,s,C,q in [(9,3900,F(2523,160),F(8046297159,8112104000)),(20,9000,F(2535,131072),F(7077818258231,7077927321600))]:
            cert=DriftCertificate.compute(CountChain(Parameters(K)),s);self.assertEqual(cert.residual,C);self.assertEqual(cert.maximum_violation,0);self.assertEqual(cert.lower(20),q)
        box=DriftCertificate.operating_box();self.assertEqual(box.residual,F(279,32));self.assertEqual(box.lower(20),F(114080769,115203200))
        slow=DriftCertificate.compute(CountChain(Parameters(9,F(1),F(100))),3900);self.assertEqual(slow.lower(20),0)

    def test_stationary_balance_and_transient_interaction(self):
        for K in (3,9,20):self.assertTrue(CountChain(Parameters(K)).verify_detailed_balance())
        a=CountChain(Parameters(3,coupling=F(0)));b=CountChain(Parameters(3,coupling=F(1)))
        self.assertEqual(a.stationary(),b.stationary());self.assertGreater(np.max(abs(a.transition_matrix(.0001)-b.transition_matrix(.0001))),.01)
        self.assertNotEqual(dict(b.transitions((1,1)))[(2,1)],dict(b.transitions((1,2)))[(2,2)])

    def test_restart_lineage_family_and_matrix_accuracy(self):
        chain=CountChain();kernel=RestartKernel(chain,20);cert=DriftCertificate.compute(chain,3900);q=float(cert.lower(20));u=float(FairDivision.joint_success((9,9)))
        rows=kernel.matrix.sum(axis=1);self.assertLess(kernel.row_error,1e-12);self.assertTrue(np.all(rows>=q));self.assertTrue(np.all(rows<=u))
        stationary=sum(p*FairDivision.joint_success(n) for n,p in chain.stationary().items());self.assertLess(abs(min(rows)-float(stationary)),1e-9)
        lineage=kernel.lineage(20);family=kernel.family(3)
        self.assertTrue(np.all(lineage[-1]>=q**20-1e-9));self.assertTrue(np.all(lineage[-1]<=u**20+1e-9))
        self.assertTrue(np.all(family[-1]>=1-7*(1-q)-1e-9));self.assertTrue(np.all(family[-1]<=u**7+1e-9))
        # At T=0, K=3 admitted (1,1) cannot successfully divide. No renormalization.
        tiny=RestartKernel(CountChain(Parameters(3)),0);self.assertEqual(tiny.matrix[0].sum(),0)
        # Both children of a successful split of (2,2) are (1,1); depth two fails.
        self.assertAlmostEqual(tiny.family(1)[1,-1],.25);self.assertEqual(tiny.family(2)[2,-1],0)
        self.assertLess(np.max(abs(chain.transition_matrix(.002)-chain.transition_matrix(.001)@chain.transition_matrix(.001))),1e-10)

    def test_sharp_geometry_budget_and_decoder(self):
        regions=CompositionRegions(4);best=min(sum(abs(a-b) for a,b in zip(p,q)) for w,ps in regions.regions.items() for v,qs in regions.regions.items() if w!=v for p in ps for q in qs)
        self.assertEqual(best,F(1,2));self.assertEqual(regions.margin(),(F(1,2),F(1,4)))
        for word,ps in regions.regions.items():
            for p in ps:self.assertEqual(regions.nearest(p)['word'],word)
        self.assertTrue(regions.nearest((F(1,4),)*4)['tie'])
        larger=CompositionRegions(9);reading=(F(1,9)-F(1,20),F(1,20),F(8,9),0)
        self.assertEqual(larger.nearest(reading)['word'],(0,0)) # L1 error 1/10 < 1/9
        self.assertEqual(max(FairDivision.joint_success((r,s)) for r in range(1,17) for s in range(1,18-r)),F(32385,32768));self.assertLess(budget_ceiling(17),F(99,100));self.assertGreater(budget_ceiling(18),F(99,100))

    def test_complete_cycle_and_unfinished_status(self):
        chemistry=FiniteFoodChemistry();cycle=ScheduledCycle(chemistry);initial=MolecularState.pure(9,(0,1),(1,1))
        stopped=cycle.run(initial,np.random.default_rng(1),max_events=0);self.assertFalse(stopped['batch']['completed']);self.assertIsNone(stopped['success']);self.assertIsNone(stopped['division'])
        short=ScheduledCycle(chemistry,F(1,100)).run(initial,np.random.default_rng(1));self.assertTrue(short['batch']['completed']);self.assertEqual(short['division']['total_food_added'],18)
        self.assertEqual(short['batch']['state'].totals,(9,9));self.assertEqual(short['batch']['state'].word,(0,1))
        with self.assertRaises(ValueError):cycle.run(MolecularState.pure(9,(0,0),(9,9)),np.random.default_rng(1))


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