"""Literal model regressions and independent exact small-case checks."""
from fractions import Fraction as F
from itertools import combinations, permutations, product
import math
import unittest

import numpy as np
from example import (CatalysedNetwork, Channel, DependencyCodec, Instruction,
                     PolymerModel, ProductiveProgram, Reference, SparseCatalysisLaw,
                     TheoremBounds, channel_count, molecule_count)


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


def independent_closure(model, channels):
    # Explicit synchronous orientation rules, separate from enabled/productive.
    state = set(model.food())
    while True:
        new = set(state)
        for c in channels:
            u, v, w = c.endpoints
            if u in state and v in state:
                new.add(w)
            if w in state:
                new.update((u,v))
        if new == state:
            return frozenset(state)
        state = new


def independent_is_raf(network, indices, only=None):
    closure = independent_closure(network.model, [network.channels[i] for i in indices])
    return bool(indices) and all(set(network.channels[i].endpoints) <= closure and
        bool(network.catalysts[i] & closure & (network.active if only is None else frozenset((only,)))) for i in indices)


class ScientificChecks(unittest.TestCase):
    def test_catalogue_and_reversible_closure(self):
        for n in range(2, 8):
            model = PolymerModel(n, 1)
            self.assertEqual(len(model.molecules()), molecule_count(n))
            self.assertEqual(len(model.channels()), channel_count(n))
        # Same product and unordered factors, distinct split-position channels.
        a, b = Channel("000",1), Channel("000",2)
        self.assertNotEqual(a,b)
        self.assertEqual(set(a.endpoints), set(b.endpoints))
        model = PolymerModel(5,2)
        channels = (Channel("010",2), Channel("01011",3), Channel("01011",4))
        self.assertIn("0101", model.closure(channels))  # obtainable by cleavage
        self.assertEqual(model.closure(channels), independent_closure(model, channels))

    def test_maximum_and_minimum_against_independent_oracle(self):
        model = PolymerModel(3,1)
        catalogue = model.channels()[:7]
        rng = np.random.default_rng(12)
        molecules = model.molecules()
        for _ in range(32):
            catalysts = tuple(frozenset(w for w in molecules if rng.random() < .12) for _ in catalogue)
            active = frozenset().union(*catalysts)
            network = CatalysedNetwork(model, catalogue, catalysts, active)
            rafs = [s for s in powerset(range(len(catalogue))) if independent_is_raf(network,s)]
            union = frozenset().union(*rafs)
            self.assertEqual(network.max_raf(), union)
            minimum = network.exact_minimum()
            self.assertEqual(None if minimum is None else len(minimum), min(map(len,rafs), default=None))
            for x in active:
                exists = any(independent_is_raf(network,s,x) for s in rafs)
                self.assertEqual(x in network.single_catalyst_witnesses(), exists)
            if union:
                with self.assertRaises(ValueError):
                    network.exact_minimum(max_channels=0)

    def test_saturation_and_empty_food_only_extraction(self):
        model = PolymerModel(3,2)
        catalogue = model.channels()[:8]
        for subset in powerset(catalogue):
            program = model.saturated_program(subset)
            self.assertEqual(program.states()[-1], independent_closure(model, subset))
            self.assertTrue(set(program.sequence) <= subset)
        food_channel = Channel("00",1)
        network = CatalysedNetwork(model, (food_channel,), (frozenset(("0",)),), frozenset(("0",)))
        self.assertEqual(network.max_raf(), frozenset((0,)))
        self.assertEqual(model.saturated_program((food_channel,)).sequence, ())
        self.assertIn("0", network.single_catalyst_witnesses())
        no_food = PolymerModel(3,0)
        self.assertEqual(no_food.closure(no_food.channels()), frozenset())

    def test_dependency_codes_count_labels_not_orders(self):
        model = PolymerModel(5,2)
        support = (Channel("010",2), Channel("01011",3), Channel("01011",4))
        valid_orders = []
        for order in permutations(support):
            p = ProductiveProgram(model, order)
            try:
                p.states()
                valid_orders.append(p)
            except ValueError:
                pass
        self.assertEqual(len(valid_orders), 1)
        codec = DependencyCodec(model)
        codes = set()
        for labels in permutations(support):
            code = codec.encode(valid_orders[0], labels)
            codes.add(code)
            self.assertEqual(codec.decode(code), labels)
        self.assertEqual(len(codes), math.factorial(3))
        cycle = (Instruction("cleavage", (Reference(label=0,endpoint=2),),1),)
        with self.assertRaisesRegex(ValueError, "Cyclic"):
            codec.decode(cycle)

    def test_exhaustive_small_productive_support_encoding(self):
        model = PolymerModel(3,1)
        codec = DependencyCodec(model)
        for depth in range(1,4):
            by_support = {}
            for program in model.programs(depth):
                by_support.setdefault(frozenset(program.sequence), program)
            seen = {}
            for support, program in by_support.items():
                for labelled in permutations(sorted(support)):
                    code = codec.encode(program, labelled)
                    decoded = codec.decode(code)
                    self.assertEqual(decoded, labelled)
                    if code in seen:
                        self.assertEqual(seen[code], labelled)
                    seen[code] = labelled
            self.assertEqual(len(seen), len(by_support)*math.factorial(depth))
            f = len(model.food())
            self.assertLessEqual(len(seen), ((f+3*depth)**2+model.n*(f+3*depth))**depth)

    def test_sparse_law_preserves_shared_activity(self):
        model, law = PolymerModel(2,1), SparseCatalysisLaw(.25)
        params = law.parameters(2)
        self.assertEqual(params["activity"], .25)
        rng = np.random.default_rng(91)
        observed = 0
        trials = 5000
        for _ in range(trials):
            network = law.sample(model,rng)
            observed += "0" in network.catalysts[0] and "0" in network.catalysts[1]
        p, q = params["activity"], params["conditional"]
        self.assertAlmostEqual(observed/trials, p*q*q, delta=.015)
        self.assertGreater(observed/trials, 2*(p*q)**2)
        self.assertEqual(SparseCatalysisLaw(0).sample(model,rng).max_raf(), frozenset())
        self.assertTrue(SparseCatalysisLaw(100).parameters(2)["clipped"])
        self.assertEqual(SparseCatalysisLaw(100).parameters(2)["expected_channels_per_molecule"], 2)
        for n in range(2,10):
            uncut = SparseCatalysisLaw(.1).parameters(n)
            self.assertAlmostEqual(uncut["expected_channels_per_molecule"], .1*n)

    def test_bounds_against_rational_formulas_and_cutoff(self):
        bounds = TheoremBounds(2, .5)
        for n in (4, 10, 50, 100):
            p = min(F(1), F(n*n, 2*channel_count(n)))
            q, count = F(1,n), molecule_count(n)
            for rank in (1,2,3):
                prefix, shallow = bounds.prefix_constants(rank+1)
                exact = shallow*p + count**rank*prefix*rank**(rank+1)*p**rank*q**(rank+1)
                self.assertAlmostEqual(bounds.log_rank_bound(n,rank), math.log(exact), places=10)
            exact_single = molecule_count(4)*p+count*channel_count(4)*channel_count(8)*p*q*q
            self.assertAlmostEqual(bounds.log_single_bound(n), math.log(exact_single), places=10)
        for budget in (0,1,2,3):
            cutoff, d = bounds.cutoff(budget)
            for k in range(cutoff+1,cutoff+101):
                self.assertLessEqual((d*k)**budget*3**k,4**k)
        self.assertFalse(bounds.linear_bound(4,1)["applicable"])
        self.assertTrue(bounds.linear_bound(1000,1)["applicable"])
        # Large-n evaluation must not lose log(N*p) to catastrophic cancellation.
        for rank in (1,2,3):
            difference = bounds.log_rank_bound(10**18,rank)-bounds.log_rank_bound(10**17,rank)
            self.assertAlmostEqual(difference,-math.log(10),places=10)
        self.assertEqual(TheoremBounds(0,1).linear_bound(100,1)["capped_bound"],0)
        self.assertEqual(TheoremBounds(2,0).log_rank_bound(100,1),-math.inf)


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