"""One replicated set, two ways of deciding a conflict.

`Replica` is a state-based OR-Set. Every operation mints a *dot* -- a pair
(replica_id, counter) -- and supersedes every assertion about that element the
replica can currently see. A per-replica vector clock (`cc`) records the
highest counter observed from each replica, and at merge time it answers the
only question that matters: is this assertion one I have already seen and
superseded, or one I have never seen at all? What survives a merge is exactly
the set of assertions that are pairwise concurrent.

`PureLWW` answers the same question with a timestamp. No dots, no vector
clock: the write with the highest (timestamp, replica_id) wins, and that is
what Cassandra and Riak's default conflict resolution actually do.

Both converge. They converge to different things, and only one of them can be
improved by talking more.

No wall clock anywhere. Every timestamp is supplied by the caller, and clock
skew is an explicit per-replica integer offset. The operation trace and the
gossip schedule draw from two separate seeded RNGs, so changing the gossip
rate cannot change which operations happen.
"""

import random

POLICIES = ("add-wins", "rm-wins", "lww")


def present(assertions, policy):
    """Is the element in the set? `assertions` is {dot: (label, ts)}."""
    if not assertions:
        return False
    if policy == "add-wins":
        return any(lbl == "add" for lbl, _ in assertions.values())
    if policy == "rm-wins":
        return all(lbl == "add" for lbl, _ in assertions.values())
    if policy == "lww":
        # Highest timestamp wins, ties broken by dot so it stays deterministic.
        _, (lbl, _) = max(assertions.items(), key=lambda kv: (kv[1][1], kv[0]))
        return lbl == "add"
    raise ValueError(policy)


class Replica:
    """An OR-Set replica: dotted assertions plus a vector clock over them."""

    def __init__(self, rid, policy="add-wins", supersede=True, causal_merge=True):
        self.rid = rid
        self.policy = policy
        self.supersede = supersede        # counterfactual knob, see commentary 5.5
        self.causal_merge = causal_merge  # counterfactual knob, see commentary 5.4
        self.seq = 0
        self.cc = {}   # rid -> highest counter observed from that replica
        self.el = {}   # elem -> {dot: (label, ts)}

    # ---- local operations -------------------------------------------------
    def _assert(self, elem, label, ts):
        """Mint a fresh dot and supersede everything this replica can see."""
        self.seq += 1
        dot = (self.rid, self.seq)
        self.cc[self.rid] = self.seq
        if self.supersede:
            self.el[elem] = {dot: (label, ts)}
        else:
            self.el.setdefault(elem, {})[dot] = (label, ts)

    def add(self, elem, ts):
        self._assert(elem, "add", ts)

    def remove(self, elem, ts):
        self._assert(elem, "rm", ts)

    # ---- merge ------------------------------------------------------------
    def covers(self, dot):
        """Have I ever observed this dot? The vector clock's whole job."""
        return self.cc.get(dot[0], 0) >= dot[1]

    def merge(self, other):
        """Absorb `other`. Keep an assertion the other side lacks only if the
        other side never saw it; if it saw it and dropped it, it superseded
        it, and the drop is news."""
        merged = {}
        for elem in set(self.el) | set(other.el):
            mine = self.el.get(elem, {})
            theirs = other.el.get(elem, {})
            if not self.causal_merge:
                keep = dict(mine)
                keep.update(theirs)
            else:
                keep = {}
                for dot, v in mine.items():
                    if dot in theirs or not other.covers(dot):
                        keep[dot] = v
                for dot, v in theirs.items():
                    if dot in mine or not self.covers(dot):
                        keep[dot] = v
            if keep:
                merged[elem] = keep
        self.el = merged
        for rid, seq in other.cc.items():
            if seq > self.cc.get(rid, 0):
                self.cc[rid] = seq

    # ---- observation ------------------------------------------------------
    def value(self):
        return frozenset(e for e, a in self.el.items() if present(a, self.policy))


class PureLWW:
    """Last write wins on the timestamp alone: no dots, no vector clock."""

    def __init__(self, rid, policy="pure-lww", **_kw):
        self.rid = rid
        self.el = {}   # elem -> (ts, rid, label)

    def _assert(self, elem, label, ts):
        cand = (ts, self.rid, label)
        cur = self.el.get(elem)
        if cur is None or cand[:2] > cur[:2]:
            self.el[elem] = cand

    def add(self, elem, ts):
        self._assert(elem, "add", ts)

    def remove(self, elem, ts):
        self._assert(elem, "rm", ts)

    def merge(self, other):
        for elem, v in other.el.items():
            cur = self.el.get(elem)
            if cur is None or v[:2] > cur[:2]:
                self.el[elem] = v

    def value(self):
        return frozenset(e for e, v in self.el.items() if v[2] == "add")


# ---- drivers --------------------------------------------------------------

def _replicas(policy, n_replicas, **kw):
    cls = PureLWW if policy == "pure-lww" else Replica
    return [cls(r, policy, **kw) for r in range(n_replicas)]


def gen_trace(seed, n_ops, n_replicas, n_elems, sync_every, p_add=0.55):
    """Events: ('op', rid, elem, label) and ('sync',). One RNG, one purpose."""
    rng = random.Random(seed)
    events = []
    for i in range(n_ops):
        rid = rng.randrange(n_replicas)
        elem = "e%d" % rng.randrange(n_elems)
        label = "add" if rng.random() < p_add else "rm"
        events.append(("op", rid, elem, label))
        if sync_every and (i + 1) % sync_every == 0:
            events.append(("sync",))
    return events


def sequential(trace):
    """Ground truth: the same operations, in order, on one machine's set."""
    s = set()
    for e in trace:
        if e[0] == "op":
            s.add(e[2]) if e[3] == "add" else s.discard(e[2])
    return frozenset(s)


def sync_all(reps):
    for a in reps:
        for b in reps:
            if a is not b:
                a.merge(b)


def run(trace, policy, n_replicas, skew=None, **kw):
    """Replay `trace` under `policy`. Returns (value, converged, replicas).

    `skew` is {rid: offset} added to that replica's timestamps -- a clock that
    runs fast. Only the timestamp-reading structures notice it at all.
    """
    skew = skew or {}
    reps = _replicas(policy, n_replicas, **kw)
    ts = 0
    for e in trace:
        if e[0] == "op":
            _, rid, elem, label = e
            op = reps[rid].add if label == "add" else reps[rid].remove
            op(elem, ts + skew.get(rid, 0))
            ts += 1
        else:
            sync_all(reps)
    for _ in range(3):          # settle: merge is idempotent, so this is safe
        sync_all(reps)
    vals = [r.value() for r in reps]
    return vals[0], all(v == vals[0] for v in vals), reps


def run_gossip(seed, n_replicas, policy, gossip_prob, skews, n_ops=120,
               n_elems=6, p_add=0.55, rounds_at_end=60, **kw):
    """Ops and random pairwise gossip interleaved -- no sync barrier, so no
    replica is ever known to be up to date. Returns (value, converged, truth).

    Two RNGs on purpose: `orng` draws the operations and `grng` draws the
    gossip pairs, so raising `gossip_prob` cannot perturb the trace. Sharing
    one RNG here produced a fabricated result; see commentary 7.5.
    """
    orng = random.Random(seed)
    grng = random.Random(seed ^ 0x5EED)
    reps = _replicas(policy, n_replicas, **kw)
    truth = set()
    for i in range(n_ops):
        rid = orng.randrange(n_replicas)
        elem = "e%d" % orng.randrange(n_elems)
        label = "add" if orng.random() < p_add else "rm"
        op = reps[rid].add if label == "add" else reps[rid].remove
        op(elem, i + skews.get(rid, 0))
        truth.add(elem) if label == "add" else truth.discard(elem)
        n_exchanges = int(gossip_prob) + (grng.random() < gossip_prob % 1)
        for _ in range(n_exchanges):
            a, b = grng.sample(range(n_replicas), 2)
            reps[a].merge(reps[b])
    for _ in range(rounds_at_end):      # let the network drain
        a, b = grng.sample(range(n_replicas), 2)
        reps[a].merge(reps[b])
        reps[b].merge(reps[a])
    for _ in range(n_replicas * 4):
        sync_all(reps)
    vals = [r.value() for r in reps]
    return vals[0], all(v == vals[0] for v in vals), frozenset(truth)


def fmt(s):
    return "{" + ", ".join(sorted(s)) + "}" if s else "{}"
