"""Scientific checks, not a replacement for the general manuscript proofs."""
import unittest
import numpy as np
import sympy as sp
from example import (SquareSource, RecyclingFamily, MassActionRealization,
                     ReverseBudget, rooted_example, vector)


class ScientificChecks(unittest.TestCase):
    def test_source_cone_graph_and_scope(self):
        for m in range(2, 9):
            family = RecyclingFamily(m); s = family.source
            self.assertEqual(s.g, sp.ones(2, 1))
            self.assertEqual(s.T, sp.Matrix([[0, 2], [sp.Rational(1, m), sp.Rational(m-1, m)]]))
            self.assertTrue(s.uniqueness_certificate()['certified'])
            for r in ('101/100', '3/2', '199/100'):
                model = family.realize(r)
                self.assertEqual(model.profile.u, sp.Matrix([1, sp.Rational(r)]))
                self.assertEqual(model.profile.phi, family.capacity+(1-sp.Rational(1, m))*(sp.Rational(r)-1))
                self.assertEqual(s.S*model.currents([1, 1]), sp.Matrix([1, 0]))
        for r in (1, 2, '1/2'):
            with self.assertRaises(ValueError): RecyclingFamily().realize(r)
        with self.assertRaises(ValueError): SquareSource([[1, 0], [0, 1]], [[0, 1], [1, 0]])
        with self.assertRaises(ValueError): RecyclingFamily(1)

    def test_global_bottleneck_and_full_stationary_curve(self):
        model = RecyclingFamily().realize('3/2')
        rng = np.random.default_rng(2201)
        saw_above_one = False
        for state in np.exp(rng.uniform(-2, 2, (100, 2))):
            b = model.bottleneck(state)
            self.assertGreaterEqual(b['delta'], -1e-12)
            self.assertGreaterEqual(b['gap'], -1e-10)
            self.assertAlmostEqual(b['gap'], b['convex_gap']+b['coupling_gap'], delta=1e-9*max(1, b['gap']))
            saw_above_one |= max(b['normalized_currents']) > 1
        self.assertTrue(saw_above_one)  # Only the selected reaction is bounded off the stationary fibre.
        for m in (2, 3, 5):
            family = RecyclingFamily(m); model = family.realize('3/2')
            curve = family.stationary_curve(model, np.geomspace(.001, 10, 71))
            for x, y, j in curve:
                field = model.field([x, y])
                self.assertAlmostEqual(field[1], 0, delta=1e-8*max(1, y**m))
                self.assertAlmostEqual(field[0], j, delta=1e-8*max(1, y**m))
                self.assertLessEqual(j, 1+1e-10)
        with self.assertRaises(ValueError): model.bottleneck([0, 1])

    def test_exact_budget_primal_dual_and_closed_optimum(self):
        for m in (2, 3, 6):
            family = RecyclingFamily(m)
            for budgets in ((2, 6), (3, 20)):
                for J in ('1/10', '1', '6/5', '3'):
                    c = ReverseBudget(family.source, budgets).certificate(J)
                    closed = family.budget_optimum(J, budgets)
                    self.assertEqual(c['feasible'], closed['feasible'])
                    D = sp.Matrix(c['D'])
                    if c['feasible']:
                        f = vector(c['f'])
                        self.assertTrue(all(v >= 1 for v in f))
                        self.assertTrue(all(v >= 0 for v in D*f))
                        self.assertTrue(all(v <= b for v, b in zip(vector(c['reverse_one_way']), budgets)))
                        optimum = family.realize(closed['r_min'], J)
                        self.assertEqual(str(optimum.profile.phi), closed['minimum_phi'])
                    else:
                        y = vector(c['dual_y']); row = y.T*D
                        self.assertTrue(all(v >= 0 for v in y))
                        self.assertTrue(all(v <= 0 for v in row)); self.assertLess(sum(row), 0)
        exact = RecyclingFamily().budget_optimum(1, (2, 6))
        self.assertEqual((exact['r_min'], exact['r_max'], exact['minimum_phi']), ('3/2', '3/2', '7/4'))
        self.assertEqual(exact['maximum_current_numeric'], 1)

    def test_fixed_equilibrium_reconstruction(self):
        family = RecyclingFamily(); model = family.realize('3/2', 1, ('9/4', '14/9'))
        self.assertEqual(model.reference, sp.Matrix([2, 3]))
        self.assertEqual(model.kp, vector(['3/2', '7/9']))
        self.assertEqual(model.km, vector(['2/3', '1/2']))
        self.assertEqual(model.currents(model.reference), sp.ones(2, 1))
        self.assertEqual([sp.simplify(a/b) for a, b in zip(model.kp, model.km)], list(vector(['9/4', '14/9'])))
        # Thermodynamic control identity: Phi=(Xeq/X*)**N.
        zeq = (family.source.S.T.inv()*sp.Matrix([sp.log(sp.Rational(9, 4)), sp.log(sp.Rational(14, 9))])).applyfunc(sp.exp)
        self.assertEqual(sp.simplify((zeq[0]/model.reference[0])**family.source.N), model.profile.phi)
        base = family.realize('3/2')
        for normalized in ([.5, .75], [1, 1], [2, 3]):
            state = np.array(normalized)*np.array(model.reference, float).ravel()
            np.testing.assert_allclose(model.field(state), base.field(normalized), atol=1e-12)

    def test_controlled_ode_against_closed_solution(self):
        model = RecyclingFamily().realize('3/2')
        y = model.z[1]
        self.assertEqual(sp.expand(model.symbolic_field[1].subs(model.z[0], 1)+(y-1)*(7*y+3)), 0)
        for initial in (.1, .5, 2, 4):
            t, values = model.controlled_trajectory([initial])
            c = (initial-1)/(7*initial+3)
            exponential = c*np.exp(-10*t)
            expected = (1+3*exponential)/(1-7*exponential)
            np.testing.assert_allclose(values[:, 0], expected, rtol=5e-10, atol=5e-11)
        self.assertLess(sp.Rational(model.local_curvature()['J_second_log_control']), 0)

    def test_rooted_two_maxima_and_dynamical_separation(self):
        model, report = rooted_example(); s = model.source
        self.assertEqual(s.components(), ([(0, 1), (2,)], [(2,)]))
        self.assertFalse(report['uniqueness_criterion']['certified'])
        self.assertEqual(report['unit_internal_determinant'], '-15/2')
        self.assertNotEqual(sp.sympify(report['second_internal_determinant']), 0)
        self.assertEqual(report['second_phi'], '96/25')
        self.assertEqual(report['second_response_f'], ['10', '18', '-6'])
        for state in (sp.ones(3, 1), sp.Matrix([5/sp.sqrt(3), sp.sqrt(3), 1])):
            self.assertEqual(model.currents(state).applyfunc(sp.simplify), s.g)
        second = sp.Matrix([5/sp.sqrt(3), sp.sqrt(3), 1]); u = sp.Matrix([1, 17, -6])
        current_log_jac = model.symbolic_currents.jacobian(model.z)*sp.diag(*model.z)
        self.assertEqual((current_log_jac.subs(dict(zip(model.z, second)))*u).applyfunc(sp.simplify), sp.zeros(3, 1))
        # Exhaustive symbolic maximum reduction, not a grid search for roots.
        p = sp.symbols('p')
        self.assertEqual(sp.expand(2*(2*p-1)-p*p-1+(p-1)*(p-3)), 0)
        for epsilon in (sp.Rational(1, 10), sp.Rational(1, 1000)):
            profile = s.profile([1, 1, epsilon])
            self.assertEqual(profile.phi, 4*(1+epsilon)**3)
            self.assertGreater(profile.phi, 4)

    def test_capacity_tradeoff_and_rescaling_current(self):
        for m in (2, 4, 10):
            family = RecyclingFamily(m)
            for gap in (sp.Rational(1, 100), sp.Rational(1, 10000)):
                model = family.realize(1+gap, 2)
                excess = model.profile.phi-family.capacity
                self.assertEqual(model.reverse[0]/model.J0, (1-sp.Rational(1, m))/excess)
                small_current = min(1, *(q-1 for q in model.profile.q))
                bounded = family.realize(1+gap, small_current)
                self.assertTrue(all(v <= 1 for v in bounded.reverse))
                self.assertLessEqual(small_current, gap)
        # An independent square source tests the reusable generic constructor.
        source = SquareSource([[1]], [[2]])
        model = MassActionRealization(source, source.profile([1]), current=3)
        self.assertEqual(model.currents([1]), sp.Matrix([3]))
        self.assertEqual(model.local_curvature()['J_second_log_control'], '-6')


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