"""Exact source/kernel checks and clearly scoped numerical diagnostics."""
import math
import unittest
from fractions import Fraction as Q
import mpmath as mp
import numpy as np
from example import (ZipfLaw, candidates, local_classification, erase_incidence, AttributedReactor,
                     ConditioningLaw, ProofBudgets, LOCAL_INCIDENCE)
from polymer import PolymerCatalogue, CatalyticEnvironment, KineticParameters, FedReactor


class ScientificChecks(unittest.TestCase):
    def test_catalogue_singleton_equivalence_and_counts(self):
        for n in (4,5):
            c=PolymerCatalogue(n);self.assertEqual(len(c.words),2**(n+1)-2)
            self.assertEqual(len(c.channels),(n-2)*2**(n+1)+4)
            self.assertEqual(len(candidates(c)),224);self.assertEqual(len(candidates(c,False)),248)
            self.assertEqual(sum(z in c.food for z,r in candidates(c)),192)
            self.assertEqual(sum(z==c.channels[r].product for z,r in candidates(c)),32)
        c=PolymerCatalogue(4);direct=set()
        for r,ch in enumerate(c.channels):
            closure=c.closure((r,))
            for z in range(len(c.words)):
                if z in closure and {ch.left,ch.right,ch.product}<=closure:direct.add((z,r))
        self.assertEqual(direct,set(candidates(c,False)))

    def test_capped_law_moments_and_conditioned_source(self):
        law=ZipfLaw(4);weights=law.weights();R=law.R
        self.assertLess(abs(mp.fsum(weights)-1),mp.mpf('1e-75'))
        self.assertLess(abs(mp.fsum(d*w for d,w in enumerate(weights))-law.mean),mp.mpf('1e-75'))
        self.assertLess(abs(mp.fsum(d*d*w for d,w in enumerate(weights))-law.second),mp.mpf('1e-73'))
        self.assertEqual(weights[0],law.q0);self.assertGreater(weights[-1],mp.power(R,-law.a)/law.zeta)
        conditional=[d*w/law.mean for d,w in enumerate(weights)]
        self.assertLess(abs(mp.fsum(conditional)-1),mp.mpf('1e-75'))
        self.assertLess(abs(mp.fsum(d*w for d,w in enumerate(conditional))-law.second/law.mean),mp.mpf('1e-73'))
        self.assertGreater(law.second/law.mean,law.mean)
        c=PolymerCatalogue(4);rng=np.random.default_rng(300)
        for _ in range(12):
            e=law.sample(c,rng,LOCAL_INCIDENCE,empty_food=True)
            self.assertTrue(e.witness_event());self.assertTrue(all(len(row)<=R-1 for row in e.rows))
        with self.assertRaises(ValueError):law.sample(c,rng,('0','00','11'),empty_food=True)

    def test_correlated_row_probabilities_and_excess_bound(self):
        law=ZipfLaw(4);R=law.R;weights=law.weights()
        row=[]
        for k in range(3):
            row.append(mp.fsum(w*mp.mpf(math.comb(2,k)*math.comb(R-2,d-k))/math.comb(R,d)
                for d,w in enumerate(weights) if 0<=d-k<=R-2 and k<=d))
        self.assertLess(abs(row[2]-law.q2),mp.mpf('1e-75'))
        self.assertLess(abs(law.row_hit(2)-(2*law.p-law.q2)),mp.mpf('1e-75'))
        exact=mp.fsum(row[i]*row[j] for i in range(3) for j in range(3) if i+j>=2)
        bounds=law.rectangle_2_by_2()
        self.assertLess(abs(exact-bounds['actual_two_or_more']),mp.mpf('1e-75'))
        self.assertLessEqual(exact,bounds['excess_count_bound']);self.assertLessEqual(exact,bounds['pair_count_bound'])
        self.assertGreater(law.q2,law.p**2)
        self.assertLess(abs(law.candidate_union((1,1))-(2*law.p-law.p**2)),mp.mpf('1e-75'))
        self.assertLess(abs(law.row_hit(R)-(1-law.q0)),mp.mpf('1e-75'))

    def test_falling_factorial_kernel_and_incidence_only_deletion(self):
        c=PolymerCatalogue(4);assignments=(('0','0','0'),LOCAL_INCIDENCE,('1111','00','11'))
        e=CatalyticEnvironment.specified(c,assignments);p=KineticParameters.sample(e,seed=30);model=FedReactor(e,p)
        counts=np.full(len(c.words),5,dtype=np.int64);V=7
        np.testing.assert_allclose(model.rates(counts,V),[event.propensity(counts,V) for event in model.events],rtol=2e-15)
        repeated=next(event for event in model.events if event.inputs==(c.index['0'],)*3)
        self.assertAlmostEqual(repeated.propensity(counts,V),repeated.coefficient*5*4*3/V**2)
        product=c.index['0011'];reverse=next(event for event in model.events if event.inputs==(product,product))
        self.assertAlmostEqual(reverse.propensity(counts,V),reverse.coefficient*5*4/V)
        off=FedReactor(erase_incidence(e,LOCAL_INCIDENCE),p)
        self.assertEqual(len(model.events)-len(off.events),2)
        self.assertEqual([v for v in model.events if v.kind=='basal'],[v for v in off.events if v.kind=='basal'])
        self.assertTrue(any(event.incidence==(c.channel_index['00','11'],c.index['1111']) for event in off.events))
        self.assertTrue(np.all(c.lengths@model.stoich==0))
        density=np.arange(1,len(c.words)+1)/100
        chemical=model.feed-density+model.stoich@model.rates(density)
        self.assertAlmostEqual(c.lengths@chemical,10-c.lengths@density,places=10)

    def test_credit_cancellation_and_exact_kinetic_budgets(self):
        c=PolymerCatalogue(4);types=set()
        for r,ch in enumerate(c.channels):
            jump=np.zeros(len(c.words),int);jump[ch.left]-=1;jump[ch.right]-=1;jump[ch.product]+=1
            for z in range(len(c.words)):
                kind,gain,credit=local_classification(c,z,r);types.add(kind)
                if kind=='mixed':
                    self.assertEqual(int(credit@jump),0);self.assertTrue(np.all(5*credit>=3*c.nonfood))
                if kind=='neutral':self.assertEqual(int(c.nonfood@jump),0)
                if kind=='outsider':self.assertEqual(jump[z],0);self.assertIn(gain,(3,4))
        self.assertEqual(types,{'mixed','neutral','outsider','productive'})
        b=ProofBudgets.algebra()
        self.assertEqual(Q(b['mixed_half_drift']),Q(101178,25000000000))
        self.assertEqual(Q(b['outsider_export_budget']),Q(83950361,937500000))
        self.assertEqual(Q(b['other_input_budget']),Q(3186831,156250000))
        self.assertFalse(ProofBudgets.evaluate(4,20)['applicable'])
        bound=ProofBudgets.evaluate(4,10**60*25)
        self.assertTrue(bound['applicable']);self.assertEqual(mp.mpf(bound['posterior_no_productive_candidate_upper']),1)
        self.assertLess(mp.mpf(bound['log10_deletion_probability_upper']),-10**10)

    def test_three_conditionings_are_not_interchangeable(self):
        law=ConditioningLaw((Q(1,2),Q(1,3),Q(1,6)),(Q(1,10),Q(1,2),Q(9,10)))
        self.assertEqual(sum(law.successful_run()),1);self.assertEqual(sum(law.reliable(Q(2,5))),1)
        self.assertNotEqual(law.successful_run(),law.reliable(Q(2,5)))
        self.assertNotEqual(law.successful_run(),law.structural((False,True,True)))
        self.assertIsNone(ConditioningLaw((1,),(0,)).successful_run())
        self.assertIsNone(law.reliable(Q(99,100)))

    def test_attribution_balance_independent_solver_and_unfinished_path(self):
        c=PolymerCatalogue(4);e=CatalyticEnvironment.specified(c,(LOCAL_INCIDENCE,));p=KineticParameters.sample(e,seed=30092026)
        original=AttributedReactor(e,p,LOCAL_INCIDENCE);off=AttributedReactor(erase_incidence(e,LOCAL_INCIDENCE),p,LOCAL_INCIDENCE)
        times=np.r_[0.,1.,np.linspace(2,100,99)];values=original.deterministic(times);reference=original.deterministic(times,method='Radau');disabled=off.deterministic(times)
        N=len(c.words);self.assertLess(np.max(np.abs(values-reference)),1e-5)
        self.assertLess(np.max(np.abs(values[:,:N]@c.lengths-10)),1e-9)
        self.assertLess(np.max(np.abs(values[:,:N]@c.nonfood+values[:,N]-values[:,N+1:N+4].sum(axis=1))),1e-8)
        export=values[-1,N]-values[1,N];self.assertGreater(export,.1)
        self.assertGreater(values[-1,N+1],.75*export)
        self.assertLess(disabled[-1,N]-disabled[1,N],.001)
        unfinished=original.stochastic(20,1,max_events=1)
        self.assertFalse(unfinished['completed']);self.assertIsNone(unfinished['productive'])


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