import unittest
from fractions import Fraction as F
from dataclasses import replace
import numpy as np
from example import exact_audit,InheritanceCertificate
from resident import ResidentChemistry,Cell,rational_root,residual
from feedback import Community,WindowReadout,uptake,load
from protocol import PopulationSource,ComplementaryDivision,UniformTransfer,SerialProtocol


class ScientificChecks(unittest.TestCase):
    def test_exact_stationary_exclusion_load_and_count_identities(self):
        audit=exact_audit();self.assertEqual(audit['load_cap'],F(2239911,97656250))
        self.assertGreater(audit['uptake_ratio_lower'],F('1.8186'))
        for rho in ['9999/1000000','1/100']:
            for tag in ['L','H']:
                lo,hi=rational_root(rho,tag,50)
                self.assertLess(residual(lo,F(rho))*residual(hi,F(rho)),0)
        self.assertLess(float(audit['load_cap']),.031)

    def test_full_error_budget_and_finite_repetition(self):
        cert=InheritanceCertificate().evaluate()
        self.assertEqual(cert['joint_two_cycle_lower'],F(24482873,24500000))
        self.assertLess(cert['nontransfer_upper'],3e-6)
        self.assertEqual(cert['finite_repetition'][2]['probability_lower'],F(2373714351,2401000000))
        self.assertEqual(cert['finite_repetition'][3]['probability_lower'],F(23957912563,29412250000))
        self.assertEqual(cert['finite_repetition'][4]['probability_lower'],0)
        self.assertAlmostEqual(cert['high_fraction_lower_numerical'],.6930453422,places=9)
        with self.assertRaises(ValueError):InheritanceCertificate(rho='.02')
        with self.assertRaises(ValueError):InheritanceCertificate(N=1000)

    def test_full_feedback_equilibria_and_resident_bridge(self):
        C=Community((.1,.2,.3,.4));low,high=C.equilibria()
        for state in [low,high]:
            o=C.observables(state);self.assertLess(max(abs(C.field(state))),1e-10)
            np.testing.assert_allclose(o['p'],C.q,atol=1e-15)
            self.assertLess(max(np.linalg.eigvals(C.jacobian(state)).real),0)
            np.testing.assert_allclose(C.field(state)[:4],ResidentChemistry().density(state[:4],o['load']))
        self.assertGreater(C.observables(high)['J']/C.observables(low)['J'],1.8186)
        self.assertLess(C.observables(high)['load'],C.observables(low)['load'])
        # Cap is stationary only: this physically positive state exceeds it.
        state=low.copy();state[4]=1;state[5:]=.1*C.q
        self.assertGreater(C.observables(state)['load'],.031)

    def test_transient_readouts_composition_and_derivative_exception(self):
        C=Community();state=C.equilibria()[0];state[5:]*=[1.1,.9,1.03]
        o=C.observables(state);f=C.field(state);co=C.composition(state)
        self.assertAlmostEqual(o['J'],sum(f[5:])+.5*o['S']+o['S']**2+o['W'],places=14)
        pdot=(f[5:]*o['S']-state[5:]*sum(f[5:]))/o['S']**2
        np.testing.assert_allclose(pdot,co['pdot'],atol=1e-16)
        self.assertAlmostEqual(co['v']@pdot/(co['v']@co['v']),o['S'],places=13)
        self.assertAlmostEqual(2*sum(o['p']*pdot/C.q),co['chi_prime'],places=14)
        self.assertAlmostEqual(-sum(C.q*pdot/o['p']),co['KL_prime'],places=14)
        run=C.integrate(state,100,501)
        self.assertLess(abs(run['abundance_balance_error']),1e-10)
        self.assertLess(abs(run['reservoir_balance_error']),1e-10)
        values=[C.observables(u) for u in run['states']];window=WindowReadout()
        js,jr=window.estimates(run['times'],[v['S'] for v in values],[v['W'] for v in values],run['states'][:,4])
        truth=run['integrals'][-1,0]/100
        self.assertLess(abs(js-truth),1e-7);self.assertLess(abs(jr-truth),1e-7)
        equilibrium=C.equilibria()[1];self.assertEqual(C.composition(equilibrium)['derivative_condition'],float('inf'))

    def test_readout_assumptions_consistency_and_error_units(self):
        W=WindowReadout();self.assertEqual(W.errors,(F('.000791'),F('.00036')))
        args=dict(calibrated=True,two_alternatives=True,recovery_bound='.001')
        self.assertEqual(W.classify('.0224','.0222',**args)['outcome'],'low')
        self.assertEqual(W.classify('.0402','.0401',**args)['outcome'],'high')
        for a,b in [('.0224','.0401'),('.031','.031')]:self.assertEqual(W.classify(a,b,**args)['outcome'],'unresolved')
        self.assertEqual(W.classify('.0224','.0222')['outcome'],'unresolved')
        self.assertIsNone(W.recovery_error(1e-12,31,100))
        self.assertGreater(W.recovery_error(1e-12,31,100,admitted_sublevel=True),0)
        with self.assertRaises(ValueError):WindowReadout(h=0)

    def test_extraction_growth_division_and_actual_recovery(self):
        chemistry=ResidentChemistry();source=PopulationSource(2,2,chemistry);rng=np.random.default_rng(5)
        cell=Cell('H',3,(4,5,6,7));pop=source.refill([cell]);before=sum(c.counts[2] for c in pop.cells)
        extracted=source.step(pop,0,13,rng)
        self.assertEqual(extracted.size,pop.size);self.assertEqual(extracted.cells[0].counts[2],before-1)
        divided=source.step(pop,0,14,rng,allocation=(1,2,3,4))
        self.assertEqual(len(divided.cells),2)
        self.assertEqual(tuple(sum(c.counts[i] for c in divided.cells) for i in range(4)),(4,5,5,7))
        self.assertEqual(divided.size+divided.precursor,pop.size+pop.precursor)
        # Recovery disables growth only; extraction propensity remains positive.
        recovering=replace(pop,precursor=0)
        self.assertEqual(source.growth_rate(recovering,cell),0);self.assertGreater(chemistry.rates(cell)[13],0)
        end,log=source.run(recovering,.1,rng,20000,collect_endpoint=False)
        self.assertEqual(log['status'],'duration')
        final=sum(c.counts[2] for c in end.cells)
        self.assertEqual(log['signed_internal_formation'],final-before+log['extracted']+log['growth_consumed'])
        tiny=Cell('L',2,(1,0,1,0))
        self.assertEqual(chemistry.channels[5].rate(tiny),0);self.assertEqual(chemistry.channels[11].rate(tiny),0)

    def test_neutral_intact_transfer_and_unfinished_law(self):
        cells=tuple(Cell(tag,2,(20,30,4,10)) for tag in ['L']*4+['H']*4);T=UniformTransfer();rng=np.random.default_rng(8)
        selected,discarded=T.sample(cells,4,rng)
        self.assertEqual(len(selected),4);self.assertEqual(len(discarded),4)
        self.assertTrue(all(any(c is old for old in cells) for c in selected))
        exact=T.enumerate(cells,4,'1/50');self.assertEqual(exact['subsets'],70)
        # Two monochromatic subsets out of 70; immutable lost tags cannot return.
        import itertools
        lost=sum(len({cells[i].tag for i in subset})==1 for subset in itertools.combinations(range(8),4))
        self.assertEqual(lost,2)
        result=SerialProtocol(PopulationSource(2,2),4,.1).run(cells[:4],2,7,0)
        self.assertFalse(result['complete']);self.assertEqual(result['stage'],'batch')


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