import unittest
from itertools import combinations,product
import numpy as np
import sympy as sp
from example import Q,MassActionSource,ChildCertificates,HopfCrossing,RelativeReactor,LogRates,PeriodicShooter
from interval_certificate import Interval,ComplexInterval,CheckedLinearSolver,certify


class ScientificChecks(unittest.TestCase):
    def test_literal_padding_and_stationary_log_rates(self):
        source=MassActionSource.paper();skeleton=source.skeleton();padding=source.Y-skeleton.Y
        self.assertEqual(source.S,skeleton.S);self.assertEqual(source.P-skeleton.P,padding)
        self.assertEqual([(i,j,int(padding[i,j])) for i in range(4) for j in range(5) if padding[i,j]],[(2,2,371),(3,0,398)])
        self.assertEqual(max(sum(source.Y[:,j]) for j in range(5)),401)
        for t in (.5,1,1.2):
            reactor=RelativeReactor.family(source,t);recovered=RelativeReactor.from_log_rates(reactor.log_rates)
            np.testing.assert_allclose(recovered.state,reactor.state,rtol=3e-13)
            self.assertLess(reactor.log_rates.values[2],-2000)
            self.assertAlmostEqual(np.linalg.det(reactor.jacobian),6*t/25,places=10)
        arbitrary=LogRates((-.3,.8,-1.2,.4,.9));logs=arbitrary.equilibrium_log();fluxlogs=np.array(arbitrary.values)+np.array(source.Y,float).T@logs
        np.testing.assert_allclose(fluxlogs,arbitrary.values[4]+np.log([1,1.5,1,1,1]),atol=1e-12)

    def test_positive_circuit_and_all_deleted_subnetworks(self):
        source=MassActionSource.paper();result=source.structural_certificate()
        self.assertEqual(result['four_column_minors'],['48','-72','48','-48','48'])
        for size in range(1,5):
            for columns in combinations(range(5),size):self.assertEqual(source.S[:,columns].rank(),size)
        for deleted,w in enumerate(result['deletion_covectors']):
            w=sp.Matrix([sp.Rational(x) for x in w]);rates=sp.Matrix([0 if j==deleted else j+1 for j in range(5)])
            self.assertEqual((w.T*source.S*rates)[0],sum(rates));self.assertGreater(sum(rates),0)
            self.assertEqual((w.T*source.S*sp.zeros(5,1))[0],0)

    def test_all_children_and_skeleton_identity(self):
        source=MassActionSource.paper();a=ChildCertificates(source).evaluate();b=ChildCertificates(source.skeleton()).evaluate()
        self.assertEqual(a,b);self.assertEqual(a['child_count'],24)
        self.assertIn('lambda*',a['covers'][-1]['scaled_characteristic_factorization'])
        self.assertEqual(a['covers'][4]['all_principal_minors'],['8','8','0'])

    def test_crossing_and_inertia(self):
        source=MassActionSource.paper();crossing=HopfCrossing();exact=crossing.exact_certificate(source)
        self.assertTrue(exact['identities_zero']);lo,hi=crossing.isolate();self.assertEqual(hi-lo,Q(1,2**111))
        tH,omega,alpha=crossing.numeric();self.assertAlmostEqual(tH,.956453654736546,places=14);self.assertGreater(alpha,0)
        for t,unstable in ((.5,0),(1,2)):
            values=np.linalg.eigvals(RelativeReactor.family(source,t).jacobian);self.assertEqual(sum(values.real>0),unstable)
        eig=np.linalg.eigvals(RelativeReactor.family(source,tH).jacobian);self.assertLess(min(abs(eig-1j*omega)),1e-12)
        with self.assertRaises(ArithmeticError):crossing.exact_certificate(source.skeleton())

    def test_falling_factorial_jets_bound_to_literal_field(self):
        source=MassActionSource.paper();z=sp.symbols('z1:5');inv=[Q(1,2),Q(1,600),Q(1,1500),Q(1,8)]
        monomials=sp.Matrix([sp.prod(z[i]**source.Y[i,j] for i in range(4)) for j in range(5)])
        field=sp.diag(*inv)*source.S*sp.diag(*source.flux)*monomials
        for order in (1,2,3):
            jet=source.jet(order,inv)
            for indices in product(range(4),repeat=order):
                direct=field
                for i in indices:direct=direct.diff(z[i])
                direct=direct.subs(dict.fromkeys(z,1))
                self.assertEqual(list(direct),jet[indices])

    def test_outward_intervals_and_lyapunov_certificate(self):
        a=Interval(Q(-2,3),Q(4,5));b=Interval(Q(2,7),Q(8,9))
        for operation in (lambda x,y:x+y,lambda x,y:x-y,lambda x,y:x*y,lambda x,y:x/y):
            interval=operation(a,b)
            for x in (a.lo,(a.lo+a.hi)/2,a.hi):
                for y in (b.lo,(b.lo+b.hi)/2,b.hi):self.assertLessEqual(interval.lo,operation(x,y));self.assertGreaterEqual(interval.hi,operation(x,y))
        with self.assertRaises(ArithmeticError):a.reciprocal()
        with self.assertRaises(ArithmeticError):CheckedLinearSolver().solve([[0]],[1])
        cert=certify(MassActionSource.paper(),HopfCrossing());lo=Q(cert['l1_unit_norm']['lower']);hi=Q(cert['l1_unit_norm']['upper'])
        self.assertLess(Q(-23,1000),lo);self.assertLess(hi,Q(-22,1000));self.assertEqual(len(cert['checked_pivots']),14)
        self.assertTrue(all(Q(p['modulus_squared_lower'])>0 for p in cert['checked_pivots']))
        with self.assertRaises(ArithmeticError):certify(MassActionSource.paper().skeleton(),HopfCrossing())

    def test_shooting_refinement_continuous_return_and_means(self):
        source=MassActionSource.paper();shoot=PeriodicShooter(source,.02);periods=[]
        for steps in (128,256,512):
            variables,trace,mids,residual=shoot.solve_midpoint(steps);periods.append(variables[3]);self.assertLess(max(abs(residual)),2e-9)
            self.assertGreater(variables[2],shoot.tH)
            reactor=RelativeReactor.family(source,variables[2]);ratios=np.exp(np.log1p(.02*mids)@reactor.Y)
            np.testing.assert_allclose(ratios.mean(axis=0),1,atol=1e-12)
        self.assertAlmostEqual((periods[0]-periods[1])/(periods[1]-periods[2]),4,places=2)
        variables,times,trace,residual,difference=shoot.continuous(variables);self.assertLess(max(abs(residual)),2e-8);self.assertLess(difference,2e-9)
        eta=.02*trace;self.assertGreater(np.min(1+eta),.97);self.assertGreater(np.ptp(eta[:,0]),.03)
        mean=np.trapezoid(eta[:,3],times)/variables[3];variance=np.trapezoid((eta[:,3]-mean)**2,times)/variables[3]
        self.assertLess(mean,0);self.assertLess(abs(2*mean+mean*mean+variance),1e-14)


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