"""Count semantics, endpoint ordering, exact source data, and analytic bound checks."""
import unittest
from dataclasses import replace
from fractions import Fraction as Q
from itertools import product
import mpmath as mp
import numpy as np
from example import (A,B,N_MIN,GAMMA_MAX,C,ALPHA,Cell,Population,BatchParameters,
    ChemicalRegions,ResidentChemistry,ComplementaryPartition,PopulationKernel,
    ErrorBudget,AncestralObservables,scalar_certificates,frozen_composition_diagnostic)


class ScientificChecks(unittest.TestCase):
    @classmethod
    def setUpClass(cls):cls.regions=ChemicalRegions()

    def test_root_enclosures_energy_matrices_and_integer_preparation(self):
        self.assertTrue(scalar_certificates()['exact_matrix_sandwiches'])
        for tag in ('L','H'):
            root=self.regions.roots[tag]
            self.assertLess(root.hi-root.lo,Q(1,10**50))
            for N in (10**12,N_MIN,111*10**20):
                cell=self.regions.newborn(tag,N)
                self.assertTrue(self.regions.admits(cell,4*A,closed=True))
                self.assertEqual(cell.readout,tag)
                self.assertLess(self.regions.energy(cell).hi,4*A)
        outside=Cell('H',100,(0,0,0,0))
        self.assertFalse(self.regions.admits(outside,B))

    def test_literal_count_rates_falling_factorials_and_weight_generator(self):
        chem=ResidentChemistry();self.assertEqual(len(chem.channels),13)
        cell=Cell('L',10,(3,5,1,7))
        self.assertEqual(chem.channels[5].rate(cell),0)
        self.assertEqual(chem.channels[11].rate(cell),Q(3*2,100000*10))
        eps=Q(1,100000);weights=(eps+Q(3,2),2*eps+1,Q(1),Q(3,2))
        obs=lambda c:sum(w*n for w,n in zip(weights,c.counts))+c.size
        for n in ((3,5,1,7),(0,0,0,0),(12,17,29,31)):
            cell=Cell('H',10,n);aa,bb,z,h=n;m=cell.size
            expected=(36+60*eps)*m-aa-bb-(eps+Q(1,2))*Q(bb*z,m)-2*eps*Q(aa*(aa-1),m)+8*z-Q(z*(z-1),m)-Q(3*h,20000)
            self.assertEqual(chem.generator(cell,obs),expected)
        with self.assertRaises(ValueError):chem.channels[5].apply(Cell('L',10,(0,0,1,0)))

    def test_complementary_allocation_probability_and_conservation(self):
        partition=ComplementaryPartition();counts=(3,2,1,0)
        total=sum(partition.allocation_probability(counts,d) for d in product(*(range(n+1) for n in counts)))
        self.assertEqual(total,1)
        rng=np.random.default_rng(21)
        for _ in range(20):
            x,y=partition.draw(counts,rng)
            self.assertEqual(tuple(a+b for a,b in zip(x,y)),counts)
            self.assertEqual(x[0]-Q(3,2),-(y[0]-Q(3,2)))
        with self.assertRaises(ArithmeticError):partition.draw((2**54,0,0,0),rng)

    def test_endpoint_division_and_chemical_flag_priority(self):
        # A deterministic acceptance double isolates event-order semantics;
        # this fixture is not presented as a chemically certified population.
        class Gate:
            def __init__(self,failed=None):self.failed=failed
            def admits(self,cell,threshold,closed=False):
                if threshold==2*A:assert closed
                return threshold!=self.failed
        p=BatchParameters(N=10)
        parent=Cell('H',19,(40,60,30,80))
        others=tuple(Cell('H' if i<3 else 'L',10,(20,30,15,40)) for i in range(6))
        pop=Population(21,(parent,)+others,5)
        rng=np.random.default_rng(1)
        for failed,status in [(None,'nutrient_endpoint'),(B,'outer_failure'),(2*A,'division_energy_failure'),(4*A,'partition_failure')]:
            kernel=PopulationKernel(p,Gate(failed))
            result=kernel.step(pop,0,13,rng,allocation=(20,30,14,40))
            self.assertEqual(result.status,status)
            self.assertEqual(result.precursor,20);self.assertEqual(result.total_size,80)
            self.assertEqual(len(result.cells),8);self.assertEqual(result.divisions,6)
            self.assertEqual(tuple(result.cells[0].counts[i]+result.cells[1].counts[i] for i in range(4)),(40,60,29,80))
            self.assertEqual(result.tag_size('H'),pop.tag_size('H')+1)
            with self.assertRaises(ValueError):kernel.step(result,0,13,rng)

    def test_small_resident_event_and_budget_stop_are_not_false_success(self):
        p=BatchParameters(N=10**12)
        kernel=PopulationKernel(p,self.regions);pop=kernel.initial()
        result=kernel.step(pop,0,6,np.random.default_rng(2))
        self.assertEqual(result.cells[0].counts[0],pop.cells[0].counts[0]+1)
        self.assertEqual(result.precursor+result.total_size,5*p.N*p.M)
        path=kernel.simulate(budget=3)
        self.assertEqual(path['status'],'event_budget');self.assertIsNone(path['success'])
        self.assertEqual(path['trace'][-1][0],path['time'])
        self.assertFalse(path['theorem_parameter_scope'])

    def test_published_bounds_envelopes_and_parameter_guards(self):
        for N,raw,simple in [(N_MIN,3.53114e9,3.32273e10),(111*10**20,.00311058,7.46769e5),(308*10**20,7.42729e-28,.00329691)]:
            p=BatchParameters(N=N);result=ErrorBudget(p).evaluate()
            self.assertAlmostEqual(float(result['chemical_error']['raw'])/raw,1,delta=2e-6)
            self.assertAlmostEqual(float(result['chemical_error']['simplified'])/simple,1,delta=2e-6)
            for kind in ('retained','double_exponent','simplified'):
                self.assertGreaterEqual(float(result['chemical_error'][kind]),float(result['chemical_error']['raw']))
        result=ErrorBudget(BatchParameters(gain=Q(1))).evaluate()
        self.assertEqual(float(result['success_lower']['raw']),0)
        result=ErrorBudget(BatchParameters(N=1000)).evaluate()
        self.assertIsNone(result['success_lower']['raw'])
        result=ErrorBudget(BatchParameters(scaled_deadline=0)).evaluate()
        self.assertEqual(float(result['success_lower']['raw']),0)
        for double in (False,True):
            N=ErrorBudget().sufficient_size(double_exponent=double)
            result=ErrorBudget(BatchParameters(N=N)).evaluate()
            kind='double_exponent' if double else 'simplified'
            self.assertGreaterEqual(float(result['success_lower'][kind]),.99)
        with self.assertRaises(ValueError):ErrorBudget(BatchParameters(gain=Q(1,5))).sufficient_size()

    def test_ancestral_odds_phase_cost_and_safe_generator(self):
        p=BatchParameters(N=1000);obs=AncestralObservables(p)
        pop=Population(8000,(Cell('H',1999,(0,0,1,0)),Cell('L',1000,(0,0,1,0))))
        values=obs.size_to_count_ratio(pop)
        self.assertEqual(values['count_odds'],1);self.assertEqual(values['size_odds'],Q(1999,1000))
        self.assertLess(values['phase_factor'],2)
        for zH,zL in product((Q(297,100),Q(3)),(Q(99,100),Q(101,100))):
            self.assertLess(obs.relative_odds_drift(pop,zH,zL),0)
        phase=frozen_composition_diagnostic(p,self.regions)
        self.assertGreater(phase[1][5],1);self.assertEqual(phase[1][6],1)
        self.assertAlmostEqual(phase[-1][7],.25,places=12)
        self.assertTrue(all(.5<r[5]/r[6]<2 for r in phase))


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