"""Whole-source audits plus kinetic and inference failure cases."""
import copy
from fractions import Fraction as Q
from pathlib import Path
import unittest
import numpy as np
from example import (CarrierPool,PairedRecovery,PlateauCertificate,symbolic_checks)
from metabolic_model import SourceModel,RationalAudit,LinearRelaxation,OBJECTIVES


class ScientificChecks(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.model,cls.cert,_=SourceModel.load(Path(__file__).parent);cls.audit=RationalAudit(cls.model)
        cls.primal=cls.audit.primal(cls.cert['primal'])
        cls.duals={name:cls.audit.upper(cls.cert[name],OBJECTIVES[name],medium='original' if name=='joint' else 'restricted') for name in ('joint','turnover','service')}

    def test_entire_source_primal_and_dual(self):
        self.assertEqual((len(self.model.columns),len(self.model.chem),len(self.model.aux)),(19620,2157,8254))
        self.assertEqual(self.primal['nonzero_extents'],1580)
        self.assertEqual(self.primal['M'],Q(4421041890341877,4273504273504270))
        self.assertEqual(self.primal['Q'],Q(9400187505686245981195327946603,1446075838260537275756274164820))
        self.assertEqual(self.duals['joint']['interval_cost'],Q(1283461673726143619667945455169487587,50452437500000000000000000000000000))
        self.assertEqual(self.duals['turnover']['interval_cost'],Q(81256006217701144379925098514861,12500000000000000000000000000000))
        self.assertEqual(self.duals['service']['interval_cost'],Q(2586309736197958070773983682319,2500000000000000000000000000000))
        self.assertEqual([self.duals[k]['nonzero_costs'] for k in ('joint','turnover','service')],[11,8,2])
        self.assertEqual(self.duals['joint']['reverse_charges'],{'R_CYSGLTH':4,'R_LHCYSTIN':4,'R_SSGTHRD':4})

    def test_reject_wrong_identifiers_signs_and_infeasible_primal(self):
        bad=dict(self.cert['primal']);bad['R_NaKt']=str(Q(bad['R_NaKt'])+1)
        with self.assertRaises(ValueError):self.audit.primal(bad)
        with self.assertRaises(ValueError):self.audit.primal({'invented_reaction':1})
        with self.assertRaises(ValueError):self.audit.primal(self.cert['primal'],floors={'M_proteome_budget':1})
        bad=copy.deepcopy(self.cert['turnover']);bad['chemical'][next(iter(bad['chemical']))]='-1'
        with self.assertRaises(ValueError):self.audit.upper(bad,OBJECTIVES['turnover'])
        bad=copy.deepcopy(self.model.data);bad['rows'][0]['kind']='wrong'
        with self.assertRaises(ValueError):SourceModel(bad)
        bad=copy.deepcopy(self.model.data);bad['columns'].append(bad['columns'][0])
        with self.assertRaises(ValueError):SourceModel(bad)

    def test_media_and_carrier_cancellation(self):
        masks=self.model.masks;self.assertEqual(len(masks['sulfur_exchange_imports']),56);self.assertEqual(len(masks['other_carbon_exchange_imports']),435)
        self.assertTrue(set(masks['sulfur_exchange_imports'])&set(masks['other_carbon_exchange_imports']))
        self.assertNotIn('R_EX_glc__D_e',masks['other_carbon_exchange_imports'])
        a=self.model.bounds('restricted');b=self.model.bounds('internal_blocked');j=self.model.ci['R_LHCYSTIN']
        self.assertLess(a[j][0],0);self.assertEqual(b[j][0],0)
        bypass=self.audit.bypass();self.assertEqual(len(bypass['auxiliary_net']),8)
        self.assertNotIn('M_gthrd_c',bypass['chemical_net'])

    def test_plateau_is_uniform_but_not_global_feasibility(self):
        p=PlateauCertificate(self.primal,self.duals['turnover'],self.duals['service'])
        for m in (0,Q('1.0345238'),p.M):self.assertEqual(p.query(m)['status'],'uniformly_enclosed')
        self.assertEqual(p.query((p.M+p.UM)/2)['status'],'unresolved_feasibility')
        self.assertEqual(p.query(p.UM+Q(1,10**20))['status'],'infeasible_by_exact_service_bound')
        self.assertLess(p.U-p.L,Q('0.000000193'))
        # The rounded displayed lower endpoint is too coarse for that tighter variation claim.
        self.assertGreater(p.U-Q('6.50048030'),Q('0.000000193'))
        self.assertLess(p.UM-Q('1.0345238'),Q('0.0000001'))

    def test_full_numerical_relaxation_and_peroxide_ledger(self):
        lp=LinearRelaxation(self.model);r=lp.solve(required_service=1.)
        self.assertTrue(r['finished']);self.assertAlmostEqual(r['Q'],6.5004804974,places=7)
        self.assertLess(r['max_auxiliary_residual'],1e-7);self.assertGreater(r['minimum_floor_slack'],-1e-7)
        ledger=r['peroxide_import']+r['peroxide_internal_production']-r['Q']-r['peroxide_other_consumption']-r['peroxide_export']
        self.assertAlmostEqual(ledger,r['terminal_peroxide'],places=7)
        self.assertFalse(lp.solve(required_service=2.)['finished'])

    def test_sharp_kinetic_bound_and_zero_carrier(self):
        symbolic_checks();pool=CarrierPool(2,1,1)
        expected=[.4555082374,2/3,1.088983525]
        for theta,value in zip((0,1/3,1),expected):
            self.assertAlmostEqual(pool.bound(1,theta),value,places=8)
            t,y=pool.evolve(theta,2)
            np.testing.assert_allclose(y[:,1],[pool.bound(float(ti),theta) for ti in t],atol=2e-9)
        t,y=pool.evolve(.2,3,lambda t,g:(.5+.4*np.sin(t)**2,.7))
        self.assertTrue(np.all(y[:,1]<=np.array([pool.bound(float(ti),.2) for ti in t])+1e-9))
        t,y=CarrierPool(2,1,0).evolve(0,1);np.testing.assert_array_equal(y,0)
        self.assertAlmostEqual(pool.bound(1e-10,0)/1e-20,1.,places=8)
        self.assertIsNone(pool.required_total(1,0,0))
        with self.assertRaises(ValueError):pool.evolve(2,1)

    def test_paired_identity_errors_and_common_r(self):
        p=PairedRecovery(1,.03);t,y=p.evolve(2,lambda t:.2+t/10,lambda t:.4+t/10)
        np.testing.assert_allclose(p.transform(y[:,0],y[:,1]),np.exp(-y[:,2]),atol=1e-9)
        v=PairedRecovery.reject_increase((Q('.8'),Q('.5')),(Q('.7'),Q('.6')),Q('.01'),(0,Q('.05')))
        self.assertTrue(v['rejected']);self.assertEqual(v['minimum_numerator_margin'],Q('.08'))
        v=PairedRecovery.reject_increase((Q('.8'),Q('.5')),(Q('.7'),Q('.6')),Q('.01'),(0,Q('.05')),discrepancy_sum=Q('.2'))
        self.assertFalse(v['rejected'])
        self.assertEqual(PairedRecovery.infer_integrated(Q('.01'),Q('.02'))['status'],'no_finite_upper_identification')
        self.assertEqual(PairedRecovery.infer_integrated(Q('1.2'),Q('.01'))['status'],'incompatible')
        v=PairedRecovery.infer_integrated(Q('.5'),Q('.02'));self.assertLess(v['K_lower'],np.log(2));self.assertGreater(v['K_upper'],np.log(2))
        with self.assertRaises(ValueError):PairedRecovery.transformed_error(Q('.1'),1,0,0,1)


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