"""Independent literal subset oracle and exhaustive reduction checks."""
from itertools import combinations, product
import unittest

from example import (AxiomReduction, CliqueReduction, CompletenessChecker,
                     Implication, ImplicationSystem, RAFOracle, Reaction,
                     ReactionSystem, coverage_example, literal_family)


def bitmask_oracle(system):
    """Literal RAF predicate using molecule bitmasks, without maxRAF pruning."""
    species = sorted(set(system.food).union(*(r.reactants | r.products | r.catalysts for r in system.reactions)))
    lookup = {s:1<<i for i,s in enumerate(species)}
    def bits(names):
        return sum(lookup[s] for s in names)
    food = bits(system.food)
    triples = [(bits(r.reactants),bits(r.products),bits(r.catalysts)) for r in system.reactions]
    raf_masks, minimal = [], []
    for mask in range(1,1<<len(triples)):
        selected = [r for i,r in enumerate(triples) if mask & (1<<i)]
        closure = food
        while True:
            new = closure
            for reactants, products, catalysts in selected:
                if reactants & new == reactants:
                    new |= products
            if new == closure:
                break
            closure = new
        if all(r & closure == r and c & closure for r,p,c in selected):
            raf_masks.append(mask)
    for mask in sorted(raf_masks,key=int.bit_count):
        if not any(previous & mask == previous for previous in minimal):
            minimal.append(mask)
    return tuple(frozenset(r.name for i,r in enumerate(system.reactions) if mask & (1<<i)) for mask in minimal),len(raf_masks)


def all_rule_systems():
    statements = ("0","1")
    premises = (frozenset(),frozenset(("0",)),frozenset(("1",)),frozenset(statements))
    possible = tuple(Implication(p,c) for p in premises for c in statements)
    for mask in range(256):
        yield ImplicationSystem(statements,tuple(r for i,r in enumerate(possible) if mask & (1<<i)))


class ScientificChecks(unittest.TestCase):
    def test_coverage_overlap_duplicates_and_empty_family(self):
        system = coverage_example()
        checker = CompletenessChecker(system)
        supplied = ({"a","b"},{"a","c"})
        report = checker.check(supplied)
        self.assertEqual(report["status"],"incomplete")
        self.assertEqual(report["missing"],["b","c"])
        self.assertEqual(report["omissions"],("a","a"))  # one deletion may hit two sets
        self.assertEqual(report["trace"][0]["deleted"],["a"])
        self.assertTrue(checker.verify_report(report))
        duplicates = checker.check(supplied+(supplied[0],))
        self.assertEqual(duplicates["k"],2)
        all_members = supplied+({"b","c"},)
        complete = checker.check(all_members)
        self.assertEqual(complete["status"],"complete")
        self.assertEqual(complete["tested_tuples"],8)
        self.assertLess(complete["distinct_residual_queries"],complete["tested_tuples"])
        self.assertTrue(checker.verify_report(complete))
        self.assertEqual(checker.check(())["status"],"incomplete")
        self.assertEqual(checker.check(all_members,max_transversals=0)["status"],"inconclusive")
        self.assertEqual(checker.check(({"a"},))["status"],"invalid")
        self.assertEqual(checker.check((set(),))["status"],"invalid")
        self.assertEqual(checker.check(({"unknown"},))["status"],"invalid")
        no_raf = ReactionSystem({"f"},(Reaction("r",{"f"},{"x"},set()),))
        self.assertEqual(CompletenessChecker(no_raf).check(())["status"],"complete")

    def test_one_deletion_requires_maximum_not_direct_raf_test(self):
        # Full source is a RAF. Every immediate deletion fails to be a RAF,
        # but {base} is a proper RAF hidden inside two of those deletions.
        system = ReactionSystem({"f"},(
            Reaction("base",{"f"},{"a"},{"a"}),
            Reaction("left",{"a"},{"b"},{"c"}),
            Reaction("right",{"a"},{"c"},{"b"})))
        self.assertTrue(system.is_raf(system.ids))
        self.assertTrue(all(not system.is_raf(system.ids-{r}) for r in system.ids))
        oracle = RAFOracle(system)
        self.assertFalse(oracle.irreducible(system.ids))
        self.assertEqual(oracle.extract(system.ids),frozenset(("base",)))
        self.assertEqual(CompletenessChecker(system).check((system.ids,))["status"],"invalid")

    def test_worked_reduction_and_source_decoding(self):
        source = ImplicationSystem(("0","1","2"),tuple(Implication({str(i)},str((i+1)%3)) for i in range(3)))
        reduction = AxiomReduction(source,1)
        irr,count = bitmask_oracle(reduction.system)
        self.assertEqual(count,33)
        self.assertEqual(len(irr),4)
        self.assertEqual(len(reduction.system.reactions),10)
        species = reduction.system.food.union(*(r.reactants | r.products | r.catalysts for r in reduction.system.reactions))
        self.assertEqual(len(species),12)
        for member in irr:
            if member not in reduction.guards:
                decoded = reduction.decode_missing(member)
                self.assertEqual(len(decoded["axioms"]),1)
                witness = reduction.source_witness(decoded["slot_choices"])
                self.assertTrue(member <= witness)
        no_seed = AxiomReduction(source,0)
        self.assertEqual(CompletenessChecker(no_seed.system).check(no_seed.guards)["status"],"complete")
        # Cyclic implications cannot invent their own first statement molecule.
        self.assertEqual(source.stages(())[-1],frozenset())

    def test_all_768_two_statement_reductions(self):
        yes_count = 0
        for source in all_rule_systems():
            for slots in range(3):
                reduction = AxiomReduction(source,slots)
                oracle = RAFOracle(reduction.system)
                self.assertTrue(all(oracle.irreducible(g) for g in reduction.guards))
                self.assertEqual(len(set(reduction.guards)),slots)
                report = CompletenessChecker(reduction.system).check(reduction.guards)
                source_yes = bool(source.generators(slots))
                yes_count += source_yes
                self.assertEqual(report["status"] == "incomplete",source_yes)
                if source_yes:
                    decoded = reduction.decode_missing(report["missing"])
                    self.assertLessEqual(len(decoded["axioms"]),slots)
                    self.assertEqual(source.stages(decoded["axioms"])[-1],frozenset(source.statements))
        self.assertEqual(yes_count,624)

    def test_independent_literal_oracle_all_512_cases(self):
        yes_count = 0
        for source in all_rule_systems():
            for slots in range(2):
                reduction = AxiomReduction(source,slots)
                irreducibles,_ = bitmask_oracle(reduction.system)
                extra = set(irreducibles)-set(reduction.guards)
                source_yes = bool(source.generators(slots))
                self.assertEqual(bool(extra),source_yes)
                yes_count += bool(extra)
                self.assertTrue(set(reduction.guards) <= set(irreducibles))
        self.assertEqual(yes_count,368)

    def test_graph_reduction_and_counts(self):
        for vertices in (2,3):
            edges = tuple(combinations(range(vertices),2))
            for mask in range(1<<len(edges)):
                present = tuple(e for i,e in enumerate(edges) if mask & (1<<i))
                present_sets = set(map(frozenset,present))
                for slots in range(4):
                    reduction = CliqueReduction(vertices,present,slots)
                    self.assertEqual(len(reduction.system.reactions),2*slots*vertices+2*(slots*vertices)**2+1)
                    source_yes = any(all(frozenset(e) in present_sets for e in combinations(c,2))
                                     for c in combinations(range(vertices),slots))
                    report = CompletenessChecker(reduction.system).check(reduction.guards)
                    self.assertEqual(report["status"] == "incomplete",source_yes)
                    if source_yes:
                        witness = reduction.decode_missing(report["missing"])
                        self.assertEqual(len(witness),slots)

    def test_normalization_serialization_and_tamper_rejection(self):
        for statements,rules in (((),()),(("a",),()),(("a",),(Implication(set(),"a"),))):
            source = ImplicationSystem(statements,rules)
            for slots in (0,1):
                reduction = AxiomReduction(source,slots)
                result = CompletenessChecker(reduction.system).check(reduction.guards)
                self.assertEqual(result["status"] == "incomplete",bool(source.generators(slots)))
                self.assertEqual(ReactionSystem.from_dict(reduction.system.to_dict()),reduction.system)
        checker = CompletenessChecker(coverage_example())
        report = checker.check(({"a","b"},{"a","c"}))
        forged = dict(report,missing=["a"])
        self.assertFalse(checker.verify_report(forged))
        forged = dict(report,status="complete")
        self.assertFalse(checker.verify_report(forged))


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