import unittest
from fractions import Fraction as F
import numpy as np
from chemistry import ChemicalSource,BatchOperations,ProductRetention
from core import PureCore,binomial_event,cluster_partition
from correction import CorrectionWalk,REGIMES,one_minority_bound,fuel_deadline,finite_fuel_bound,selection_bound


class ScientificTests(unittest.TestCase):
    def test_literal_source_and_disabled_channels(self):
        s=ChemicalSource();self.assertEqual(len(s.reactions),14);self.assertEqual(s.equilibrium_activities()[:5],(100,100,1,100,100))
        z=(8,1,71,0,0,1,0);p=s.propensities(z);self.assertEqual(p[0],p[2]);self.assertEqual(p[10],56*s.rates.correction);self.assertEqual(p[12],0)
        np.testing.assert_array_equal(np.array(z)+s.changes[10],[9,0,71,0,0,0,1])
        pure=(8,0,72,0,0,64,0)
        for r in s.reactions:
            self.assertEqual(sum(r.change[:5]),0);self.assertEqual(sum(r.change[5:]),0)
            if r.change[1]>0 and r.propensity(pure)>0:self.assertEqual(r.name,'leak_X_to_Y')
        self.assertEqual(s.reactions[1].propensity((1,0,79,0,0,1,0)),0)

    def test_exact_core_recurrence_and_payoff(self):
        c=PureCore();a=c.certify();b=c.certify(F(12,25),F(9,10))
        self.assertEqual(a['lower'],F(2147232289,2**31));self.assertEqual(b['lower'],F(2147153442,2**31));self.assertEqual(a['minimum_state'],(8,0));self.assertEqual(a['state_count'],3240)
        self.assertLess(a['weighted_accumulator_bound'],2**63);self.assertEqual(a['lambda_'],3216)
        self.assertEqual(c.payoff()[c.index[(15,20)]],0);self.assertEqual(c.payoff()[c.index[(30,3)]],0)
        for n in range(16,81):self.assertEqual(binomial_event(n,8,n-8,F(12,25)),binomial_event(n,8,n-8,F(13,25)))
        with self.assertRaises(OverflowError):c.certify(scale=2**40)

    def test_consensus_harmonicity_and_fuel_mass(self):
        for r in range(3,18):
            w=CorrectionWalk(r)
            for j in range(1,r):
                down,up=w.probabilities(j);self.assertEqual(w.wrong_consensus(j),down*w.wrong_consensus(j-1)+up*w.wrong_consensus(j+1))
        w=CorrectionWalk(10);a=w.absorption(2,64)
        self.assertEqual(w.wrong_consensus(2),F(1,128));self.assertEqual(sum(v['mass'] for v in a['moments'].values())+sum(a['transient'].values()),1)
        self.assertEqual(a['moments'][0]['mass'],F(675247821429526785890499475507935979615,680564733841876926926749214863536422912))
        self.assertEqual(a['moments'][0]['spent'],F(24981573789484567970908456334429570891,10633823966279326983230456482242756608))
        self.assertEqual(w.absorption(2,0)['transient'],{2:F(1)})

    def test_joint_budgets_and_state_dependent_variation(self):
        c,co=F(2147232289,2**31),F(2147153442,2**31)
        for r in REGIMES[:4]:self.assertTrue(one_minority_bound(r,co if r.imperfect else c)['passes'])
        r=REGIMES[-1];d=fuel_deadline(r);self.assertLess(d['erlang_upper'],F(244,10**16))
        a=CorrectionWalk(10).absorption(2,64);correct=finite_fuel_bound(co,a,0,r);opposite=finite_fuel_bound(co,a,10,r)
        self.assertGreater(correct,F(2479,2500));self.assertLess(F(127,128)-correct,F(55,10**5));self.assertGreater(opposite,F(37,5000));self.assertGreater(opposite/(80*r.epsilon_max*(1+r.prefix_duration)),935)
        self.assertEqual(selection_bound(F(4999,5000)),F(16297339,16384000));self.assertEqual(selection_bound(F(9997,10000)),F(16264571,16384000))

    def test_complementary_operations_and_refill_ledger(self):
        ops=BatchOperations();z=np.array([32,1,7,35,5,0,1]);o=ops.apply(z,np.random.default_rng(19))
        before=sum(o['partition_before_refill']);np.testing.assert_array_equal(before,[32,1,7,0,0,0,1])
        self.assertEqual(o['food_supplied'],120);self.assertEqual(o['fuel_supplied'],2);self.assertEqual(o['waste_removed'],1)
        self.assertTrue(np.all(o['credited']<=o['harvested']))
        for d in o['daughters']:self.assertEqual(sum(d[:5]),80);self.assertEqual(d[5],1);self.assertEqual(d[6],0)
        gate=ProductRetention();fake=dict(credited=np.array([4,0]),daughters=o['daughters']);self.assertEqual(len(gate.retain(fake,np.random.default_rng(20))),2)

    def test_literal_trajectory_keeps_material_and_terminal_stock(self):
        source=ChemicalSource();run=source.simulate((8,1,71,0,0,1,0),1,np.random.default_rng(7))
        self.assertEqual(run['status'],'completed')
        for row in run['trace']:self.assertEqual(sum(row[1:6]),80);self.assertEqual(sum(row[6:8]),1)
        counts=run['counts'];self.assertEqual(run['state'][3],counts[2]-counts[3]);self.assertEqual(run['state'][4],counts[6]-counts[7])
        self.assertEqual(source.simulate((8,1,71,0,0,1,0),1,np.random.default_rng(7),event_budget=1)['status'],'event_budget_exhausted')

    def test_clusters_and_invalid_inputs(self):
        self.assertEqual(cluster_partition([1]*32),F(2142968775,2**31));self.assertEqual(cluster_partition([2]*16),F(32071,32768));self.assertLess(cluster_partition([4]*8),cluster_partition([2]*16))
        with self.assertRaises(ValueError):BatchOperations(partition=F(11,10))
        with self.assertRaises(ValueError):BatchOperations().apply([8.2,1,70.8,0,0,1,0],np.random.default_rng(2))
        with self.assertRaises(ValueError):CorrectionWalk(2)
        with self.assertRaises(ValueError):PureCore().certify(duration=F(1,7))


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