"""Scientific checks of chemistry, reduction, count semantics and certificate scope."""
import unittest
from fractions import Fraction as F
import numpy as np
from example import exact_checks,MissionCertificate,literal_initial
from reactor import (Chemistry,CommonExchange,Intervention,HistoryPolicy,ReactorNetwork,
    DeterministicMission,MolecularMission,ready,A,B,I,J,Y)


class ScientificChecks(unittest.TestCase):
    def network(self,refined=True):
        return ReactorNetwork([Chemistry(refined=refined),Chemistry(refined=refined)],CommonExchange([[0,1],[1,0]]))

    def test_exact_source_phase_and_witness(self):
        d=exact_checks();self.assertEqual(d['net_synthesis_positive_from_cycle'],55)
        net=self.network();V=224000000000000;bound=MissionCertificate(net).bound(V,100,literal_initial(V))
        self.assertGreater(bound['lower_probability'],.99)
        self.assertEqual(bound['net_synthesis_lower_worst_ready'],365200000000000)
        self.assertGreaterEqual(MissionCertificate(net).sufficient_scale(100,'1/100'),212332534120360)
        self.assertEqual(MissionCertificate(net).bound(V,0,literal_initial(V))['lower_probability'],1)
        self.assertEqual(MissionCertificate(net).bound(14000,100,literal_initial(14000))['lower_probability'],0)

    def test_literal_duplex_and_mark_semantics(self):
        m=Chemistry();self.assertEqual(len(m.channels),23)
        channel=next(c for c in m.channels if c.name=='pair4-')
        N=[0,0,1,0,0,0,0];self.assertEqual(channel.propensity(N,100),0)
        N[2]=2;self.assertEqual(channel.propensity(N,100),F(2,5))
        np.testing.assert_allclose(m.rates(N,100),[float(c.propensity(N,100)) for c in m.channels])
        washD=next(c for c in m.channels if c.wash==6);self.assertEqual(washD.marks(True),(0,0,0,1,0,0,1))
        washZ=next(c for c in m.channels if c.wash==5);self.assertEqual(washZ.marks(True)[:2],(2,0))
        self.assertEqual(sum(c.marks(False)[2] for c in m.channels),2)
        self.assertEqual(next(c for c in m.channels if c.name=='pair6+').marks(True)[2],0)

    def test_material_reduction_after_species_specific_pulse(self):
        net=self.network(False);c=literal_initial(1400,False)/1400
        pulse=Intervention('3/5',('1','49/50','1','49/50','1','49/50'),('-1/200','1/200'))
        post=np.array([pulse.deterministic(row)[0] for row in c])
        self.assertGreater(abs(post@A[:6]-1).max(),.005)
        full=net.evolve(post);reduced=net.donor_reduced(post)
        np.testing.assert_allclose(full['end'],reduced,atol=1e-10,rtol=1e-9)
        self.assertLess(full['material_error'],1e-10)
        # Graph exchange cancels globally, but not at individual nodes.
        np.testing.assert_allclose((net.exchange.D@post).sum(axis=0),0,atol=1e-15)
        self.assertGreater(np.max(abs(net.exchange.D@post)),0)

    def test_retained_intermediate_nonclosure_and_filter(self):
        m=Chemistry();a=np.array([.95,.95,.05,0,0,0,0]);b=a.copy();b[6]=.005
        self.assertTrue(ready(a));self.assertTrue(ready(b))
        np.testing.assert_allclose((m.field(b)-m.field(a))[:6],[.015,.015,.00015,0,0,0],atol=1e-14)
        b[6]=m.filter_equilibrium(b);self.assertAlmostEqual(m.field(b)[6],0,places=14)
        j5=float(m.d*(1+m.beta))*b[2]-float(m.d*m.beta/m.theta)*b[6]
        j6=float(m.d/m.theta)*b[6]-float(m.d*F(1,8000000000)*(1+m.beta)/m.beta)*b[0]*b[1]
        self.assertAlmostEqual(j5-j6,b[6],places=14)
        self.assertGreater(b[6],0)

    def test_repeated_ode_ledgers_and_independent_accuracy(self):
        net=self.network();c=literal_initial(1400)/1400
        run=DeterministicMission(net,HistoryPolicy()).run(c,2)
        self.assertTrue(all(h['ready'] for h in run['history']))
        self.assertLess(abs(run['inventory_telescope_residual']),1e-10)
        self.assertLess(abs(run['storage_telescope_residual']),1e-10)
        fine=DeterministicMission(net,HistoryPolicy()).run(c,1,rtol=2e-11)
        np.testing.assert_allclose(run['history'][0]['accounts'],fine['history'][0]['accounts'],atol=2e-10,rtol=1e-8)
        acc=np.array(run['history'][0]['accounts']);self.assertAlmostEqual(acc[0,0],.275164,places=6)
        self.assertAlmostEqual(acc[0,1],.128711,places=6)
        self.assertAlmostEqual(acc[0,2],.007416,places=6)

    def test_count_pulse_law_actual_endpoints_and_ledgers(self):
        pulse=Intervention('1/4',('49/50',)*7,('-1/200','1/200'));N=np.array([20,20,8,2,3,4,1]);rng=np.random.default_rng(7)
        after,removed,lost,dose=pulse.molecular(N,20,rng)
        recovered=after+removed+lost;recovered[:2]-=dose;np.testing.assert_array_equal(N,recovered)
        self.assertEqual(dose.tolist(),[14,15])
        net=ReactorNetwork([Chemistry()],CommonExchange([[0]]));initial=np.array([[19,19,1,0,0,0,0]])
        run=MolecularMission(net,HistoryPolicy(),20,17).run(initial,2,20000)
        self.assertEqual(run['status'],'completed');self.assertEqual(len(run['history']),2)
        self.assertEqual(run['inventory_telescope_residual'],0);self.assertEqual(run['storage_telescope_residual'],0)
        stopped=MolecularMission(net,HistoryPolicy(),20,17).run(initial,1,0)
        self.assertEqual(stopped['status'],'unfinished');self.assertNotIn('all_success',stopped)

    def test_invalid_scope_and_exact_ready_boundary(self):
        with self.assertRaises(ValueError):CommonExchange([[0,1],[0,0]])
        with self.assertRaises(ValueError):Intervention('1/5')
        with self.assertRaises(ValueError):MissionCertificate(self.network()).bound(9999,1,np.zeros((2,7),int))
        net=ReactorNetwork([Chemistry(theta='1/1000')],CommonExchange([[0]]))
        with self.assertRaises(ValueError):MissionCertificate(net).bound(10000,1,np.array([[9500,9500,500,0,0,0,0]]))
        self.assertTrue(net.chemistries[0].scope(False))
        self.assertFalse(net.chemistries[0].scope(True))
        good=literal_initial(1400);self.assertTrue(ready(good,V=1400));bad=good.copy();bad[0,2]-=1
        self.assertFalse(ready(bad,V=1400))


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