import unittest
from fractions import Fraction as F
from math import log
import numpy as np
from branching import SisterKernel,FamilySource,RetainedBranchExperiment,Phase,A,B
from certificates import ReferenceCertificate,boundary,pulse_extinction,constant_extinction
from assay import binomial_weights,tail,SymmetricReadout,PairedDecisionRule,binomial_confidence_interval,project_covariance,covariance_sign,minimax_regret
from example import symbolic_checks


class ScientificTests(unittest.TestCase):
    def test_joint_law_is_hidden_from_retained_records(self):
        models=[FamilySource(SisterKernel(F(1,2),c),F(1,30)) for c in [F(-1,5),F(1,5)]]
        self.assertEqual(models[0].observation_signature([A,B]),models[1].observation_signature([A,B]))
        self.assertEqual(models[0].mean_operator(A),models[1].mean_operator(A))
        seq=[Phase(A,1),Phase(B,1)]
        self.assertEqual(RetainedBranchExperiment(models[0]).sample(seq,20),RetainedBranchExperiment(models[1]).sample(seq,20))
        self.assertNotEqual(models[0].founder.generating_function(.2,.7),models[1].founder.generating_function(.2,.7))
        for p in [F(3,5),F(2,5)]:
            for c in [F(-1,10),F(1,10)]:self.assertEqual(SisterKernel(p,c).retained_marginal(),(p,1-p))

    def test_exact_boundary_and_backward_order(self):
        symbolic_checks();cp,slope=boundary(F(7,8));self.assertEqual(cp,F(8967219299,65423564404))
        h=log(8/7)
        for c in [F(-1,5),F(1,5)]:
            m=FamilySource(SisterKernel(F(1,2),c))
            ab=m.extinction([Phase(A,h),Phase(B,h)])[0];ba=m.extinction([Phase(B,h),Phase(A,h)])[0]
            self.assertAlmostEqual(ab,float(pulse_extinction(F(7,8),c,'AB')),places=12)
            self.assertAlmostEqual(ba,float(pulse_extinction(F(7,8),c,'BA')),places=12)
            self.assertEqual(pulse_extinction(F(7,8),c,'AB')-pulse_extinction(F(7,8),c,'BA'),slope*(cp-c))
        self.assertGreater(boundary(F(4,5))[0],F(1,4))

    def test_recurring_division_reserves_and_refinement(self):
        for x,c,e in [(F(7,8),F(1,5),F(1,30)),(F(7,8),F(1,4),F(1,16)),(F(49,50),F(1,4),F(1))]:
            r=ReferenceCertificate(x,e,arbitrary_descendants=True);h=-log(float(x));band=r.band('pulses')
            for cc in [-c,c]:
                self.assertEqual(band.decide(cc),'AB' if cc<0 else 'BA')
                m=FamilySource(SisterKernel(F(1,2),cc),e,[SisterKernel(F(3,5),F(-1,10)),SisterKernel(F(2,5),F(1,10))])
                ab=m.extinction([Phase(A,h),Phase(B,h)])[0];ba=m.extinction([Phase(B,h),Phase(A,h)])[0];lo,hi=band.low_minus_high_extinction_interval(cc)
                self.assertLessEqual(float(lo),ab-ba+1e-12);self.assertGreaterEqual(float(hi),ab-ba-1e-12)
                self.assertLessEqual(ab,float(pulse_extinction(x,cc,'AB'))+1e-12)
        with self.assertRaises(ValueError):ReferenceCertificate(F(3,4),arbitrary_descendants=True)

    def test_exact_assay_cutoffs_and_error_control(self):
        r=ReferenceCertificate()
        for task,n,cuts in [('constant',200,(106,137)),('pulses',800,(496,584)),('pulses',1600,(1009,1153))]:
            rule=PairedDecisionRule(r.assay_band(task),n);self.assertEqual((rule.low_cut,rule.high_cut),cuts)
            self.assertLessEqual(tail(n,rule.p_minus,rule.low_cut),F(1,40));self.assertGreater(tail(n,rule.p_minus,rule.low_cut+1),F(1,40))
            self.assertLessEqual(tail(n,rule.p_plus,rule.high_cut,True),F(1,40));self.assertGreater(tail(n,rule.p_plus,rule.high_cut-1,True),F(1,40))
            self.assertEqual(rule.decide((rule.low_cut+rule.high_cut)//2),'unresolved');self.assertGreater(rule.probabilities(F(1,5))['resolution'],F(19,20))
        t=r.band('constant').threshold;read=SymmetricReadout()
        self.assertLess(tail(40,read.agreement(t-F(1,10)),25,True),F(1,20));self.assertLess(tail(40,read.agreement(t+F(1,10)),24),F(1,20))

    def test_calibration_and_outward_confidence_interval(self):
        lo,hi=binomial_confidence_interval(40,25,bits=24)
        self.assertLessEqual(tail(40,lo,25,True),F(1,40));self.assertLessEqual(tail(40,hi,25),F(1,40))
        self.assertEqual(binomial_confidence_interval(4,0)[0],0);self.assertEqual(binomial_confidence_interval(4,4)[1],1)
        self.assertEqual(SymmetricReadout(0).agreement(F(1,20)),SymmetricReadout(F(1,4)).agreement(F(1,5)))
        cov=project_covariance((F(3,5),F(3,5)),(F(1,4),F(1)));self.assertEqual(cov,(F(1,20),F(1,5)))
        self.assertEqual(covariance_sign((F(2,100),F(3,100)),F(3,100)),'unresolved')
        for p in [F(0),F(1),F(1,3)]:w,d=binomial_weights(10,p);self.assertEqual(sum(w),d)

    def test_mean_operator_and_risk_signs(self):
        m=FamilySource(SisterKernel(F(1,2),F(1,5)),F(1,10));step=1e-6
        J=np.column_stack([(m.field(A,np.ones(3)+step*np.eye(3)[j])-m.field(A,np.ones(3)-step*np.eye(3)[j]))/(2*step) for j in range(3)])
        np.testing.assert_allclose(J,np.array(m.mean_operator(A),float),atol=1e-9)
        result=minimax_regret((F(-2,1000),F(1,1000)));self.assertEqual(result['probability_first'],F(2,3));self.assertEqual(result['regret'],F(1,1500))
        self.assertEqual(minimax_regret((F(1,100),F(2,100)))['probability_first'],0)
        h=log(8/7);c=F(1,5)
        for action in [A,B]:self.assertAlmostEqual(FamilySource(SisterKernel(F(1,2),c)).extinction([Phase(action,h)])[0],float(constant_extinction(F(7,8),c,action.label)),places=12)

    def test_invalid_models_are_rejected(self):
        with self.assertRaises(ValueError):SisterKernel(F(1,2),F(3,10))
        with self.assertRaises(ValueError):FamilySource(descendant_division=-1)
        with self.assertRaises(ValueError):PairedDecisionRule(ReferenceCertificate().band('pulses'),40,SymmetricReadout(F(1,2)))
        with self.assertRaises(ValueError):ReferenceCertificate(F(1,2))
        with self.assertRaises(ValueError):FamilySource().founder_configuration_extinction([Phase(A,1)],(-1,0,0))
        with self.assertRaises(ValueError):project_covariance((F(9,10),F(1)),(F(1,10),F(1,10)))


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