#!/usr/bin/env python3
"""S19 finite, serialized, in-memory bag maintenance; stdout only.
Integers are exact. Input batches are net changes, not SQL statement sequences.
Internal normalized-map primitives assume pure total key/projection functions.
No NULL, concurrent calls, durable storage, transport replay, or SQL parser.
"""
from collections import Counter
from copy import deepcopy
from dataclasses import dataclass
from fractions import Fraction
import json
import random


class Rejected(ValueError):
    pass


def require(condition, message):
    if not condition:
        raise Rejected(message)


def integer(x):
    return type(x) is int


def add_weight(out, row, weight):
    value = out.get(row, 0) + weight
    if value:
        out[row] = value
    else:
        out.pop(row, None)


def normalize(entries, schema):
    """Read every raw entry, consolidate signed weights, remove zero entries."""
    out = {}
    source = entries.items() if isinstance(entries, dict) else entries
    for row, weight in source:
        require(isinstance(row, tuple) and schema(row), 'row schema')
        require(integer(weight), 'integer weight required')
        add_weight(out, row, weight)
    return out


def order_schema(r):
    return (len(r) == 3 and type(r[0]) is str and type(r[1]) is str
            and integer(r[2]))


def tag_schema(r):
    return len(r) == 2 and all(type(x) is str for x in r)


def grouped_schema(r):
    return (len(r) == 3 and type(r[0]) is str and type(r[1]) is str
            and integer(r[2]))


def changed(base, delta):
    out = base.copy()
    for row, weight in delta.items():
        require(base.get(row, 0) + weight >= 0, ('negative snapshot', row))
        add_weight(out, row, weight)
    return out


def plus(*maps):
    out = {}
    for rows in maps:
        for row, weight in rows.items():
            add_weight(out, row, weight)
    return out


def opposite(rows):
    return {row: -weight for row, weight in rows.items()}


def signed_join(left, right, left_key, right_key, project, stats=None):
    """Bilinear primitive on normalized finite signed maps; nested support scan."""
    out = {}
    for r, nr in left.items():
        for s, ns in right.items():
            if stats is not None:
                stats['candidate_pairs'] += 1
            if left_key(r) == right_key(s):
                if stats is not None:
                    stats['matching_products'] += 1
                add_weight(out, project(r, s), nr * ns)
    return out


def query(r, s, stats=None):
    return signed_join(r, s, lambda x: x[1], lambda x: x[0],
                       lambda x, y: (x[0], y[1], x[2]), stats)


@dataclass(frozen=True)
class Group:
    n: int
    total: int
    freq: dict  # immutable by contract once published

    def row(self, key):
        require(self.n > 0 and bool(self.freq), 'nonempty group required')
        return (*key, self.n, self.total, Fraction(self.total, self.n),
                len(self.freq), min(self.freq))


def groups_from(rows):
    groups = {}
    for (customer, tag, value), weight in rows.items():
        require(weight > 0, 'positive initial bag')
        key = (customer, tag)
        n, total, freq = groups.get(key, (0, 0, {}))
        add_weight(freq, value, weight)
        groups[key] = (n + weight, total + weight * value, freq)
    return {key: Group(n, total, freq) for key, (n, total, freq) in groups.items()}


def prepare_groups(groups, delta):
    """Candidate state and consolidated old-row/new-row delta; no source edits."""
    by_key = {}
    for (customer, tag, value), weight in delta.items():
        add_weight(by_key.setdefault((customer, tag), {}), value, weight)
    candidate = groups.copy()
    output = {}
    for key, edits in by_key.items():
        old = groups.get(key)
        freq = old.freq.copy() if old else {}
        n, total = (old.n, old.total) if old else (0, 0)
        for value, weight in edits.items():
            require(freq.get(value, 0) + weight >= 0,
                    ('negative group frequency', key, value))
            add_weight(freq, value, weight)
            n += weight
            total += value * weight
        require(n >= 0, 'negative group count')
        if old:
            add_weight(output, old.row(key), -1)
        if n:
            new = Group(n, total, freq)
            candidate[key] = new
            add_weight(output, new.row(key), 1)
        else:
            require(not freq and total == 0, 'empty group invariant')
            candidate.pop(key, None)
    return candidate, output


@dataclass(frozen=True)
class State:
    epoch: int
    orders: dict
    tags: dict
    view: dict
    groups: dict


class Maintainer:
    def __init__(self, orders=(), tags=()):
        r = normalize(orders, order_schema)
        s = normalize(tags, tag_schema)
        require(all(w > 0 for w in r.values()) and all(w > 0 for w in s.values()),
                'negative initial snapshot')
        v = query(r, s)
        self.state = State(0, r, s, v, groups_from(v))

    def apply(self, expected_epoch, orders_delta=(), tags_delta=()):
        old = self.state
        require(integer(expected_epoch) and expected_epoch == old.epoch,
                'stale epoch')
        dr = normalize(orders_delta, order_schema)
        ds = normalize(tags_delta, tag_schema)
        r = changed(old.orders, dr)
        s = changed(old.tags, ds)
        stats = {'candidate_pairs': 0, 'matching_products': 0}
        terms = (query(dr, old.tags, stats), query(old.orders, ds, stats),
                 query(dr, ds, stats))
        dv = plus(*terms)
        v = changed(old.view, dv)
        groups, dg = prepare_groups(old.groups, dv)
        new = State(old.epoch + 1, r, s, v, groups)
        result = {'epoch': new.epoch, 'terms': terms, 'view_delta': dv,
                  'summary_delta': dg, 'stats': stats}
        self.state = new  # sole publication point, after all candidate work
        return result


def summaries(groups):
    return {group.row(key): 1 for key, group in groups.items()}


def global_summary(values):
    """Explicit mathematical global aggregate: sum(empty)=0; avg/min undefined."""
    values = list(values)
    n, total = len(values), sum(values)
    return (n, total, Fraction(total, n) if n else None,
            len(set(values)), min(values) if n else None)


def expanded_oracle(r, s):
    """Independent self-test path: enumerate literal occurrences, then group lists."""
    rr = [row for row, n in r.items() for _ in range(n)]
    ss = [row for row, n in s.items() for _ in range(n)]
    result = [(a, tag, amount) for a, product, amount in rr
              for other, tag in ss if product == other]
    view = dict(Counter(result))
    values = {}
    for a, tag, amount in result:
        values.setdefault((a, tag), []).append(amount)
    summary = {(*key, *global_summary(v)): 1 for key, v in values.items()}
    return view, summary


def verify(state):
    view, summary = expanded_oracle(state.orders, state.tags)
    require(state.view == view, 'occurrence join mismatch')
    require(summaries(state.groups) == summary, 'occurrence aggregate mismatch')
    for g in state.groups.values():
        require(g.n == sum(g.freq.values()) and
                g.total == sum(v * w for v, w in g.freq.items()) and
                all(w > 0 for w in g.freq.values()), 'group invariant')


def encode(value):
    if isinstance(value, Fraction):
        return str(value)
    if isinstance(value, dict):
        if all(isinstance(k, str) for k in value):
            return {k: encode(v) for k, v in value.items()}
        return [[encode(k), encode(v)] for k, v in sorted(value.items())]
    if isinstance(value, (tuple, list)):
        return [encode(x) for x in value]
    return value


def main():
    r = {('a', 'p', 10): 2, ('a', 'p', 30): 1, ('b', 'q', 5): 1}
    s = {('p', 'hot'): 2, ('p', 'sale'): 1, ('q', 'sale'): 1}
    dr = {('a', 'p', 10): -1, ('a', 'p', 20): 1, ('c', 'r', 7): 1}
    ds = {('p', 'hot'): -1, ('p', 'sale'): 1, ('r', 'sale'): 1}
    machine = Maintainer(r, s)
    initial = summaries(machine.state.groups)
    first = machine.apply(0, dr, ds)
    verify(machine.state)
    require(sum(machine.state.view.values()) == 11 and
            sum(row[2] * n for row, n in machine.state.view.items()) == 192,
            'main total')
    after_first = summaries(machine.state.groups)
    second = machine.apply(1, {('a', 'p', 10): -1})
    verify(machine.state)
    after_second = summaries(machine.state.groups)
    third = machine.apply(2, tags_delta={('p', 'hot'): -1})
    verify(machine.state)
    require(('a', 'hot') not in machine.state.groups, 'empty group retained')
    failed = 0
    for epoch, bad in [(3, {('a', 'p', 99): -1}), (2, {}),
                       (3, [(('a', 'p', 1), True)])]:
        old = machine.state
        image = deepcopy(old)
        try:
            machine.apply(epoch, bad)
        except Rejected:
            failed += 1
        else:
            raise RuntimeError('invalid batch accepted')
        require(machine.state is old and old == image, 'partial publication')
    mini = Maintainer({('g', 'p', v): 1 for v in [1, 2, 5]}, {('p', 't'): 1})
    silent = mini.apply(0, {('g', 'p', 2): -1, ('g', 'p', 5): -1,
                            ('g', 'p', 3): 1, ('g', 'p', 4): 1})
    require(silent['summary_delta'] == {}, 'equal summary should cancel')
    mini.apply(1, {('g', 'p', 1): -1})
    require(min(mini.state.groups[('g', 't')].freq) == 3, 'silent state not advanced')
    old = machine.state
    image = deepcopy(old)
    try:
        machine.apply(3, {('z', 'p', 8): 1}, {('missing', 't'): -1})
    except Rejected:
        failed += 1
    else:
        raise RuntimeError('late invalid side accepted')
    require(machine.state is old and old == image, 'partial base-table change')
    rng = random.Random(1919)
    rdomain = [(c, p, v) for c in ['a', 'b'] for p in ['p', 'q']
               for v in [-2, 0, 3]]
    sdomain = [(p, tag) for p in ['p', 'q'] for tag in ['hot', 'sale']]
    checks = 0
    for _ in range(300):
        m = Maintainer()
        for step in range(20):
            old = m.state
            nr, ns = old.orders.copy(), old.tags.copy()
            for target, domain in [(nr, rdomain), (ns, sdomain)]:
                for _ in range(rng.randrange(4)):
                    row = rng.choice(domain)
                    weight = rng.randrange(4)
                    if weight:
                        target[row] = weight
                    else:
                        target.pop(row, None)
            dr = plus(nr, opposite(old.orders))
            ds = plus(ns, opposite(old.tags))
            event = m.apply(step, dr, ds)
            verify(m.state)
            require(plus(query(dr, old.tags), query(nr, ds)) == event['view_delta'],
                    'sequential formula')
            require(changed(summaries(old.groups), event['summary_delta']) ==
                    summaries(m.state.groups), 'summary retraction mismatch')
            checks += 1
    # Self-join uses the same base relation in two logical input slots.
    self_tests = 0
    for n in range(4):
        for new_n in range(4):
            base = {('p', 1): n} if n else {}
            updated = {('p', 1): new_n} if new_n else {}
            d = plus(updated, opposite(base))
            j = lambda a, b: signed_join(a, b, lambda x: x[0], lambda x: x[0],
                                        lambda x, y: (x[0], x[1], y[1]))
            delta = plus(j(d, base), j(base, d), j(d, d))
            require(changed(j(base, base), delta) == j(updated, updated), 'self join')
            self_tests += 1
    print(json.dumps(encode({'status': 'PASS', 'initial': initial, 'batch1': first,
          'after1': after_first, 'batch2': second, 'after2': after_second,
          'batch3': third, 'after3': summaries(machine.state.groups),
          'rejected_without_publication': failed, 'random_committed_batches': checks,
          'self_join_cases': self_tests, 'silent_summary_delta': silent['summary_delta'],
          'after_silent_then_delete_min': min(mini.state.groups[('g', 't')].freq),
          'zero_sum_nonempty': global_summary([-5, 5]),
          'global_empty': global_summary([])}),
          ensure_ascii=False, indent=2))


if __name__ == '__main__':
    main()
