import unittest
import numpy as np
import sympy as sp
from example import (Source, Reaction, OperatingFamily, MassAction, Quartic,
                     ChildLattice, DilutionCompletion, local_departure)


class ScientificChecks(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.family=OperatingFamily();cls.model=cls.family.model()
        cls.children=ChildLattice(cls.family.source);cls.records=cls.children.certificates()

    def test_literal_source_rate_convention_and_complete_flux_cone(self):
        s=self.family.source;m=self.model
        self.assertEqual(list(m.k),[10000,4,100,10000,99,2])
        self.assertEqual(s.molecularity(),{'reactants':[2,2,2,2,0,0],'products':[2,0,1,9,1,1],'catalyst_free':True})
        self.assertEqual(m.exact_field(m.equilibrium),sp.zeros(4,1))
        # Recover all four dependent fluxes from the literal balance equations.
        v=sp.Matrix(sp.symbols('v0:6'))
        sol=sp.solve(s.S*v,[v[1],v[2],v[4],v[5]],dict=True)
        self.assertEqual(sol,[{v[1]:4*v[3],v[2]:v[0],v[4]:v[0]-v[3],v[5]:2*v[3]}])
        flux=sp.diag(*self.family.flux(100))*s.Y.T*sp.diag(1,1,1,100)
        self.assertEqual(s.S*flux,m.jacobian())
        # Homodimer derivative includes two; event rate has no extra 1/2.
        self.assertEqual(m.symbolic_flux[2],100*m.x[2]**2)
        self.assertEqual(sp.diff(m.symbolic_flux[2],m.x[2]),200*m.x[2])
        for T in (1,0):
            with self.assertRaises(ValueError):self.family.flux(T)
        with self.assertRaises(ValueError):s.reconstruct([1]*4,[1]*6)

    def test_exact_unstable_shift_and_stable_control(self):
        q=Quartic(self.model.jacobian());r=q.classification();root=q.shifted_root_certificate()
        self.assertEqual(r['coefficients'],['10908','185600','1280000','64000000'])
        self.assertEqual(r['H'],'-5025252352000000')
        self.assertTrue(root['certified'])
        self.assertEqual(root['H_at_endpoints'],['-2133097221521408','2676722101092352'])
        stable=Quartic(self.family.model(2,1).jacobian()).classification()
        self.assertEqual(stable['status'],'strictly_stable');self.assertEqual(stable['H'],'626688')
        # Verify the algebra behind the root construction at H=0.
        a,b,c,d=sp.symbols('a b c d',nonzero=True)
        residual=c*c/a**2-b*c/a+d
        self.assertEqual(sp.simplify(residual+Quartic.hurwitz((a,b,c,d))/a**2),0)
        self.assertFalse(Quartic(self.family.model(2,1).jacobian()).shifted_root_certificate()['certified'])

    def test_all_children_and_principal_coverage(self):
        self.assertEqual(len(self.records),25)
        self.assertEqual(sum(r['singular'] for r in self.records),4)
        self.assertTrue(all(r['certified_D_nonunstable'] for r in self.records))
        maximal=self.children.maximal_patterns()
        self.assertEqual(set(maximal),{(0,1,2,3),(None,1,2,0),(0,None,1,3),(None,None,1,0)})
        for p in self.children.patterns():
            extending=next(q for q in maximal if all(r is None or r==q[i] for i,r in enumerate(p)))
            kept=[i for i,r in enumerate(extending) if r is not None]
            selected=[kept.index(i) for i,r in enumerate(p) if r is not None]
            self.assertEqual(self.children.matrix(extending).extract(selected,selected),self.children.matrix(p))
        for record in self.records:
            self.assertTrue(all(sp.Integer(t['coefficient'])>0 for t in record['hurwitz_monomials']))
        # The certificate is a sufficient test, not a fake general D-stability oracle.
        toy=Source(('A','B','C','D'),[Reaction('A -> 2A',(1,0,0,0),(2,0,0,0))])
        self.assertFalse(all(r['certified_D_nonunstable'] for r in ChildLattice(toy).certificates()))
        with self.assertRaises(RuntimeError):self.children.patterns(budget=1)

    def test_operating_family_ray_and_schur_reduction(self):
        cert=self.family.exact_family_certificate()
        self.assertEqual(cert['negative_ray_coefficients_descending'],['-5254246400','-1562660812800','-154010347520000','-5025252352000000'])
        for T,L in ((2,1),(10,3),(100,100),(100,10000)):
            model=self.family.model(T,L,s='3/2')
            base=self.family.model(T,L)
            self.assertEqual(model.jacobian(),sp.Rational(3,2)*base.jacobian())
            self.assertEqual(base.jacobian().charpoly().all_coeffs()[1:],list(self.family.coefficients(T,L)))
            J=base.jacobian();red=J[:3,:3]-J[:3,3]/J[3,3]*J[3,:3]
            self.assertEqual(red.charpoly().all_coeffs()[1:],
                             [sp.Rational(14*T+32,T+4),sp.Rational(112*T,T+4),sp.Rational(64*T*T,T+4)])
        self.assertAlmostEqual(cert['large_L_threshold_numeric'],22.94104015,places=5)

    def test_dilution_source_and_all_child_transport(self):
        completed=DilutionCompletion(self.model,1);m=completed.model
        self.assertEqual(list(completed.feed_concentrations),[100,1,1,sp.Rational(201,100)])
        self.assertEqual(m.source.m,12)
        self.assertEqual(m.jacobian(),self.model.jacobian()-sp.eye(4))
        self.assertEqual((m.symbolic_field-self.model.symbolic_field-(self.model.equilibrium-self.model.x)).applyfunc(sp.expand),sp.zeros(4,1))
        rows=completed.transport_certificates(self.records)
        self.assertEqual(len(rows),121)
        self.assertTrue(all(r['certified_D_nonunstable'] for r in rows))
        # Independent scaled determinant factorization for every completed child.
        lattice=ChildLattice(m.source);old=ChildLattice(self.model.source);z=sp.Symbol('z')
        for record in rows:
            p=tuple(record['assignment']);I=[i for i,r in enumerate(p) if r is not None]
            scales={i:sp.Rational(i+2,i+1) for i in I}
            C=lattice.matrix(p)*sp.diag(*[scales[i] for i in I]) if I else sp.zeros(0,0)
            op=tuple(record['old_assignment']);oldI=[i for i,r in enumerate(op) if r is not None]
            O=old.matrix(op)*sp.diag(*[scales[i] for i in oldI]) if oldI else sp.zeros(0,0)
            expected=O.charpoly(z).as_expr()*sp.prod(z+scales[i] for i in record['loss_species'])
            self.assertEqual(sp.expand(C.charpoly(z).as_expr()-expected),0)
        for delta in ('1/10','199/100'):
            self.assertEqual(Quartic(DilutionCompletion(self.model,delta).model.jacobian()).classification()['status'],'unstable_complex_pair')

    def test_rate_perturbation_equilibria_and_mass_obstruction(self):
        for j in range(6):
            rates=self.model.k.copy();rates[j]*=sp.Rational(1001,1000)
            model=self.family.from_rates(rates)
            self.assertEqual(model.exact_field(model.equilibrium).applyfunc(sp.simplify),sp.zeros(4,1))
            self.assertGreater(max(np.linalg.eigvals(np.array(model.jacobian(),float)).real),2.9)
        self.assertEqual(self.family.from_rates([2,4,2,1,1,2]).equilibrium,sp.ones(4,1))
        # Independent arbitrary concentration reconstruction round trip.
        state=sp.Matrix([2,3,5,sp.Rational(1,7)])
        custom=self.family.source.reconstruct(state,self.family.flux('7/2','2/3'))
        self.assertEqual(self.family.from_rates(custom.k).equilibrium,state)
        mass=sp.Matrix(sp.symbols('mA mB mC mD'))
        obstruction=(mass.T*(2*self.family.source.S[:,2]+self.family.source.S[:,3]))[0]
        self.assertEqual(obstruction,mass[0]+4*mass[1])
        with self.assertRaises(ValueError):self.family.from_rates([0]*6)

    def test_local_numerical_departure(self):
        times,dev,linear,report=local_departure(self.model)
        self.assertAlmostEqual(report['eigenvalue_real'],2.99974676,places=7)
        self.assertAlmostEqual(report['eigenvalue_imag'],15.68952573,places=7)
        self.assertLess(report['relative_linear_error'],1e-3)
        self.assertLess(report['independent_solver_error_relative_to_mode'],1e-3)
        self.assertGreater(np.max(np.linalg.norm(dev,axis=1)),5*np.linalg.norm(dev[0]))
        np.testing.assert_allclose(dev[0],linear[0],atol=1e-16)


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