"""Exhaustive small oracles and independently checked positive/negative certificates."""
import copy
import itertools
import unittest
from example import (SetFamily, SupportGraph, HornTheory, IntervalRecognizer, check_trace,
                     check_rejection, hall_matching, hall_counterexample, StructuralCRS)


class ScientificChecks(unittest.TestCase):
    def test_recognizer_against_every_three_vertex_family(self):
        graphs=[SupportGraph(tuple('abc'),p) for p in itertools.product(range(8),repeat=3)]
        oracle={g.family().members for g in graphs}
        self.assertEqual(len(oracle),55)
        for mask in range(256):
            family=SetFamily(tuple('abc'),frozenset(s for s in range(8) if mask>>s&1))
            for policy in ('largest','largest_reverse'):
                result=IntervalRecognizer(family,policy).solve()
                self.assertEqual(result['status']=='success',family.members in oracle)
                self.assertLessEqual(result['unique_states'],result['unmemoized_node_bound'])
                if result['status']=='success':
                    roots=[r for s,r in result['trace']]
                    self.assertEqual(set(roots),{r for r in range(3) if 1<<r not in family.members})
                    self.assertEqual(len(roots),len(set(roots)))
                    graph=check_trace(family,result['trace']);self.assertEqual(graph.family(),family)
                    for s in range(8):self.assertEqual(graph.interior(s),family.interior(s))
                else:self.assertTrue(check_rejection(family,result)['valid'])

    def test_root_branching_trace_tampering_and_budget(self):
        family=SupportGraph(tuple('abc'),(1,1,2)).family()
        graph=check_trace(family,[(6,1),(5,2)])
        self.assertEqual(graph.predecessors,(1,1,2))
        with self.assertRaises(ValueError):check_trace(family,[(6,2),(5,2)])
        with self.assertRaises(ValueError):check_trace(family,[(2,1),(5,2)])
        partial=IntervalRecognizer(family,node_budget=1).solve()
        self.assertEqual(partial['status'],'unresolved')
        with self.assertRaises(ValueError):check_rejection(family,partial)
        bad=SetFamily(tuple('abc'),frozenset((0,1,2,3,7)));receipt=IntervalRecognizer(bad).solve()
        changed=copy.deepcopy(receipt);changed['nodes'][0]['children']=[]
        with self.assertRaises(ValueError):check_rejection(bad,changed)
        self.assertEqual(sum(SupportGraph(tuple('abc'),p).family()==family for p in itertools.product(range(8),repeat=3)),4)
        self.assertEqual(sum(SupportGraph(tuple('abc'),p).family().members==frozenset((0,7)) for p in itertools.product(range(8),repeat=3)),2)

    def test_hall_is_necessary_but_not_sufficient(self):
        conjunction=SetFamily(tuple('abc'),frozenset((0,1,2,3,7)));h=hall_matching(conjunction)
        self.assertFalse(h['matching_exists']);self.assertGreater(len(h['deficient_residual_sets']),len(h['their_eligible_roots']))
        family=hall_counterexample();self.assertTrue(hall_matching(family)['matching_exists'])
        result=IntervalRecognizer(family).solve();self.assertEqual(result['status'],'rejected');self.assertEqual(result['unique_states'],322)
        self.assertTrue(check_rejection(family,result)['valid'])

    def test_next_closure_and_explicit_list_certificate(self):
        for p in itertools.product(range(8),repeat=3):
            graph=SupportGraph(tuple('abc'),p);horn=graph.horn();models=list(horn.next_closure())
            expected={s for s in range(8) if all(not (body&~s==0) or s>>head&1 for body,head in horn.rules)}
            self.assertEqual(set(models),expected);self.assertEqual(len(models),len(expected))
            family=graph.family();certificate=graph.verify_explicit_family(family)
            self.assertTrue(certificate['equal']);self.assertEqual(certificate['models_checked'],len(family.members))
            truncated=SetFamily(family.labels,family.members-{max(family.members)})
            negative=graph.verify_explicit_family(truncated);self.assertFalse(negative['equal'])
            self.assertLessEqual(negative['models_checked'],len(truncated.members)+1)
        # This verifies an explicit list on 40 coordinates without a 2^40 scan.
        large=SupportGraph(tuple(f'x{i}' for i in range(40)),(0,)*40)
        self.assertEqual(large.verify_explicit_family(SetFamily(large.labels,frozenset((0,))))['models_checked'],1)
        with self.assertRaises(ValueError):large.family()

    def test_projection_deletion_and_disjoint_products(self):
        for p in itertools.product(range(8),repeat=3):
            graph=SupportGraph(tuple('abc'),p);family=graph.family()
            for removed in range(3):
                kept=[i for i in range(3) if i!=removed];visible=[graph.labels[i] for i in kept]
                mapped=lambda s:sum(1<<j for j,k in enumerate(kept) if s>>k&1)
                expected=frozenset(mapped(s) for s in family.members)
                projected=graph.project(visible)
                self.assertEqual(projected.family().members,expected)
                self.assertEqual(graph.delete(visible).family().members,frozenset(mapped(s) for s in family.members if not s>>removed&1))
                for small in range(4):
                    original=sum(1<<kept[j] for j in range(2) if small>>j&1)|(1<<removed)
                    self.assertEqual(projected.interior(small),mapped(family.interior(original)))
            self.assertEqual(graph.project(('a',)).family(),graph.eliminate('c').eliminate('b').family())
        for a,b in itertools.product(itertools.product(range(4),repeat=2),repeat=2):
            first=SupportGraph(('a','b'),a);second=SupportGraph(('c','d'),b)
            expected=frozenset(x|(y<<2) for x in first.family().members for y in second.family().members)
            self.assertEqual(first.independent_product(second).family().members,expected)

    def test_horn_completion_normalization(self):
        accepted=HornTheory(3,((1,1),));completion=HornTheory(3,((1,1),(1,2)))
        normalized,receipt=completion.normalize_completion(accepted,3)
        self.assertEqual(receipt['old_body'],1);self.assertEqual(receipt['new_body'],3)
        self.assertEqual(set(normalized.next_closure()),set(completion.next_closure()))
        self.assertNotEqual(set(HornTheory(3,((1,2),)).next_closure()),set(HornTheory(3,((3,2),)).next_closure()))
        with self.assertRaises(ValueError):completion.normalize_completion(accepted,7)

    def test_boundaries_elementary_realization_and_conjunction(self):
        for n in range(4):
            labels=tuple(range(n))
            for members in (frozenset(range(1<<n)),frozenset((0,))):
                family=SetFamily(labels,members);result=IntervalRecognizer(family).solve()
                self.assertEqual(result['status'],'success');graph=check_trace(family,result['trace'])
                self.assertEqual(StructuralCRS.elementary(graph).fixed_family(),family)
        self.assertEqual(IntervalRecognizer(SetFamily((),frozenset())).solve()['status'],'rejected')
        first=SupportGraph(tuple('abc'),(1,2,1));second=SupportGraph(tuple('abc'),(1,2,2))
        conjunction=SetFamily(tuple('abc'),first.family().members&second.family().members)
        self.assertEqual(conjunction.members,frozenset((0,1,2,3,7)))
        general=StructuralCRS(tuple('abc'),frozenset(('f',)),(frozenset(('f',)),frozenset(('f',)),frozenset(('x','y'))),
            (frozenset(('x',)),frozenset(('y',)),frozenset(('z',))),(frozenset(('f',)),)*3)
        self.assertEqual(general.fixed_family(),conjunction)
        self.assertEqual(IntervalRecognizer(conjunction).solve()['status'],'rejected')


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