import unittest
from fractions import Fraction as F
import numpy as np
import sympy as sp
from reactor import *
from geometry import EquilibriumGeometry,capacity_bounds
from certificates import PublishedSource,PublishedOperatingContract,driving_ratios,load


def rates(n):return [EdgeKinetics(*([F(1)]*6)) for _ in range(n)]


class ScientificTests(unittest.TestCase):
    def test_full_network_inventories_and_analytic_jacobian(self):
        model=PublishedSource().network
        self.assertEqual((model.size,len(model.reactions),model.size-3),(34,72,31))
        np.testing.assert_array_equal(model.inventory@model.stoichiometry,0)
        x=np.linspace(.1,.4,34);v=np.linspace(-.2,.3,34);step=1e-6
        finite=(model.rhs(0,x+step*v)-model.rhs(0,x-step*v))/(2*step)
        np.testing.assert_allclose(model.jacobian(x)@v,finite,rtol=2e-8,atol=2e-7)
        chart=model.concentrations(x[model.chart_indices],model.inventory@x)
        np.testing.assert_allclose(chart,x,atol=1e-14)

    def test_exact_symmetric_drift_and_generator_aggregation(self):
        rs=rates(3);cube=PhosphorylationNetwork.symmetric_lift(rs);chain=PhosphorylationNetwork.chain(rs);A=cube.cube_aggregation()
        # Nonuniform microstates: equality is not restricted to symmetric preparations.
        counts=np.arange(1,cube.size+1);aggregated=A@counts
        np.testing.assert_allclose(A@cube.rhs(0,counts),chain.rhs(0,aggregated),atol=1e-11)
        def jumps(network,state,aggregation):
            out={}
            for r,nu in zip(network.reactions,network.stoichiometry.T):
                key=tuple(aggregation@nu);prop=r.rate*math.prod(int(state[i]) for i in r.reactants)/F(7)**(len(r.reactants)-1)
                out[key]=out.get(key,F(0))+prop
            return out
        self.assertEqual(jumps(cube,counts,A),jumps(chain,aggregated,np.eye(chain.size,dtype=int)))
        self.assertTrue(all(x[-1]==1 for x in cube.square_ratios()))

    def test_exceptional_ratio_and_positive_root_admissibility(self):
        one=PhosphorylationNetwork.symmetric_lift(rates(1));g=EquilibriumGeometry(one)
        rows=g.census([1,1,4]);self.assertEqual(len(rows),1);self.assertTrue(rows[0]['exceptional_ratio'] and rows[0]['admissible'])
        x=np.array([float(v) for v in g.state_on_curve('1',[1,1],4)])
        np.testing.assert_allclose(one.inventory@x,[1,1,4],atol=1e-12);np.testing.assert_allclose(one.rhs(0,x),0,atol=1e-12)
        source=PublishedSource();u=sp.Symbol('u');tau=[sp.Poly.from_list(row,u).as_expr() for row in load('tree_weights.json')['weights_descending_coefficients']]
        geometry=EquilibriumGeometry(source.network,tau);rows=geometry.census(source.totals)
        self.assertEqual((len(rows),sum(r['admissible'] for r in rows)),(8,7));self.assertFalse(rows[-1]['admissible'])

    def test_exact_target_kernel_and_distinct_driving_affinity(self):
        source=PublishedSource();shifted=PublishedSource(F(1,100));u=sp.Symbol('u');tau=[sp.Poly.from_list(row,u).as_expr() for row in load('tree_weights.json')['weights_descending_coefficients']]
        a,b=[EquilibriumGeometry(s.network,tau) for s in [source,shifted]]
        self.assertEqual(sp.expand(a.B-10*a.D-b.B+10*b.D),0)
        self.assertTrue(any(r[-1]!=1 for r in source.network.square_ratios()))
        self.assertEqual(len(driving_ratios(source.network)),12)
        reverse=PublishedSource(reverse_activity=F(1,10)).network
        self.assertEqual(len(reverse.reactions),96);np.testing.assert_array_equal(reverse.inventory@reverse.stoichiometry,0)
        # Every reversible pair balances at free activities one, C=a/b and Y=alpha/beta.
        x=[F(1)]*reverse.size
        for i,r in enumerate(reverse.kinetics):x[10+i]=r.association/r.dissociation;x[22+i]=r.reverse_association/r.reverse_dissociation
        unit=PublishedSource(reverse_activity=F(1)).network
        v=[r.rate*math.prod(x[k] for k in r.reactants) for r in unit.reactions]
        self.assertTrue(all(value==0 for value in unit.stoichiometry@np.array(v,dtype=object)))

    def test_operating_clock_invariance_and_composite_contract(self):
        one=PublishedOperatingContract().calculate();slow=PublishedOperatingContract().calculate(clock_multiplier=F(1,100))
        self.assertEqual(one['required_size'],slow['required_size']);self.assertEqual(one['rows'],slow['rows'])
        self.assertEqual(slow['recovery_time'],100*one['recovery_time'])
        self.assertTrue(one['sufficient_size_passes'])
        self.assertLess(max(r['composite_failure_upper'] for r in one['rows']),F(11,1000))
        small=PublishedOperatingContract().calculate(F(1,10**100));self.assertFalse(small['sufficient_size_passes'])
        self.assertTrue(any(r['composite_failure_upper']>1 for r in small['rows']))

    def test_full_ode_and_count_path_preserve_retained_substrate(self):
        model=PhosphorylationNetwork.symmetric_lift(rates(2));x=np.zeros(model.size);x[0]=10;x[model.E]=2;x[model.F]=1
        y=model.simulate(x,np.linspace(0,3,16));np.testing.assert_allclose(y@model.inventory.T,np.tile([2,1,10],(len(y),1)),atol=1e-12)
        t,counts,events=model.ssa([int(v*10) for v in x],10,3,72)
        self.assertEqual(t[-1],3);self.assertGreater(len(events),0)
        for c in counts:self.assertEqual(list(model.inventory@c),[20,10,100])
        self.assertGreater(y[-1,model.q+2:].sum(),0)
        with self.assertRaises(RuntimeError):model.ssa([int(v*10) for v in x],10,3,event_budget=1)

    def test_invalid_inputs_and_capacity_is_an_interval(self):
        self.assertEqual(capacity_bounds(3)['constructive_lower'],4);self.assertEqual(capacity_bounds(3)['universal_upper'],8)
        for fn in [lambda:EdgeKinetics(catalysis=0),lambda:PhosphorylationNetwork.cube(0,[]),lambda:PublishedSource(F(-1)),lambda:PublishedOperatingContract().calculate(0),lambda:capacity_bounds(1)]:
            with self.assertRaises(ValueError):fn()
        with self.assertRaises(ValueError):EquilibriumGeometry(PublishedSource(reverse_activity=F(1,10)).network)


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