from dataclasses import replace
from fractions import Fraction as Q
import unittest
import numpy as np
from example import (BindingReactor,ReactorParameters,OperatingPolicy,OperatingMonitor,
                     ProbabilityBudget,symbolic_bookkeeping,RESOURCE_A,RESOURCE_B,
                     CATALYTIC_WEIGHTS,COVALENT_STOCK)


class ScientificChecks(unittest.TestCase):
    def test_literal_channels_and_pair_equilibrium(self):
        for K in (8,10,12):
            model=BindingReactor(ReactorParameters(Q(K),Q(20)))
            self.assertEqual(len(model.channels),18)
            state=np.array((1,1,K,K,K,K*K),dtype=float)
            rates=model.rates(state)
            for i in range(0,10,2):self.assertAlmostEqual(rates[i],rates[i+1],10)
            for channel in model.channels[:10]:
                self.assertEqual(sum(w*d for w,d in zip(RESOURCE_A,channel.jump)),0)
                self.assertEqual(sum(w*d for w,d in zip(RESOURCE_B,channel.jump)),0)
        on,off=BindingReactor(),BindingReactor(enabled=False)
        self.assertEqual({c.name for c in on.channels}-{c.name for c in off.channels},
                         {'templated_ligation_forward','templated_ligation_reverse'})
        self.assertEqual(len(off.channels),16)

    def test_count_support_generator_and_falling_factorial(self):
        model=BindingReactor()
        for counts in ((100,100,0,0,0,0),(80,80,1,3,2,4),(77,83,4,3,5,6)):
            rates=model.rates(counts,100)
            for i,channel in enumerate(model.channels):
                self.assertAlmostEqual(rates[i],float(channel.count_rate(counts,100)),10)
                if rates[i] > 0:self.assertTrue(all(n+d >= 0 for n,d in zip(counts,channel.jump)))
            for weights in (RESOURCE_A,RESOURCE_B):
                drift,quad=model.generator(counts,100,weights)
                stock=sum(w*n for w,n in zip(weights,counts))
                self.assertEqual(drift,100-stock)
                self.assertLessEqual(quad,100+2*stock)
        self.assertEqual(model.rates((80,80,1,0,0,0),100)[9],0)
        self.assertAlmostEqual(model.rates((80,80,2,0,0,0),100)[9],.4)

    def test_symbolic_and_rational_probability_certificates(self):
        self.assertTrue(symbolic_bookkeeping()['disabled_covalent_balance'])
        bounds=ProbabilityBudget.evaluate()
        self.assertGreater(bounds['enabled_success_lower'],Q(9,10))
        self.assertLess(bounds['disabled_conservative_event_upper'],Q(1,5000))
        self.assertEqual(bounds['paper_coarse_disabled_upper'],Q(123,625000))
        self.assertTrue(ProbabilityBudget.applies(ReactorParameters(),OperatingPolicy()))
        self.assertFalse(ProbabilityBudget.applies(ReactorParameters(Q(7)),OperatingPolicy()))
        self.assertFalse(ProbabilityBudget.applies(ReactorParameters(),replace(OperatingPolicy(),volume=1000)))

    def test_entry_return_and_deadline_gate(self):
        p=OperatingPolicy(volume=100,entry=4,floor=2,deadline=5,end=10,export=8)
        monitor=OperatingMonitor(p)
        monitor.jump(2,(96,96,4,0,0,0),0)
        self.assertEqual(monitor.phase,1)
        self.assertEqual(monitor.entry_time,2)
        monitor.jump(3,(97,97,3,0,0,0),0)
        monitor.jump(4,(99,99,1,0,0,0),0)
        self.assertEqual(monitor.status,'return_failure')
        self.assertFalse(monitor.report()['event_verdict'])
        missed=OperatingMonitor(p)
        missed.jump(5,(100,100,0,0,0,0),0)
        missed.jump(5.1,(96,96,4,0,0,0),0)
        self.assertEqual(missed.status,'missed_deadline')
        self.assertEqual(missed.counts.tolist(),[100,100,0,0,0,0])

    def test_export_window_saturation_and_disabled_event(self):
        p=OperatingPolicy(volume=100,entry=4,floor=2,deadline=5,end=10,export=8)
        m=OperatingMonitor(p)
        m.jump(1,(96,96,4,0,0,0),0)
        m.jump(5,(96,96,4,0,0,0),8)
        self.assertEqual(m.counter,0)
        m.jump(6,(96,96,4,0,0,0),8)
        self.assertEqual(m.counter,8)
        self.assertEqual(m.status,'active')
        m.jump(7,(120,96,4,0,0,0),0)
        self.assertEqual(m.status,'resource_exit')
        self.assertFalse(m.report()['event_verdict'])
        off=OperatingMonitor(p,False)
        off.advance(5)
        self.assertEqual(off.status,'active')
        off.jump(6,(111,100,0,0,0,0),0)
        self.assertTrue(off.report()['event_verdict'])  # deliberately conservative disabled event

    def test_budget_limited_ssa_is_unfinished(self):
        run=BindingReactor().stochastic(limit=3)
        self.assertEqual(run['status'],'unfinished')
        self.assertIsNone(run['event_verdict'])
        self.assertLess(run['time'],1)
        self.assertTrue(all(n >= 0 for n in run['counts']))

    def test_deterministic_reproduction_and_resources(self):
        for enabled in (True,False):
            model=BindingReactor(enabled=enabled)
            solution=model.deterministic()
            sample=solution.sol(np.linspace(0,1000,101))
            self.assertLess(np.max(np.abs(np.array(RESOURCE_A)@sample[:6]-1)),1e-9)
            self.assertLess(np.max(np.abs(np.array(RESOURCE_B)@sample[:6]-1)),1e-9)
            export=solution.sol(1000)[6]-solution.sol(500)[6]
            if enabled:
                self.assertTrue(700 < export < 800)
                self.assertTrue(8 < solution.t_events[0][0] < 11)
            else:
                self.assertAlmostEqual(export,4e-6,delta=1e-10)
                self.assertEqual(len(solution.t_events[0]),0)


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