"""Independent scientific identities, extremal witnesses and boundary checks."""
import unittest
from fractions import Fraction as F
from dataclasses import replace
import math
import numpy as np
from scipy.integrate import quad
from kinetics import *
from population import *
from example import individual_design, asymptotic_rows, resource_counterexample


class ScientificChecks(unittest.TestCase):
    def test_01_literal_rate_bijection_and_dual_deadline(self):
        rng=np.random.default_rng(5801)
        for _ in range(30):
            beta=tuple(F(int(x),10) for x in rng.integers(1,101,6));p=EnzymeParameters.from_coefficients(beta)
            self.assertEqual(p.coefficients(),beta)
            inputs=tuple(F(int(x)) for x in rng.integers(1,100,5))
            self.assertEqual(1/p.rate(*inputs),sum(a*b for a,b in zip(features(*inputs),beta)))
            task=ScalarTask(p)
            for g in map(F,[0,10,28,56]):self.assertEqual(task.rate(g),p.rate(56-g,F(7),g))
        design=individual_design();self.assertEqual(design['dual']['reciprocal_upper'],F(243,112))
        self.assertEqual(design['deadline']['drift_lower'],F(391,2430));self.assertEqual(design['deadline']['arrival_upper'],F(43740,391))
        self.assertEqual(design['solver_proposal']['status'],'exact certificate verified')
        self.assertEqual(design['missing_inhibition_observation']['status'],'not certified')
        self.assertTrue(design['literal_rate_band_dual']['nonempty_verified'])
        sequence=design['near_boundary_sequence'];self.assertTrue(all(x['target_drift']>0 for x in sequence))
        self.assertTrue(all(x['passage_numeric']<y['passage_numeric'] for x,y in zip(sequence,sequence[1:])))
        poly=ObservationPolyhedron([features(28,7,28)],[F(3)])
        with self.assertRaises(ValueError):poly.verify_dual(features(28,7,28),[F(999,1000)])
        with self.assertRaises(ValueError):EnzymeParameters.from_coefficients([0]*6)

    def test_02_actual_trajectories_rectangle_bounds_and_no_arrival(self):
        family=CapacityFamily();self.assertEqual(family.eventual_threshold(),F(363,560))
        for V in [F(19,20),F(951,1000),F(1),F(9,8),F(11,8)]:
            task=family.at(V);lohi=task.passage_rectangles();numeric=task.passage_numeric()
            self.assertLess(float(lohi['lower']),numeric);self.assertGreater(float(lohi['upper']),numeric)
            integral=quad(lambda g:1/float(task.drift(g)),10,28,epsabs=1e-10)[0]
            self.assertAlmostEqual(numeric,integral,8)
            t,g,hit=task.trajectory()
            self.assertGreaterEqual(min(g),0);self.assertLessEqual(max(g),56)
            if numeric<120:self.assertAlmostEqual(hit,numeric,6)
            else:self.assertIsNone(hit)
        self.assertEqual(family.at(F(19,20)).classify()['status'],'failure')
        self.assertEqual(family.at(F(951,1000)).classify()['status'],'success')
        self.assertEqual(family.at(family.eventual_threshold()).classify()['status'],'failure')
        self.assertTrue(math.isinf(family.at(family.eventual_threshold()).passage_numeric()))
        self.assertEqual(family.at(F(1,10)).classify()['status'],'not certified')
        near=family.at(F(9506,10000));self.assertEqual(near.classify()['status'],'unresolved')
        self.assertAlmostEqual(family.numerical_cutoff(),.9506018463,9)

    def test_03_sharp_populations_readout_joint_witnesses_and_surface(self):
        family=CapacityFamily();bands=TwoBandClass();readout=ReadoutCalibration()
        self.assertEqual(bands.bounds(family)['interval'],(F(20,41),F(1)))
        self.assertEqual(readout.interval(bands.bounds(family)['interval']),(F(33,49),F(4,5)))
        for p in [F(20,41),F(1,2),F(33,49),F(17,25),F(7,10),F(18,25),F(4,5),F(1)]:
            pop=bands.witness(p);self.assertEqual(pop.mean,1);self.assertEqual(pop.recovery_enclosure(family)['lower'],p);self.assertEqual(pop.recovery_enclosure(family)['upper'],p)
            counts=pop.count_realization();self.assertEqual(F(sum(n for n,v in zip(counts,pop.capacities) if v>=1),sum(counts)),p)
            for inputs in [(F(28),F(7),F(28),F(0),F(0)),(F(11),F(2),F(30),F(40),F(9))]:
                self.assertEqual(pop.pooled_rate(family,inputs),family.at(1).parameters.rate(*inputs))
            if F(33,49)<=p<=F(4,5):
                w=readout.witness(p);self.assertEqual(w['successful_positive']+w['failing_positive'],w['observed'])
                self.assertEqual(w['sensitivity']*p+w['false_positive']*(1-p),w['observed'])
                self.assertTrue(F(9,10)<=w['sensitivity']<=1 and 0<=w['false_positive']<=F(1,50))
        short=CapacityFamily(replace(ScalarTask(),deadline=F(50)))
        self.assertIsNone(bands.bounds(short)['interval'])
        self.assertIsNone(replace(readout,observed=(F(0),F(0))).interval((F(20,41),F(1))))
        # Distinct realizations: reciprocal of a pooled rate is not pooled beta.
        p1=EnzymeParameters();p2=replace(p1,inhibition_H=F(2));phi=features(28,7,28)
        pooled=(p1.rate(F(28),F(7),F(28))+p2.rate(F(28),F(7),F(28)))/2
        self.assertNotEqual(1/pooled,sum(f*(b+c)/2 for f,b,c in zip(phi,p1.coefficients(),p2.coefficients())))

    def test_04_measurement_margin_weights_and_capacity_observation(self):
        bulk=(F(20,41),F(1));repair=(F(33,49),F(4,5))
        self.assertEqual(minimax(bulk),dict(midpoint=F(61,82),radius=F(21,82)))
        self.assertEqual(minimax(repair),dict(midpoint=F(361,490),radius=F(31,490)))
        self.assertEqual(convert_weights(repair,F(6,5)),(F(55,87),F(24,29)))
        for R in [F(1),F(51,50),F(33,32)]:
            e=error_normalization_frontier(R)['maximum_additional_error'];alpha=(33-50*e)/49
            self.assertEqual(convert_weights((alpha,F(4,5)),R)[0],F(2,3))
        self.assertIsNone(error_normalization_frontier(F(6,5))['maximum_additional_error'])
        pop=TwoBandClass().witness(F(7,10));e=F(1,100);measured=[v+(-1)**i*e for i,v in enumerate(pop.capacities)]
        b=resolved_capacity_bounds(pop,measured,e,(F(19,20),F(951,1000)))
        self.assertLessEqual(b['lower'],F(7,10));self.assertGreaterEqual(b['upper'],F(7,10));self.assertEqual(b['ambiguity_band'],(F(94,100),F(961,1000)))

    def test_05_connected_support_attainment_offband_and_variance(self):
        family=CapacityFamily();cutoff=family.numerical_cutoff();limits=connected_limits(cutoff)
        self.assertAlmostEqual(limits['lower'],.1163957789,9);self.assertFalse(limits['lower_attained'])
        self.assertFalse(connected_limits(F(1))['lower_attained']);self.assertTrue(connected_limits(F(9,8))['lower_attained'])
        self.assertEqual(connected_limits(F(11,8))['upper'],F(1,2));self.assertEqual(connected_limits(2)['upper'],0);self.assertEqual(connected_limits(F(5,8))['lower'],1)
        for eta in [F(0),F(1,10),F(9,10),F(1)]:
            witness=offband_witness(eta,F(19,20));self.assertEqual(witness.mean,1);self.assertLessEqual(witness.weights[1],eta)
            self.assertEqual(witness.recovery_enclosure(family)['lower'],witness.weights[2])
            self.assertGreaterEqual(float(witness.weights[2]),float(offband_infimum(eta,cutoff)['infimum'])-1e-15)
        pops=[Population((F(7,8),F(11,8)),(F(3,4),F(1,4))),Population((F(5,8),F(9,8)),(F(1,4),F(3,4)))]
        for p in pops:self.assertEqual(p.mean,1);self.assertEqual(p.variance,F(3,64))
        self.assertEqual([p.recovery_enclosure(family)['lower'] for p in pops],[F(1,4),F(3,4)])

    def test_06_row_span_null_direction_and_useless_second_assay(self):
        atoms=list(map(F,['1/20','1/10','1/5','4/5']));rows=[[1]*4,atoms,[2*x/(1+x) for x in atoms]];h=[0,0,1,1]
        cert=row_span_certificate(rows,h);self.assertFalse(cert['identified_for_all_weights'])
        z=[-98,165,-70,3]
        self.assertEqual([sum(a*b for a,b in zip(row,z)) for row in rows],[0,0,0]);self.assertEqual(sum(a*b for a,b in zip(h,z)),-67)
        endpoints=[[0,F(55,56),0,F(1,56)],[F(7,12),0,F(5,12),0]]
        for w in endpoints:self.assertEqual([sum(a*b for a,b in zip(row,w)) for row in rows],[1,F(9,80),F(7,36)])
        self.assertEqual([sum(a*b for a,b in zip(h,w)) for w in endpoints],[F(1,56),F(5,12)])
        identified=row_span_certificate(rows+[h],h);self.assertTrue(identified['identified_for_all_weights'])
        lp=finite_support_bounds(rows,[1,F(9,80),F(7,36)],h);self.assertAlmostEqual(lp['lower'],1/56,13);self.assertAlmostEqual(lp['upper'],5/12,13)
        # A particular simplex face can identify the target despite global failure.
        specific=finite_support_bounds([[1,1],[0,1]],[1,0],[1,0]);self.assertEqual(specific['lower'],specific['upper'])

    def test_07_eventual_limit_asymptotic_and_multidimensional_counterexample(self):
        family=CapacityFamily();vi=family.eventual_threshold();d=F(11,8)
        self.assertEqual((1-vi)/(d-vi),F(197,407))
        self.assertEqual(F(210,407)*vi+F(197,407)*d,1)
        rows=asymptotic_rows();self.assertGreater(rows[0]['relative_error_numeric'],.01);self.assertLess(rows[-1]['relative_error_numeric'],1e-7)
        examples=resource_counterexample();self.assertTrue(examples[0]['sampled_success']);self.assertFalse(examples[1]['sampled_success'])
        self.assertLess(F(4,17),F(1,4))
        for row in examples:self.assertLess(row['maximum_difference'],1e-10)


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