import itertools
import unittest
from dataclasses import replace
import sympy as sp
from example import (Q, Box, PairCore, PairAssembly, ShortcutFamily, Branch,
                     Fan, reference_fan, robustness_certificate)


class ScientificChecks(unittest.TestCase):
    def test_pair_response_identity_and_productivity(self):
        x = sp.Symbol('x')
        f = lambda z:(z+z*z)/2
        g = lambda z:(z+2*z*z)/3
        self.assertEqual(sp.expand(f(f(x))-g(x)-x*(x-1)*(3*x*x+9*x+2)/24),0)
        core = PairCore(0,1)
        for a,b in itertools.product((Q(i,20) for i in range(1,31)),repeat=2):
            self.assertEqual(min(core.residuals((a,b))) > 0,g(a) < b < f(a))

    def test_shortcuts_and_global_relaxations(self):
        for length in range(2,15):
            source = ShortcutFamily(length)
            report = source.report()
            self.assertGreater(report['minimum_forward_current'],0)
            self.assertGreater(report['minimum_relaxed_production'],0)
            self.assertIsNone(source.assembly.graded_state())
            for deletion in report['single_deletions']:
                self.assertTrue(all(Q(9,10) <= a <= 1 for a in deletion['state']))
                self.assertGreater(deletion['minimum_production'],0)
        report = ShortcutFamily(2).report()
        self.assertEqual(report['single_deletions'][2]['state'],(Q(251,256),Q(31,32),Q(19,20)))

    def test_graded_diamond_disconnected_and_ring(self):
        edges = tuple(PairCore(s,t) for s,t in ((0,1),(0,2),(1,3),(2,3)))
        state = PairAssembly(5,edges).graded_state()
        self.assertEqual(state[:4],(Q(251,256),Q(31,32),Q(31,32),Q(19,20)))
        for n in range(3,12):
            ring = tuple(PairCore(i,(i+1)%n) for i in range(n))
            self.assertIsNone(PairAssembly(n,ring).graded_state())
            for omitted in range(n):
                self.assertIsNotNone(PairAssembly(n,ring[:omitted]+ring[omitted+1:]).graded_state())

    def test_private_elimination_against_direct_affine_constraints(self):
        # Independent feasibility: intersect the three residual half-lines in c.
        for m,x,y,a,b in itertools.product((2,3,4),(Q(1,2),Q(3)),(Q(1),Q(2)),
                                          (Q(1,4),Q(3,4),Q(9,10)),(Q(1,3),Q(2,3),Q(4,5))):
            branch = Branch(m,x,y,Box(Q(1,5),Q(4,5)))
            j = a-b
            lo = max((x*b-j)/x,(j+m*y*a**m)/(m*y))
            hi = (x*b+y*a**m)/(x+y)
            expected = lo < hi and branch.box.lower < hi and lo < branch.box.upper
            lower,upper = branch.responses(a,Q(1))
            self.assertEqual(expected,max(lower) < b < min(upper))
            c = branch.private_state(a,b,Q(1))
            self.assertEqual(c is not None,expected)
            if c is not None:
                self.assertGreater(min(branch.residuals(a,b,c,Q(1))),0)

    def test_fan_exact_decision_and_unchanged_omitted_boxes(self):
        fan = reference_fan()
        for k in range(4):
            for selected in itertools.combinations(range(3),k):
                result = fan.decide(selected)
                self.assertEqual(result['status'],'incompatible' if k == 3 else 'compatible')
                if k < 3:
                    self.assertTrue(fan.audit(result['state'],selected))
        self.assertEqual(fan.decide(max_degree=1)['status'],'not_computed')
        with self.assertRaises(ValueError):
            fan.decide((0,0))

    def test_closed_point_boxes_and_strict_boundaries(self):
        branch = Branch(2,Q(1),Q(1),Box(Q(1,5),Q(4,5)))
        fan = Fan(Q(1),Box(Q(4,5),Q(4,5)),Box(Q(73,100),Q(73,100)),(branch,))
        result = fan.decide()
        self.assertEqual(result['status'],'compatible')
        a,b,c = result['state']
        self.assertTrue(replace(fan,branches=(replace(branch,box=Box(c,c)),)).audit((a,b,c)))
        boundary = (a*a+(a-b)/2)
        bad = replace(fan,branches=(replace(branch,box=Box(boundary,boundary)),))
        self.assertEqual(bad.decide()['status'],'incompatible')
        self.assertIsNone(Box(Q(1),Q(2)).intersect_open(Q(2),Q(3)))

    def test_uniform_robustness_certificate(self):
        result = robustness_certificate()
        self.assertEqual(result['deletion_reference_margins'],
                         [Q(311253,3276800000),Q(10809,8000000),Q(429509,25000000)])
        self.assertEqual(result['independent_complex_reference_margin'],Q(201,100000))
        self.assertLess(result['full_family_production_upper_bound'],0)
        self.assertGreater(result['deletion_uniform_lower_bound'],0)
        self.assertEqual(robustness_certificate(Q(1,1000))['status'],'not_certified')


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