import unittest
from fractions import Fraction as F
import numpy as np
from branching import *
from example import cone_check,subtract,winner,decision_consequences
from validated import certify
from consequence_kernel import scalar_floor


class ScientificChecks(unittest.TestCase):
    def test_offspring_source_and_complete_mean_identity(self):
        a=Action(F(3,10));J=BranchingSource();I=BranchingSource(IndependentSisters())
        self.assertEqual(J.mean_matrix(a),I.mean_matrix(a))
        for model,law in [(J,'J'),(I,'I')]:
            self.assertEqual(model.polynomial(a),source.exact_source(F(3,10),F(0),F(1,10),law))
            d,B,C=model.polynomial(a)
            self.assertEqual([d[i]+sum(B[i])+sum(p for j,k,p in C[i]) for i in range(7)],[0]*7)
        for i,row in enumerate(J.pairs):
            for j,k,p in row:self.assertEqual(tuple(x+y for x,y in zip(source.STATES[j],source.STATES[k])),source.STATES[i])
        self.assertTrue(any(tuple(x+y for x,y in zip(source.STATES[j],source.STATES[k]))!=source.STATES[i] for i,row in enumerate(I.pairs) for j,k,p in row))
        with self.assertRaises(ValueError):Acquisition(F(3))

    def test_clearing_endpoint_and_event_distinctions(self):
        model=BranchingSource();terminal,slack=model.clearing_terminal()
        self.assertEqual(terminal,[F(1)]*6+[F(1,5)])
        self.assertEqual(slack,list(map(F,['.2','.2','.2','1.56','.78','1.49'])))
        phases=[Phase(Action(F(3,10))),Phase(Action(F(1,100),F(120553,10**6)))]
        eventual=1-model.compose(phases)[5];present=1-model.compose(phases,[1]*6+[0])[5]
        appeared=1-model.compose(phases,[1]*6+[0],killed_acquisition=True)[5]
        live=1-model.compose(phases,[0]*7)[5]
        self.assertLess(eventual,present);self.assertLess(present,appeared);self.assertLess(appeared,live)
        zero=BranchingSource(acquisition=Acquisition(F(0)));self.assertAlmostEqual(zero.compose(phases)[5],1,places=12)

    def test_fresh_central_reversal_and_founder_consequences(self):
        a=(F(3,10),F(0),F(10));b=(F(1,100),F(120553,10**6),F(10));q={}
        for law in ['J','I']:
            for order,ph in [('AB',[a,b]),('BA',[b,a])]:q[law+order]=tuple(map(F,certify(ph,law)['bounds'][5]))
        gaps={law:subtract(q[law+'AB'],q[law+'BA']) for law in ['J','I']}
        self.assertEqual(winner(gaps['J']),'BA');self.assertEqual(winner(gaps['I']),'AB')
        d=decision_consequences(q,gaps)
        self.assertEqual(d['founders']['J']['unique_maximum'],32);self.assertEqual(d['founders']['I']['unique_maximum'],32)
        self.assertEqual(d['symmetric_decision'],'unresolved');self.assertEqual(d['corrected_decision'],'BA')

    def test_covariance_cone_and_finite_response_identity(self):
        c=cone_check();self.assertEqual(len(c['checks']['A'])+len(c['checks']['B']),64)
        for phases in [[Phase(Action(F(3,10))),Phase(Action(F(1,100),F(120553,10**6)))],
            [Phase(Action(F(1,100),F(120553,10**6))),Phase(Action(F(3,10)))]]:
            m,_,_=mechanism(phases);self.assertLess(m['identity_residual'],1e-10)
            self.assertGreaterEqual(min(m['covariance_transport']),-1e-12)
            self.assertTrue(np.all(abs(m['finite_dependence_error'])<=m['absolute_response_bound']+1e-10))

    def test_chronology_and_mean_scope(self):
        a=Phase(Action(F(3,10)));b=Phase(Action(F(1,100),F(120553,10**6)));J=BranchingSource();I=BranchingSource(IndependentSisters())
        z=[0]*7;z[5]=1
        np.testing.assert_allclose(J.means([a,b],z),I.means([a,b],z),atol=0)
        self.assertGreater(max(abs(J.means([a,b],z)[-1]-J.means([b,a],z)[-1])),1e-4)
        nested=J.compose([a],J.compose([b]));np.testing.assert_allclose(nested,J.compose([a,b]),atol=0)

    def test_exact_feedback_scalar_enclosure(self):
        floor=scalar_floor();lo,hi=map(F,floor['bounds'])
        self.assertGreater(lo,F('.0024725782'));self.assertLess(hi,F('.0024725784'))
        rho=F(1,5);b=F(1,10);dmax=F(421,1000);mumin=F(1,100)
        for q in [rho,F(1,2),F(1)]:
            for death,mu in [(F(1,100),F(1,25)),(F(3,10),F(1,100))]:
                difference=(dmax-death)*(1-q)+b*(mu-mumin)*q*(q-rho)
                self.assertGreaterEqual(difference,0)

    def test_modified_models_do_not_inherit_reference_claims(self):
        small=BranchingSource(acquisition=Acquisition(F(0)),resistant=ResistantLineage(F(1,10),F(1,5)))
        self.assertEqual(small.resistant.extinction,1)
        self.assertEqual(ResistantLineage(F(0),F(0)).extinction,0)
        with self.assertRaises(ValueError):BranchingSource(acquisition=Acquisition(F(-1)))
        self.assertEqual(winner((-F(1,1000000),F(1,1000000))),'unresolved')


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