"""Exact chemical/transfer checks and actual-state protocol semantics."""
from fractions import Fraction as Q
from dataclasses import replace
import unittest
import numpy as np
from mpmath import mp
from example import (ChemicalModel,ChemicalRegions,Cell,A,Population,PopulationSource,UniformTransfer,
    SerialProtocol,TwoCycleCertificate,OddsAccounting,NEWBORN_SIZE)

class ScienceTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):cls.regions=ChemicalRegions();cls.chemistry=ChemicalModel()
    def test_completion_source_and_ready_preparation(self):
        self.assertTrue(self.chemistry.certificates()['stationary_curve_identity'])
        for tag in ('L','H'):
            cell=self.regions.newborn(tag,NEWBORN_SIZE)
            self.assertTrue(self.regions.admits(cell,A,closed=True));self.assertEqual(cell.readout,tag)
            changed=replace(cell,tag='H' if tag=='L' else 'L')
            np.testing.assert_array_equal(self.chemistry.rates(cell),self.chemistry.rates(changed))
        for n in ((2,3,1,4),(0,0,0,0),(12,21,32,11)):
            cell=Cell('L',10,n)
            np.testing.assert_allclose(self.chemistry.rates(cell),[float(c.rate(cell)) for c in self.chemistry.channels],rtol=1e-15,atol=0)
        self.assertEqual(Cell('H',10,(1,99,20,1)).readout,'L')
        self.assertEqual(Cell('L',10,(1,999,21,1)).readout,'H')

    def test_complementary_division_at_endpoint_and_arbitrary_phase_cap(self):
        source=PopulationSource(2,1);parent=Cell('H',3,(10,12,50,15));pop=Population((parent,),1,4,0)
        after=source.step(pop,0,13,np.random.default_rng(1),allocation=(3,4,17,6))
        self.assertEqual(after.precursor,0);self.assertEqual([c.size for c in after.cells],[2,2]);self.assertEqual(after.divisions,1)
        self.assertEqual(tuple(a+b for a,b in zip(after.cells[0].counts,after.cells[1].counts)),(10,12,49,15))
        self.assertEqual(after.size+after.precursor,pop.size+pop.precursor)
        # A non-newborn founder can exceed the old 4M cell cap at the 4W0 endpoint.
        pop=source.refill((replace(parent,counts=(10,12,5000,15)),));rng=np.random.default_rng(21)
        for _ in range(9):
            index=max(range(len(pop.cells)),key=lambda i:pop.cells[i].size)
            grown=list(pop.cells[index].counts);grown[2]-=1
            allocation=tuple(n//2 for n in grown) if pop.cells[index].size==3 else None
            pop=source.step(pop,index,13,rng,allocation)
        self.assertEqual(pop.size,12);self.assertEqual(pop.precursor,pop.endpoint);self.assertGreater(len(pop.cells),4)

    def test_exact_transfer_and_inclusion_covariance(self):
        transfer=UniformTransfer();expected=[('equal',[10]*8,('17/35','1/35')),('mixed',[19,10,19,10,10,19,10,19],('27/35','13/35'))]
        for name,sizes,failures in expected:
            cells=tuple(Cell('H' if i<4 else 'L',m,(1,2,3,4)) for i,m in enumerate(sizes))
            for eps,want in zip(('1/50','1/2'),failures):
                result=transfer.enumerate(cells,4,eps);self.assertEqual(result['subsets'],70);self.assertEqual(result['failure'],want)
            weights=[Q(c.size,20) if c.tag=='H' else Q(0) for c in cells];mom=transfer.weighted_moments(weights,4)
            direct=transfer.enumerate(cells,4,'1/50')['moments']['H']
            self.assertEqual(Q(mom['mean']),Q(direct['mean'])/20);self.assertEqual(Q(mom['variance']),Q(direct['variance'])/400);self.assertLess(Q(mom['pair_covariance']),0)
            selected,discarded=transfer.sample(cells,4,np.random.default_rng(17));swapped=tuple(replace(c,tag='H' if c.tag=='L' else 'L') for c in cells)
            selected2,_=transfer.sample(swapped,4,np.random.default_rng(17));self.assertEqual([c.size for c in selected],[c.size for c in selected2]);self.assertEqual(sum(c.size for c in selected+discarded),sum(sizes))
        self.assertEqual(transfer.weighted_moments([Q(2,3)],1)['variance'],'0')

    def test_capped_service_projection_hard_cutoff_and_frozen_recovery(self):
        source=PopulationSource(10,1);cells=(Cell('H',10,(20,15,30,40)),Cell('L',13,(22,24,13,30)));pop=replace(source.refill(cells),precursor=0)
        a,log_a=source.run(pop,.05,np.random.default_rng(4),5000,False)
        b,log_b=source.run(pop,.05,np.random.default_rng(4),5000,False,quota=1)
        self.assertEqual(a,b);self.assertEqual(log_a['gross_exchange'],log_b['gross_exchange']);self.assertEqual(log_a['events'],log_b['events']);self.assertTrue(log_b['saturated']);self.assertEqual(log_b['status'],'duration')
        self.assertEqual([c.size for c in a.cells],[10,13]);self.assertEqual([c.tag for c in a.cells],['H','L'])
        c,hard=source.run(pop,.05,np.random.default_rng(4),5000,False,hard_limit=10000);self.assertEqual(a,c);self.assertEqual(hard['status'],'duration')
        _,small=source.run(pop,.05,np.random.default_rng(4),5000,False,hard_limit=0);self.assertEqual(small['status'],'hard_service_cutoff');self.assertEqual(small['gross_exchange'],[0]*5)
        _,limited=source.run(pop,5376,np.random.default_rng(4),1,False);self.assertFalse(limited['complete']);self.assertEqual(limited['status'],'event_budget')

    def test_certificate_witnesses_and_scope(self):
        for N,M,target in [(65536*10**18,4*10**9,Q(495999,500000)),(131072*10**18,10**15,Q(497999,500000))]:
            cert=TwoCycleCertificate(N,M);rational=cert.rational_witness();self.assertGreaterEqual(Q(rational['success_lower']),target)
            result=cert.evaluate();self.assertGreaterEqual(Q(result['success_lower_certified']),target)
            self.assertLess(cert.T*cert.qB/cert.JB,Q(1,1000));self.assertLess(Q(M*5376*cert.qR,cert.JR),Q(1,1000))
        cert=TwoCycleCertificate();self.assertEqual(cert.JB,58720256*10**45+1);self.assertEqual(cert.JR,47915728896*10**33+1)
        self.assertEqual(cert.evaluate()['success_lower_certified'],'0.993959999')
        self.assertEqual(TwoCycleCertificate(14*10**19,4*10**9).evaluate()['success_lower_certified'],'0.000000000')
        for args in [(10**12,4,'1e-11'),(NEWBORN_SIZE,3,'1e-11'),(NEWBORN_SIZE,4,'1e-10')]:
            with self.assertRaises(ValueError):TwoCycleCertificate(*args)

    def test_endpoint_phase_and_distinct_count_readout(self):
        newborn=(Cell('H',10,(1,2,30,4)),Cell('L',10,(1,2,10,4)))
        phased=(replace(newborn[0],size=19),newborn[1]);a=OddsAccounting.coordinates(newborn);b=OddsAccounting.coordinates(phased)
        self.assertEqual(a['phase'],'0.0');self.assertEqual(b['log_count'],'0.0');self.assertGreater(mp.mpf(b['log_size']),0)
        self.assertLess(abs(mp.mpf(b['log_count'])-(mp.mpf(b['log_size'])-mp.mpf(b['phase']))),mp.mpf('1e-28'))
        thresholds=OddsAccounting.theorem_thresholds();self.assertEqual(thresholds['two_fraction_certified_lower'],'0.6930453')
        self.assertGreater(Q('49/1600'),Q(1,100));self.assertLess(Q('2401/1280000'),Q(1,100))
        g=mp.mpf(thresholds['cycle_size_gain']);self.assertLess(abs(2*g-mp.log(2)-mp.mpf(thresholds['two_count_gain'])),mp.mpf('1e-28'))

    def test_actual_selected_recovery_and_protocol_continuation(self):
        cell=self.regions.newborn('H',NEWBORN_SIZE);perturbed=replace(cell,counts=tuple(n+(NEWBORN_SIZE//50000 if j==2 else 0) for j,n in enumerate(cell.counts)))
        self.assertFalse(self.regions.admits(perturbed,A,closed=True));self.assertTrue(self.regions.admits(perturbed,8*A))
        a=self.chemistry.recovery_density(perturbed);b=self.chemistry.recovery_density(perturbed,method='BDF');times=np.r_[0,np.geomspace(.01,5376,70)]
        self.assertLess(np.max(abs(a.sol(times)-b.sol(times))),2e-8)
        initial=np.array(perturbed.counts,float)/perturbed.size;np.testing.assert_array_equal(a.sol(0),initial)
        class TraceSource:
            gamma=Q(1)
            def __init__(self):self.starts=[];self.physical=PopulationSource(10,1)
            def refill(self,cells):self.starts.append(tuple(cells));return self.physical.refill(cells)
            def run(self,pop,duration,rng,event_budget,collect_endpoint=True):
                if collect_endpoint:
                    copies=tuple(c for c in pop.cells for _ in range(4));result=replace(pop,cells=copies,precursor=pop.endpoint)
                    return result,dict(status='endpoint',events=1,time=1)
                recovered=tuple(replace(c,counts=(c.counts[0]+1,*c.counts[1:])) for c in pop.cells)
                return replace(pop,cells=recovered),dict(status='duration',events=1,time=duration)
        source=TraceSource();small=(Cell('H',10,(1,2,30,4)),Cell('L',10,(1,2,10,4)))
        result=SerialProtocol(source,2).run(small,2,7,10);self.assertTrue(result['complete'])
        self.assertEqual([c.counts[0] for c in source.starts[1]],[2,2]);self.assertEqual([c.counts[0] for c in source.starts[2]],[3,3])

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