"""CNF correspondence, literal family source and trustworthy completion status."""
from itertools import combinations, product
import unittest

from pysat.solvers import Solver
from example import (BooleanFormula, ModelEnumerator, PySATBackend, RAFOracle,
                     Reaction, ReactionSystem, SATFamilySource, SupportedTraceCNF,
                     essential_reactions, output_bits)


def subsets(values):
    values = tuple(values)
    return (frozenset(s) for k in range(len(values)+1) for s in combinations(values,k))


def literal_irr(system):
    result = []
    for subset in subsets(sorted(system.ids)):
        if system.is_raf(subset) and not any(i <= subset for i in result):
            result.append(subset)
    return frozenset(result)


class ScientificChecks(unittest.TestCase):
    def test_supported_trace_all_selected_sets_on_tiny_sources(self):
        fixtures = (
            ReactionSystem(set(),()),
            ReactionSystem(set(),(Reaction("empty",set(),set(),set()),)),
            ReactionSystem({"f"},(Reaction("self",{"f"},{"x"},{"x"}),)),
            ReactionSystem({"f"},(Reaction("a",{"f"},{"x"},{"y"}),Reaction("b",{"x"},{"y"},{"x"}))),
            ReactionSystem({"f"},(Reaction("cycle1",{"y"},{"x"},{"f"}),Reaction("cycle2",{"x"},{"y"},{"f"}))),
            ReactionSystem({"f"},(Reaction("free",set(),{"x"},{"f"}),Reaction("unsupported",{"missing"},{"y"},{"f"}))),
        )
        for system in fixtures:
            formula = SupportedTraceCNF(system)
            with Solver(name="g3",bootstrap_with=formula.clauses) as solver:
                for selected in subsets(system.ids):
                    assumptions = [formula.selection(i) if name in selected else -formula.selection(i)
                                   for i,name in enumerate(formula.reactions)]
                    sat = solver.solve(assumptions=assumptions)
                    self.assertEqual(sat,system.is_raf(selected))
                    if sat:
                        self.assertEqual(formula.check_model(solver.get_model()),selected)
                        self.assertEqual(formula.check_model(formula.witness(selected)),selected)

    def test_containers_avoidance_and_blocking_reaction_sets(self):
        system = ReactionSystem({"f"},tuple(Reaction(r,{"f"},{r},set("abc")-{r}) for r in "abc"))
        irr = literal_irr(system)
        for known in subsets(irr):
            for container in subsets(system.ids):
                formula = SupportedTraceCNF(system,known,container)
                expected = any(system.is_raf(s) and all(not i <= s for i in known) for s in subsets(container))
                with Solver(name="g3",bootstrap_with=formula.clauses) as solver:
                    self.assertEqual(solver.solve(),expected)
                    if expected:
                        selected = formula.check_model(solver.get_model())
                        self.assertTrue(all(not i <= selected for i in known))
        # Known {a,b} is already avoided because a lies outside this container.
        query = SupportedTraceCNF(system,({"a","b"},),{"b","c"})
        with Solver(name="g3",bootstrap_with=query.clauses) as solver:
            self.assertTrue(solver.solve())

    def test_whole_family_source_exhaustive_choice_sets(self):
        # Every subset of the nine non-tautological clauses on two variables,
        # including the empty clause; all 16 choice sets for each source.
        possible = []
        for signs in product((-1,0,1),repeat=2):
            possible.append(tuple((i+1)*s for i,s in enumerate(signs) if s))
        for mask in range(1<<len(possible)):
            boolean = BooleanFormula(2,tuple(c for i,c in enumerate(possible) if mask & (1<<i)))
            source = SATFamilySource(boolean)
            accepted = []
            for choices in subsets(source.inputs):
                expected = boolean.accepts_choices(choices)
                self.assertEqual(source.system.is_raf(source.canonical(choices)),expected)
                if expected:
                    accepted.append(choices)
            minimal = [s for s in accepted if not any(t < s for t in accepted)]
            family = frozenset(source.canonical(s) for s in minimal)
            self.assertEqual(family,source.predicted_family())
            self.assertEqual(len(family),2+len(boolean.satisfying_assignments()))
            # The reset alone, or the full auxiliary cycle, cannot bootstrap out.
            self.assertNotIn("out",source.system.closure_stages(source.auxiliary)[-1])

    def test_enumerator_call_count_and_worked_dimensions(self):
        for clauses,expected,bits in ((((1,),(-1,)),3,91),(((1,),(2,)),5,151)):
            source = SATFamilySource(BooleanFormula(3,clauses))
            self.assertEqual(source.dimensions(),{"molecules":36,"reactions":29,"input_bits":3235,"baseline_output_bits":91})
            for known in ((),source.baseline):
                report = ModelEnumerator(source.system).run(known)
                self.assertEqual(report["status"],"complete")
                self.assertEqual(frozenset(map(frozenset,report["family"])),source.predicted_family())
                self.assertEqual(report["solver_calls"],expected-len(known)+1)
                self.assertLessEqual(report["minimization_maxraf_calls"],29*(expected-len(known)))
                encoded = output_bits(tuple(r.name for r in source.system.reactions),report["family"])
                self.assertEqual(len(encoded),bits)
                self.assertTrue(encoded.startswith("1"*expected+"0"))
            essentials = essential_reactions(source.system)
            self.assertEqual(frozenset(essentials["reactions"]),source.auxiliary)

    def test_limits_unknown_and_invalid_solver_model(self):
        system = ReactionSystem({"f"},(Reaction("self",{"f"},{"x"},{"x"}),))
        self.assertEqual(ModelEnumerator(system).run(max_calls=0)["status"],"incomplete")
        self.assertEqual(ModelEnumerator(system).run(max_literals=0)["solver_calls"],0)
        class Unknown:
            def __init__(self,clauses): pass
            def solve(self): return "unknown",None
            def close(self): pass
        self.assertEqual(ModelEnumerator(system,Unknown).run()["status"],"incomplete")
        class Malformed(Unknown):
            def solve(self): return "sat",[1]
        with self.assertRaises(ValueError):
            ModelEnumerator(system,Malformed).run()
        with self.assertRaises(ValueError):
            ModelEnumerator(system).run((frozenset(),))
        # A solver call with one found member still needs a final UNSAT call.
        limited = ModelEnumerator(system).run(max_calls=1)
        self.assertEqual(len(limited["family"]),1)
        self.assertEqual(limited["status"],"incomplete")
        self.assertEqual(ModelEnumerator(system).run(limited["family"])["solver_calls"],1)

    def test_low_order_deletions_and_nonessential_maxraf_members(self):
        source = SATFamilySource(BooleanFormula(3,((1,),(2,))))
        oracle = RAFOracle(source.system)
        for size in range(3):
            for deleted in combinations(source.inputs.values(),size):
                residual = source.system.ids-set(deleted)
                self.assertEqual(oracle.maximum(residual),residual)
        for assignment in product((False,True),repeat=3):
            deleted = {source.inputs[(i,not b)] for i,b in enumerate(assignment)}
            residual = source.system.ids-deleted
            self.assertEqual(bool(oracle.maximum(residual)),source.formula.evaluate(assignment))
        chain = ReactionSystem({"f"},(Reaction("a",{"f"},{"x"},{"f"}),Reaction("b",{"x"},{"y"},{"f"})))
        self.assertEqual(literal_irr(chain),frozenset((frozenset(("a",)),)))
        self.assertEqual(RAFOracle(chain).maximum(),frozenset(("a","b")))
        self.assertEqual(essential_reactions(chain)["reactions"],["a"])
        no_raf = ReactionSystem({"f"},(Reaction("a",{"f"},{"x"},set()),))
        self.assertEqual(essential_reactions(no_raf)["status"],"no RAF")

    def test_distinct_ids_same_chemistry_and_formula_export(self):
        system = ReactionSystem({"food"},tuple(Reaction(name,{"food"},{"x"},{"food"}) for name in ("first","second")))
        report = ModelEnumerator(system).run()
        self.assertEqual(frozenset(map(frozenset,report["family"])),frozenset((frozenset(("first",)),frozenset(("second",)))))
        formula = SupportedTraceCNF(system)
        from pysat.formula import CNF
        parsed = CNF(from_string=formula.dimacs())
        self.assertEqual(tuple(map(tuple,parsed.clauses)),formula.clauses)
        self.assertEqual(ReactionSystem.from_dict(system.to_dict()),system)
        self.assertEqual(output_bits((),()),"0")
        with self.assertRaises(ValueError):
            output_bits(("first",),(("first",),("first",)))


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