import unittest
import numpy as np
from example import (Q, Description, ReversibleNetwork, ChannelQuotient,
    CatalyticPeeler, IndependentCatalysis, ExactCatalyticLaw, ChannelCensus,
    seed_cutoff, toy, probability, wilson)


class ScientificChecks(unittest.TestCase):
    def test_canonical_identity_counts_and_gateways(self):
        self.assertEqual(Description('AA','A').canonical(),Description('A','AA'))
        self.assertNotEqual(Description('A','B').canonical(),Description('B','A').canonical())
        self.assertEqual(Description('AB','AB').canonical(),Description('AB','AB'))
        for n in range(2,9):
            q=ChannelQuotient.binary(n);row=ChannelCensus.formula(n)
            self.assertEqual(len(q.split.channels),row['split']);self.assertEqual(len(q.quotient.channels),row['quotient'])
            self.assertLessEqual(row['deleted'],row['elementary_loss_bound']);self.assertLessEqual(max(map(len,q.fibres)),2)
            if n>=4:self.assertEqual((len(q.quotient.gateways()),len(q.quotient.gateways(True))),(34,30))
        self.assertEqual(ChannelCensus.formula(12)['deleted'],96)
        self.assertEqual(ChannelCensus.formula(24)['quotient'],738196506)
        with self.assertRaises(ValueError):ChannelQuotient(ReversibleNetwork(('A','AA'),('A',),(Description('A','A'),Description('A','A'))))

    def test_closure_projection_and_witness_lifting(self):
        q=toy()
        for mask in range(8):
            active=[i for i in range(3) if mask>>i&1];projected=set(q.mapping[i] for i in active)
            self.assertEqual(q.split.closure(active),q.quotient.closure(projected))
        for marks in ExactCatalyticLaw(q.split).configurations():
            a=CatalyticPeeler(q.split).run(marks);b=CatalyticPeeler(q.quotient).run(q.project_marks(marks))
            brute=any(q.split.is_raf([j for j in range(3) if mask>>j&1],marks) for mask in range(1,8))
            self.assertEqual(a['has_raf'],brute);self.assertEqual(a['has_raf'],b['has_raf'])
            if b['has_raf']:self.assertTrue(q.split.is_raf(q.lift_witness(b['usable'],marks),marks))
        # Only the noncanonical duplicate is catalysed: fixed representative loses RAF.
        marks=(0,1,0);b=CatalyticPeeler(q.quotient).run(q.project_marks(marks));self.assertTrue(b['has_raf'])
        self.assertEqual(q.lift_witness(b['usable'],marks),(1,))
        self.assertFalse(q.split.is_raf((0,),marks))

    def test_exact_probability_laws_and_domination(self):
        q=toy();sp=ExactCatalyticLaw(q.split);qt=ExactCatalyticLaw(q.quotient)
        for p,a,b in [(Q(1,2),Q(4077,4096),Q(239,256)),(Q(1,3),Q(505585,531441),Q(5137,6561))]:
            self.assertEqual(sp.probability([p]*3),a);self.assertEqual(qt.probability(q.or_parameters(p)),a);self.assertEqual(qt.probability([p]*2),b)
        for a in (Q(0),Q(1,10),Q(1,2),Q(1)):
            self.assertTrue(all(p<=a for p in q.or_parameters(a/2)))
            self.assertTrue(all(p>=a for p in q.or_parameters(a)))
        with self.assertRaises(ValueError):probability('11/10')
        with self.assertRaises(ValueError):probability(.1)

    def test_complete_history_identity_and_exact_column_weights(self):
        law=ExactCatalyticLaw(toy().quotient);histories=law.history_polynomials()
        self.assertEqual(histories,law.history_polynomials(True));self.assertEqual(len(histories),6)
        for p in (Q(0),Q(1,3),Q(1,2),Q(1)):
            total=Q(0)
            for history,counts in histories.items():
                expected=sum(count*p**k*(1-p)**(law.bits-k) for k,count in counts.items())
                self.assertEqual(law.history_probability(history,p),expected);total+=expected
            self.assertEqual(total,1)

    def test_reversible_cleavage_food_only_and_unusable_channels(self):
        network=ReversibleNetwork(('A','B','AB'),('AB',),(Description('A','B'),))
        self.assertEqual(network.closure((0,)),7);self.assertTrue(network.is_raf((0,),(4,)))
        food_only=ReversibleNetwork(('A','AA'),('A','AA'),(Description('A','A'),))
        result=CatalyticPeeler(food_only).run((1,));self.assertTrue(result['has_raf']);self.assertEqual(result['closure'],food_only.food)
        unusable=ReversibleNetwork(('A','AA','AAA','AAAA'),('A',),(Description('AA','AA'),))
        result=CatalyticPeeler(unusable).run((1,));self.assertFalse(result['has_raf']);self.assertEqual(result['history'],((0,),))
        # Correlated duplicate observations keep p; independent OR would change it.
        q=toy()
        for marks in ExactCatalyticLaw(q.quotient).configurations():
            duplicated=tuple(marks[j] for j in q.mapping);self.assertEqual(q.project_marks(duplicated),marks)

    def test_normalization_cutoffs_and_finite_event_scope(self):
        row=ChannelCensus.formula(12);p=Q(12,row['quotient'])
        self.assertAlmostEqual(ChannelCensus.openness(p,row['M']),.699150,places=6)
        self.assertEqual(seed_cutoff('1/10',10)['L'],536670);self.assertEqual(seed_cutoff('1/2',10)['L'],20820)
        net=ChannelQuotient.binary(4).quotient
        for j in net.gateways():
            H=net.closure((j,));escape=any(len(w)>2 for w in net.generated_words(H))
            self.assertEqual(escape,j in net.gateways(True))
        inactive=set(range(len(net.channels)))-set(net.gateways(True));self.assertEqual(net.closure(inactive),net.food)
        self.assertEqual(ChannelCensus.openness(0,100),0);self.assertEqual(ChannelCensus.openness(1,100),1)

    def test_sampling_reproducibility_boundaries_and_limits(self):
        net=toy().quotient;a=IndependentCatalysis(np.random.default_rng(123));b=IndependentCatalysis(np.random.default_rng(123))
        for _ in range(20):self.assertEqual(a.sample(net,'1/3'),b.sample(net,'1/3'))
        self.assertEqual(a.sample(net,0),(0,0));self.assertEqual(a.sample(net,1),(15,15))
        self.assertEqual(wilson(0,100)[0],0);self.assertEqual(wilson(100,100)[1],1);self.assertIsNone(wilson(0,0))
        self.assertGreater(wilson(0,100)[1],0);self.assertLess(wilson(100,100)[0],1)
        with self.assertRaises(ValueError):ExactCatalyticLaw(ChannelQuotient.binary(3).quotient)
        with self.assertRaises(ValueError):ChannelQuotient.binary(13)
        with self.assertRaises(ValueError):CatalyticPeeler(net).run((16,0))


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