"""Scientific checks independent of saved outputs."""
import unittest
from fractions import Fraction as Q
import numpy as np
import sympy as sp
from example import (Parameters,BalancedRealization,SupportingCertificate,CatalyticAttachment,
    currents,slice_certificate,PARAMETERS,EXACT_REFERENCE,COMPLEXES)

class ScienceTests(unittest.TestCase):
    def test_literal_polynomial_and_reconstruction(self):
        p=Parameters(*PARAMETERS);x=tuple(map(Q,EXACT_REFERENCE));network=p.literal()
        self.assertEqual(len(network.reactions),13);self.assertEqual(network.exact_field(x),(0,0,0,0))
        for y in [(Q(2),Q(1,3),Q(3,5),Q(7,8)),x]:
            A,B,z,H=y;expected=(p.a-2*A+B*z+2*p.e*(B-A*A),p.b+A-(1+z)*B-p.e*(B-A*A),A-B*z-p.u*z-2*p.v*z*z+3*H,p.u*z+p.v*z*z-(2+p.d)*H)
            self.assertEqual(network.exact_field(y),expected)
        self.assertEqual(p.reduced(x[2])[0],x)
        self.assertEqual(p.reduced(x[2])[2],0)

    def test_exact_parameter_boundaries_and_root_enclosure(self):
        x=('2/5','1/5','1/5','2/25')
        for eta,inside in [('499999999999999999/100000000000000000',True),('5',True),('500000000000000001/100000000000000000',False)]:
            p=Parameters.slice(eta);self.assertEqual(p.locus()['inside'],inside)
            self.assertEqual(p.literal().exact_field(x),(0,0,0,0))
        p=Parameters('7/20','43/40',2,1,'1/10','1/2');self.assertEqual(p.locus()['E_z0'],0)
        self.assertTrue(p.locus()['inside']);self.assertEqual(p.admissible_root(),(Q(1,2),Q(1,2)))
        self.assertEqual(currents(p,('1/2',1,'1/2','1/2'))['K'],0)
        BalancedRealization(p,('1/2',1,'1/2','1/2'))
        p=Parameters(*PARAMETERS);lo,hi=p.admissible_root(80)
        self.assertLessEqual(lo,Q(5,8));self.assertGreaterEqual(hi,Q(5,8));self.assertLess(hi-lo,Q(1,10**20))
        rejected=Parameters('1/10','1/10',5,1,'1/2','1/10');self.assertFalse(rejected.locus()['inside'])
        with self.assertRaises(ValueError):rejected.admissible_root()

    def test_supporting_planes_and_stationary_pairings(self):
        p=Parameters(*PARAMETERS);x=tuple(map(Q,EXACT_REFERENCE));rng=np.random.default_rng(39)
        extras=[tuple(map(int,row)) for row in rng.integers(0,30,(25,4))]
        for kind,zero_count in [('K',45),('J',29)]:
            c=SupportingCertificate(kind);s=c.extension();self.assertEqual(sum(s[i][j]==0 for i in range(8) for j in range(8) if i!=j),zero_count)
            self.assertGreaterEqual(min(map(min,c.extension(extras))),0)
            self.assertEqual(c.pairing(p,x),currents(p,x)['K'] if kind=='K' else 2*currents(p,x)['B_minus_J'])
        q=Parameters('41729/400000','115951/800000','1/100',5,5,10);y=('703/2000','9/50','19/100','19/1250')
        self.assertEqual(q.literal().exact_field(y),(0,0,0,0));self.assertEqual(SupportingCertificate('J').pairing(q,y),Q(-81791,400000));self.assertFalse(q.locus()['inside'])
        with self.assertRaises(ValueError):BalancedRealization(q,y)

    def test_full_field_equality_flux_balance_and_negative_current(self):
        fixtures=[(Parameters(*PARAMETERS),EXACT_REFERENCE),(Parameters.slice(5),('2/5','1/5','1/5','2/25'))]
        p,x=Parameters.from_stationary_design(1,'1/2','19/10','10/19','10/361','1/10');self.assertLess(currents(p,x)['J'],0);fixtures.append((p,x))
        for p,x in fixtures:
            witness=BalancedRealization(p,x)
            self.assertEqual(witness.network.coefficients(),p.literal().coefficients());self.assertFalse(any(witness.network.balance(x).values()))
            self.assertEqual(len(witness.fluxes),19)
        boundary=BalancedRealization(Parameters.slice(5),('2/5','1/5','1/5','2/25'))
        self.assertEqual(boundary.fluxes[6,0],0);self.assertEqual(boundary.fluxes[1,7],0);self.assertEqual(boundary.fluxes[4,3],0)
        with self.assertRaises(ValueError):BalancedRealization(Parameters(*PARAMETERS),('.50000001','1/2','5/8','13/48'))

    def test_permanence_identity_and_slice_hurwitz(self):
        p=Parameters(*PARAMETERS);bounds=p.permanence_bounds();self.assertGreater(bounds['eventual_lower'],0)
        rng=np.random.default_rng(1939)
        for x in np.exp(rng.uniform(-5,3,(50,4))):
            A,B,z,H=x;L=(float(p.e)+1.5)*A+(2*float(p.e)+1)*B+z+1.5*H
            ld=np.array([float(p.e)+1.5,2*float(p.e)+1,1,1.5])@p.literal().field(x)
            self.assertLessEqual(ld,float(bounds['C'])-float(bounds['rho'])*L+1e-10)
        cert=slice_certificate();self.assertEqual(cert['boundary_energy_minors'],['25','2199/16','55245/16','31001025/256'])
        for eta in ('9/2','5','11/2','9'):
            p=Parameters.slice(eta);self.assertLess(max(np.linalg.eigvals(p.literal().jacobian(np.array([.4,.2,.2,.08]))).real),0)

    def test_entropy_and_equivalent_trajectories(self):
        p=Parameters(*PARAMETERS);w=BalancedRealization(p,EXACT_REFERENCE);rng=np.random.default_rng(239)
        for x in np.exp(rng.uniform(-3,2,(100,4))):self.assertLess(w.entropy_derivative(x),0)
        times=np.linspace(0,15,51);x=(.1,1.2,.2,.8);a=p.literal().integrate(x,times);b=w.network.integrate(x,times,'BDF')
        self.assertLess(np.max(abs(a-b)),5e-9);self.assertTrue(np.all(np.diff([w.entropy(y) for y in a])<=1e-12))

    def test_private_attachments_and_exact_balance(self):
        p=Parameters(*PARAMETERS);w=BalancedRealization(p,EXACT_REFERENCE)
        attachments=[CatalyticAttachment('W',(1,0,1,0),2,3),CatalyticAttachment('V',(0,0,0,0),1,2)]
        literal=CatalyticAttachment.compose(p.literal(),attachments);realized=CatalyticAttachment.compose(w.network,attachments)
        self.assertEqual(literal.coefficients(),realized.coefficients());self.assertFalse(any(realized.balance(w.state+(Q(2,3),Q(1,2))).values()))
        times=np.linspace(0,5,41);initial=list(map(float,w.state))+[.1,2.];states=literal.integrate(initial,times)
        self.assertLess(np.max(abs(states[:,:4]-np.array(w.state,float))),1e-10)
        expected=2/3+(.1-2/3)*np.exp(-3*float(w.state[0]*w.state[2])*times)
        self.assertLess(np.max(abs(states[:,4]-expected)),2e-9)
        self.assertLess(np.max(abs(states[:,5]-(.5+1.5*np.exp(-2*times)))) ,2e-9)

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