import unittest
from dataclasses import replace
from fractions import Fraction as F
import numpy as np
import sympy as sp
import example as e


class ScientificChecks(unittest.TestCase):
    def setUp(self):
        self.source=e.SandwichSource(e.BindingSite(1,1),e.BindingSite(1,1))
        self.certs=e.paper_certificates()

    def test_source_mass_balance_collision_and_structure(self):
        for c,d,k,j in [(1,1,1,1),(10,10,.1,.1),(.3,5,2,.01)]:
            src=e.SandwichSource(e.BindingSite(c,k),e.BindingSite(d,j))
            u=np.geomspace(1e-4,1e5,501);p=src.capture.occupancy(u);q=src.detector.occupancy(u)
            np.testing.assert_allclose(u*p+k*p/(1-p),c,rtol=1e-10)
            np.testing.assert_allclose(u*q+j*q/(1-q),d,rtol=1e-10)
            self.assertTrue(np.all(np.diff(src.log_slope(u))<0))
            self.assertTrue(np.all(np.diff(src.signal(u/7.3)/src.signal(u))>0))
            a=src.species(7);self.assertAlmostEqual(sum(a),7)
            self.assertTrue(np.all(a>=0));self.assertAlmostEqual(float(src.log_slope(src.peak())),0,places=10)
        for u,p in [(F(200,2499),F(49,100)),(F(2499,50),F(1,51))]:
            self.assertEqual(u*p+p/(1-p),1);self.assertEqual(u*p*p,F(49,2550))
        self.assertAlmostEqual(float(self.source.signal(200/2499)),float(self.source.signal(2499/50)),places=15)
        for u in np.geomspace(.001,10000,40):self.assertAlmostEqual(float(self.source.signal(u)/self.source.signal(4/u)),1,places=12)

    def test_fresh_exact_bernstein_and_endpoint_identities(self):
        t,u,d,dl,dh,A,H,c,r,m=sp.symbols('t u d dl dh A H c r m')
        Q=lambda dd:c*r*dd*u*u-(c+m*u)*(u+A*dd)*(u+H*dd)
        self.assertEqual(sp.expand((dh-dl)*Q(d)-(dh-d)*Q(dl)-(d-dl)*Q(dh)-(dh-dl)*A*H*(c+m*u)*(d-dl)*(dh-d)),0)
        for cert in self.certs.values():
            check=cert.check();self.assertTrue(check['low_sign']);self.assertTrue(check['high_margin'])
            for dd in [cert.box.dilution.lo,cert.box.dilution.hi]:
                b=cert.coefficients(dd);l,h=cert.high.lo,cert.high.hi
                poly=cert.polynomial(l+(h-l)*t,dd)
                bern=b[0]*(1-t)**3+3*b[1]*t*(1-t)**2+3*b[2]*t*t*(1-t)+b[3]*t**3
                self.assertEqual(sp.expand(poly-bern),0)
        self.assertEqual(self.certs['main'].coefficients(F(9)),[F(549739,2500),F(6249899,2500),F(23318059,2500),F(554219,2500)])
        self.assertFalse(replace(self.certs['main'],margin=F(1)).check()['high_margin'])
        for dd,want in [(9,F(3591,10000)),(11,F(431,10000))]:
            self.assertEqual(4*dd-9*(1+F(11,100)*dd)**2,want)
        self.assertEqual(F(81,100)*(F(19,20)*F(9,4)-1),F(7371,8000))

    def test_outer_enclosures_and_true_state_retention(self):
        inf=e.OuterInference(e.CalibrationBox(),e.ObservationBudget(),self.certs['main'])
        pt=e.Interval.point
        blank=inf.enclose([pt(0),pt(0)],calibrated=True)
        self.assertEqual(blank['intervals'],[(F(0),F(125,16384))]);self.assertEqual(blank['status'],'LOW_CONDITIONAL')
        self.assertEqual(inf.enclose([pt(0),pt(0)],ceiling=None,calibrated=True)['status'],'UNRESOLVED')
        self.assertEqual(inf.enclose([pt(-1),pt(0)],calibrated=True)['status'],'MODEL_OR_MEASUREMENT_INCOMPATIBLE')
        self.assertEqual(inf.enclose([pt(0),pt(0)])['status'],'UNRESOLVED')
        self.assertEqual(inf.enclose([pt(F(49,2550)),None],calibrated=True)['status'],'UNRESOLVED')
        rng=np.random.default_rng(54092026)
        for x in [.002,.08,.4,1,10,25,50,99]:
            c,d,k,j=rng.uniform(.9,1.1,4);rho=rng.uniform(.8,1);dil=rng.uniform(9,11);drift=rng.uniform(.95,1.05);gain=rng.uniform(1,5)
            src=e.SandwichSource(e.BindingSite(c,k),e.BindingSite(d,j))
            ys=src.paired(x,rho,dil,gain,drift)+rng.uniform(-.0001,.0001,2)
            result=inf.enclose(list(map(pt,ys)),depth=10,calibrated=True)
            self.assertTrue(any(float(a)<=x<=float(b) for a,b in result['intervals']))
        limited=inf.enclose([None,None],calibrated=True,cell_budget=2)
        self.assertEqual(limited['intervals'],[(F(0),F(100))]);self.assertLess(limited['depth'],14)
        # Same-gain intersection rejects disjoint required gain intervals.
        self.assertFalse(inf.feasible(F(1),F(1),[pt(100),pt(0)]))

    def test_error_budgets_extensions_and_promise(self):
        budget=e.ObservationBudget();cert=self.certs['main']
        self.assertEqual(budget.bounds(cert.margin),(F(1,5000),F(199,5000)))
        self.assertEqual(budget.report(F(1,50),cert,promised=True,calibrated=True)['status'],'MODEL_OR_PROMISE_INCOMPATIBLE')
        self.assertEqual(budget.report(F(1,10),cert,calibrated=True)['status'],'ONE_SIDED_EXCLUSIONS_ONLY')
        self.assertEqual(budget.report(F(0),cert,promised=True,calibrated=True)['status'],'LOW_UNDER_PROMISE')
        self.assertEqual(budget.report(F(1),cert)['status'],'UNRESOLVED')
        self.assertAlmostEqual(budget.gaussian_allowance(.001,.002,.01),.006614067536687622,places=12)
        rel=self.certs['relative_error'].box
        self.assertEqual(rel.drift,e.Interval(F(1843,2060),F(2163,1940)))
        self.assertEqual(e.ObservationBudget(F(1,1000),F(1,1000),F(97,100)).bounds(F(3,100)),(F(1,500),F(271,10000)))
        self.assertEqual(self.certs['availability_drift'].box.dilution,e.Interval(F(60,7),F(220,19)))

    def test_full_kinetics_product_and_attenuation(self):
        src=e.SandwichSource(e.BindingSite(1.3,.5),e.BindingSite(.8,2))
        kin=e.BindingKinetics(src,1.1,.6);times=np.linspace(0,5,61);u=7
        full=kin.trajectory(u,times);red=kin.trajectory(u,times,False)
        np.testing.assert_allclose(full,red,rtol=2e-8,atol=2e-10)
        np.testing.assert_allclose(full.sum(axis=1),u,rtol=1e-12)
        self.assertGreaterEqual(full.min(),-1e-12)
        self.assertTrue(np.all(full[:,3]+1e-10 >= kin.attenuation(u,times)*src.signal(u)))
        self.assertLessEqual(full[:,3].max(),src.signal(u)+1e-10)
        self.assertTrue(np.all(full[:,1]+full[:,3]<=1.3+1e-10))
        self.assertTrue(np.all(full[:,2]+full[:,3]<=.8+1e-10))
        slow=e.BindingKinetics(src,.011,.006)
        np.testing.assert_allclose(slow.trajectory(u,times*100),full,atol=2e-9)
        self.assertGreater((1-np.exp(-3))**2,.9)

    def test_sequential_wash_and_tail(self):
        src=e.SandwichSource(e.BindingSite(10,.1),e.BindingSite(10,.1))
        ideal=e.SequentialSource(src);carry=e.SequentialSource(src,.001)
        u=np.geomspace(.1,1e6,301);w=src.capture.captured(u);b0=ideal.signal(u);b=carry.signal(u)
        self.assertTrue(np.all(np.diff(b0)>0));self.assertTrue(np.all(b<=b0))
        self.assertTrue(np.all(b>=10/(10+.001*(u-w))*b0-1e-12))
        hi=u>10
        self.assertTrue(np.all(ideal.plateau-b0[hi]<=1/(u[hi]-10)))
        self.assertTrue(np.all(b[hi]*.001*(u[hi]-10)<=100))
        self.assertAlmostEqual(float(carry.signal(1e6)),.099001,places=6)
        self.assertEqual(e.SequentialSource.wash_limit(F(10),F(1000),F(9,10)),F(1,900))

    def test_obstructions_and_invalid_contracts(self):
        for dilution in [1,3,10,100]:
            for spike in [0,.1,100]:
                self.assertEqual(F(1,10)/dilution+F(str(spike)),F(100)*F(1,1000)/dilution+F(str(spike)))
        self.assertEqual(F(121,100)*11/F(1,10000),133100)
        for d in [1,5,11]:self.assertLess(float(self.source.signal(133100/d)),.0001)
        for constructor in [lambda:e.BindingSite(0,1),lambda:e.CalibrationBox(rho=e.Interval(0,1)),
                            lambda:e.SequentialSource(self.source,-1),lambda:e.ObservationBudget(F(-1)),
                            lambda:e.CalibrationBox().relative_error(F(1))]:
            with self.assertRaises(ValueError):constructor()


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