import unittest,json
from pathlib import Path
import numpy as np
import sympy as sp
from clock import finite_model,RATE_NAMES
from reactor import MassActionReactor,ResonantPulse,ZeroSignal,SwitchingExperiment
from local_certificate import certify
import input_rank

class ScientificChecks(unittest.TestCase):
    def test_exact_equilibrium_and_conservation_chart(self):
        m=finite_model();self.assertTrue(all(v==0 for v in m.field()));self.assertEqual(m.Rm*m.P,sp.eye(9))
        c=sp.zeros(3,12)
        for row,indices in enumerate([[4,6,7,8],[5,9,10,11],[0,1,2,3,6,7,8,9,10,11]]):
            for j in indices:c[row,j]=1
        self.assertEqual(c*m.N,sp.zeros(3,18));self.assertEqual(c*m.P,sp.zeros(3,9))
    def test_all_eighteen_inputs_change_one_flux(self):
        m=finite_model();base=MassActionReactor(m);y=np.zeros(9);y[0]=.001;x=base.species(y);v=base.flux(x)
        for j,name in enumerate(RATE_NAMES):
            reactor=MassActionReactor(m,input_rate=name);change=reactor.flux(x,.1)-v
            expected=np.zeros(18);expected[j]=.1*v[j];np.testing.assert_allclose(change,expected,atol=1e-14,rtol=2e-14)
            np.testing.assert_allclose(reactor.f(0,y,.1)-base.f(0,y),reactor.RN[:,j]*.1*v[j],atol=2e-14)
    def test_nonlinear_jacobian_and_moved_equilibrium(self):
        m=finite_model();reactor=MassActionReactor(m);y=np.zeros(9);y[2]=.03;h=1e-6
        numeric=np.column_stack([(reactor.f(0,y+np.eye(9)[j]*h,.07)-reactor.f(0,y-np.eye(9)[j]*h,.07))/(2*h) for j in range(9)])
        np.testing.assert_allclose(reactor.jac(0,y,.07),numeric,atol=2e-7,rtol=2e-7)
        rates=reactor.rates.copy();rates[9]*=1.00001;changed=MassActionReactor(m,rates);eq,res=changed.equilibrium();self.assertLess(res,1e-9);self.assertGreater(np.linalg.norm(eq),1e-8)
        np.testing.assert_allclose(changed.conservation@changed.species(eq),changed.conservation@changed.x0,atol=1e-13)
    def test_pulse_support_positive_rate_and_state_response(self):
        signal=ResonantPulse(.1,28.7,86.1,np.pi);self.assertEqual(signal.value(-1),0);self.assertEqual(signal.value(86.1),0)
        self.assertLessEqual(max(abs(signal.value(t)) for t in np.linspace(0,86.1,1001)),.1)
        base=MassActionReactor(finite_model());controlled=base.controlled(signal)
        self.assertGreater(max(abs(controlled.f(10,np.zeros(9)))),1e-5);self.assertLess(max(abs(base.f(10,np.zeros(9)))),1e-12)
        run=controlled.integrate(np.zeros(9),1,samples=21);self.assertLess(run['conservation_drift'],1e-12)
        with self.assertRaises(ValueError):ResonantPulse(1,1,1,0)
    def test_fresh_generalized_hopf_and_rank(self):
        cert=certify();self.assertTrue(cert['strict_inclusion']);self.assertTrue(cert['contraction']);self.assertTrue(cert['complement_hurwitz']);self.assertTrue(input_rank.certify(cert)['all_eighteen_rates'])
    def test_same_source_has_numeric_sink_and_attracting_cycle(self):
        reactor=MassActionReactor(finite_model());seeds=json.loads((Path(__file__).parent/'orbit_seeds.json').read_text());o=reactor.shoot(seeds['stable_y0'],seeds['stable_T'])
        self.assertLess(max(np.linalg.eigvals(reactor.jac(0,np.zeros(9))).real),0);self.assertLess(o['largest_nontrivial'],.98);self.assertGreater(o['largest_nontrivial'],.88);self.assertLess(o['residual'],1e-8)
        run=reactor.integrate(o['anchor'],o['period'],samples=301,ledger=True);self.assertGreater(np.ptp(run['species'][3]+run['species'][11]),2.34);self.assertLess(abs(sum(run['ledger'][[2,5,8]])-sum(run['ledger'][[11,14,17]])),1e-6)
    def test_phase_trigger_and_invalid_domain(self):
        reactor=MassActionReactor(finite_model());seed=json.loads((Path(__file__).parent/'orbit_seeds.json').read_text());experiment=SwitchingExperiment(reactor,seed['stable_T'],seed['XI'],seed['ZETA'])
        y,runs=experiment.phase_trigger(seed['stable_y0']);self.assertGreater(np.array(seed['XI'])@y,0);self.assertLess(abs(np.array(seed['ZETA'])@y),1e-8)
        with self.assertRaises(ValueError):MassActionReactor(finite_model(),rates=[1]*17)
        with self.assertRaises(ValueError):reactor.integrate(np.ones(9)*100,1)

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