"""Tests for tpc.py.

  python3 -m unittest test_tpc -v      (or: python3 test_tpc.py)

The interesting ones are the invariants rather than the transcripts:
cooperative termination must NEVER produce a split, on any vote vector, at
any crash step -- and the closed form blocked = 3N+1 must hold for every N.
"""

import itertools
import unittest

from tpc import (Sim, Participant, Coordinator, run_crash_at, blocked_timings,
                 total_steps, INIT, PREPARED, COMMITTED, ABORTED, C)

ALL_YES = ["yes", "yes", "yes"]
ONE_NO = ["yes", "no", "yes"]


def vote_vectors(n):
    return [list(v) for v in itertools.product(["yes", "no"], repeat=n)]


class CleanRun(unittest.TestCase):
    def test_all_yes_commits_everyone(self):
        s = Sim(ALL_YES)
        s.run()
        self.assertEqual(set(s.outcomes().values()), {COMMITTED})
        self.assertEqual(s.coord.log, ["commit"])

    def test_one_no_aborts_everyone(self):
        s = Sim(ONE_NO)
        s.run()
        self.assertEqual(set(s.outcomes().values()), {ABORTED})
        self.assertEqual(s.coord.log, ["abort"])

    def test_clean_run_releases_every_lock(self):
        for votes in vote_vectors(3):
            s = Sim(votes)
            s.run()
            self.assertEqual(s.locks_held(), 0, votes)

    def test_step_count_is_8N_when_unanimous(self):
        for n in range(1, 9):
            self.assertEqual(total_steps(["yes"] * n), 8 * n)

    def test_a_yes_voter_forces_prepare_before_replying(self):
        p = Participant(1, "yes")
        out = p.on_prepare({})
        self.assertEqual(p.log, ["prepare"])       # forced first
        self.assertEqual(out[0]["vote"], "yes")    # then the reply
        self.assertEqual(p.locks, 1)


class TheNoVoter(unittest.TestCase):
    """The load-bearing line: a NO voter never enters the uncertainty period."""

    def test_no_voter_never_becomes_prepared(self):
        p = Participant(2, "no")
        p.on_prepare({})
        self.assertEqual(p.state, ABORTED)
        self.assertNotEqual(p.state, PREPARED)
        self.assertEqual(p.locks, 0)
        self.assertFalse(p.uncertain())

    def test_no_voter_is_never_uncertain_at_any_crash_step(self):
        n = total_steps(ONE_NO)
        for k in range(n + 1):
            sim, _ = run_crash_at(ONE_NO, k, recovery="none")
            self.assertFalse(sim.parts[2].uncertain(), f"k={k}")


class CrashSweep(unittest.TestCase):
    def test_headline_all_yes(self):
        blocked, n = blocked_timings(ALL_YES)
        self.assertEqual(n, 24)
        self.assertEqual(blocked, list(range(3, 13)))
        self.assertEqual((len(blocked), n + 1), (10, 25))

    def test_headline_one_no(self):
        blocked, n = blocked_timings(ONE_NO)
        self.assertEqual(blocked, [])
        self.assertEqual(n + 1, 23)

    def test_one_no_anywhere_never_blocks(self):
        """Not an artefact of P2 being the vetoer."""
        for i in range(3):
            votes = ["yes"] * 3
            votes[i] = "no"
            self.assertEqual(blocked_timings(votes)[0], [], votes)

    def test_any_number_of_no_voters_never_blocks(self):
        for votes in vote_vectors(4):
            if "no" in votes:
                self.assertEqual(blocked_timings(votes)[0], [], votes)

    def test_closed_form_3N_plus_1(self):
        for n in range(1, 9):
            blocked, steps = blocked_timings(["yes"] * n)
            self.assertEqual(steps, 8 * n)
            self.assertEqual(len(blocked), 3 * n + 1, f"N={n}")

    def test_blocked_window_is_contiguous(self):
        for n in range(1, 7):
            blocked, _ = blocked_timings(["yes"] * n)
            self.assertEqual(blocked, list(range(blocked[0], blocked[-1] + 1)))

    def test_blocked_means_every_participant_uncertain_holding_locks(self):
        blocked, _ = blocked_timings(ALL_YES)
        for k in blocked:
            sim, _ = run_crash_at(ALL_YES, k, recovery="none")
            self.assertTrue(all(p.uncertain() for p in sim.parts.values()))
            self.assertEqual(sim.locks_held(), 3, f"k={k}")

    def test_nine_of_ten_blocked_timings_had_no_decision_at_all(self):
        blocked, _ = blocked_timings(ALL_YES)
        logged = [k for k in blocked
                  if run_crash_at(ALL_YES, k, recovery="none")[0].coord.log]
        self.assertEqual(logged, [12])
        self.assertEqual(len(blocked) - len(logged), 9)


class Safety(unittest.TestCase):
    """Cooperative termination is allowed to block. It is not allowed to lie."""

    @staticmethod
    def _effective(sim):
        return {(ABORTED if p.state == INIT else p.state)
                for p in sim.parts.values()}

    def test_termination_never_splits(self):
        for n in (2, 3, 4):
            for votes in vote_vectors(n):
                for k in range(total_steps(votes) + 1):
                    sim, res = run_crash_at(votes, k)
                    if res == "BLOCKED":
                        continue
                    self.assertEqual(len(self._effective(sim)), 1,
                                     f"{votes} k={k}")

    def test_termination_never_contradicts_a_forced_decision(self):
        for votes in vote_vectors(3):
            for k in range(total_steps(votes) + 1):
                sim, res = run_crash_at(votes, k)
                if res == "BLOCKED" or not sim.coord.log:
                    continue
                want = COMMITTED if sim.coord.log[0] == "commit" else ABORTED
                self.assertEqual(self._effective(sim), {want}, f"{votes} k={k}")

    def test_resolution_releases_every_lock(self):
        for k in range(total_steps(ALL_YES) + 1):
            sim, res = run_crash_at(ALL_YES, k)
            if res != "BLOCKED":
                self.assertEqual(sim.locks_held(), 0, f"k={k}")


class Uncertainty(unittest.TestCase):
    def test_two_worlds_are_locally_indistinguishable(self):
        """The heart of the toy: identical local state, opposite answers."""
        def local(sim):
            p = sim.parts[1]
            return (p.state, p.vote, tuple(p.log), p.locks)

        a, _ = run_crash_at(ALL_YES, 12, recovery="none")
        b, _ = run_crash_at(ONE_NO, 12, recovery="none")
        self.assertEqual(local(a), local(b))
        self.assertEqual(a.coord.log, ["commit"])
        self.assertEqual(b.coord.log, ["abort"])

    def test_peers_add_no_information_when_blocked(self):
        sim, res = run_crash_at(ALL_YES, 12)
        self.assertEqual(res, "BLOCKED")
        states = {p.state for p in sim.parts.values()}
        self.assertEqual(states, {PREPARED})


class Heuristics(unittest.TestCase):
    """Refusing to block relocates the damage; it does not remove it."""

    @staticmethod
    def _split(sim):
        return len({(ABORTED if p.state == INIT else p.state)
                    for p in sim.parts.values()}) > 1

    def test_presumed_abort_splits_at_13_and_14(self):
        splits = [k for k in range(total_steps(ALL_YES) + 1)
                  if self._split(run_crash_at(ALL_YES, k, recovery="unilateral",
                                              heuristic="abort")[0])]
        self.assertEqual(splits, [13, 14])

    def test_presumed_commit_wrecks_the_vetoed_transaction(self):
        splits = [k for k in range(total_steps(ONE_NO) + 1)
                  if self._split(run_crash_at(ONE_NO, k, recovery="unilateral",
                                              heuristic="commit")[0])]
        self.assertEqual(splits, list(range(1, 15)))

    def test_heuristics_never_block(self):
        for h in ("abort", "commit"):
            for k in range(total_steps(ALL_YES) + 1):
                sim, res = run_crash_at(ALL_YES, k, recovery="unilateral",
                                        heuristic=h)
                self.assertIsNone(res)
                self.assertFalse(any(p.uncertain() for p in sim.parts.values()))


class Mechanics(unittest.TestCase):
    def test_a_crashed_coordinator_transmits_nothing(self):
        sim, _ = run_crash_at(ALL_YES, 12, recovery="none")
        self.assertTrue(any("NEVER SENT" in line for line in sim.trace))
        self.assertFalse(any("DELIVER C->" in line and "decision" in line
                             for line in sim.trace))

    def test_messages_already_sent_survive_the_crash(self):
        sim, _ = run_crash_at(ALL_YES, 13, recovery="none")
        self.assertEqual(sim.parts[1].state, COMMITTED)
        self.assertEqual(sim.parts[2].state, PREPARED)

    def test_decisions_are_idempotent(self):
        p = Participant(1, "yes")
        p.on_prepare({})
        p.on_decision(dict(decision="commit"))
        self.assertEqual(p.on_decision(dict(decision="abort")), [])
        self.assertEqual(p.state, COMMITTED)

    def test_coordinator_needs_every_vote_before_deciding(self):
        c = Coordinator([1, 2, 3])
        c.start()
        self.assertEqual(c.on_vote(dict(src=1, vote="yes")), [])
        self.assertEqual(c.on_vote(dict(src=2, vote="yes")), [])
        self.assertEqual(c.log, [])
        out = c.on_vote(dict(src=3, vote="yes"))
        self.assertEqual(c.log, ["commit"])
        self.assertEqual(len(out), 3)


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