#!/usr/bin/env python3
"""Exact finite checks for permutation weights. Python 3, standard library only.

Run without arguments for all finite tests and a JSON report. These are tests of
stated ranges, not a substitute for the proofs in the accompanying pages.
Polynomials are lists of integer coefficients in increasing power order.
"""
from collections import Counter
from fractions import Fraction
from functools import lru_cache
from itertools import combinations, permutations, product
import json


def descents(w):
    return tuple(i + 1 for i in range(len(w) - 1) if w[i] > w[i + 1])


def major(w):
    return sum(descents(w))


def inversions(w):
    return sum(a > b for i, a in enumerate(w) for b in w[i + 1:])


def gamma(r, x):
    """Cut after letters on the same threshold side as the last letter."""
    if not r:
        return ()
    low = r[-1] <= x
    blocks, start = [], 0
    for i, y in enumerate(r):
        if (y <= x) == low:
            blocks.append(r[start:i + 1])
            start = i + 1
    assert start == len(r)
    return tuple(y for block in blocks for y in block[-1:] + block[:-1])


def ungamma(v, x):
    """The first letter identifies the branch; cut before its threshold side."""
    if not v:
        return ()
    low = v[0] <= x
    cuts = [i for i, y in enumerate(v) if (y <= x) == low] + [len(v)]
    blocks = [v[cuts[i]:cuts[i + 1]] for i in range(len(cuts) - 1)]
    return tuple(y for block in blocks for y in block[1:] + block[:1])


def foata(w):
    r = ()
    for x in w:
        r = gamma(r, x) + (x,)
    return r


def inverse_foata(w):
    r, removed = tuple(w), []
    while r:
        x = r[-1]
        removed.append(x)
        r = ungamma(r[:-1], x)
    return tuple(reversed(removed))


def inverse_descents(w):
    pos = {x: i for i, x in enumerate(w)}
    letters = sorted(w)
    return tuple((a, b) for a, b in zip(letters, letters[1:]) if pos[a] > pos[b])


def trim(p):
    p = list(p)
    while len(p) > 1 and p[-1] == 0:
        p.pop()
    return p


def add(a, b):
    return trim([(a[i] if i < len(a) else 0) + (b[i] if i < len(b) else 0)
                 for i in range(max(len(a), len(b)))])


def multiply(a, b):
    c = [0] * (len(a) + len(b) - 1)
    for i, x in enumerate(a):
        for j, y in enumerate(b):
            c[i + j] += x * y
    return trim(c)


def shift(a, k):
    return [0] * k + list(a) if any(a) else [0]


def polynomial_from_values(values):
    c = Counter(values)
    return [c[i] for i in range(max(c, default=0) + 1)]


@lru_cache(None)
def gaussian(n, k):
    if k < 0 or k > n:
        return (0,)
    if k == 0 or k == n:
        return (1,)
    return tuple(add(gaussian(n - 1, k), shift(gaussian(n - 1, k - 1), n - k)))


def qmultinomial(counts):
    n, p = sum(counts), [1]
    for a in counts:
        p = multiply(p, gaussian(n, a))
        n -= a
    return p


def eulerian_row(n):
    if n == 0:
        return [[1]]
    row = [[1]]
    for size in range(2, n + 1):
        nxt = []
        for k in range(size):
            keep = multiply([1] * (k + 1), row[k]) if k < len(row) else [0]
            gain = shift(multiply([1] * (size - k), row[k - 1]), k) if k else [0]
            nxt.append(add(keep, gain))
        row = nxt
    return row


def binomial_polynomial(z, n):
    p = Fraction(1)
    for i in range(n):
        p *= Fraction(z - i, i + 1)
    return p


def closure(n, edges):
    r = set(edges)
    for k in range(n):
        for i in range(n):
            for j in range(n):
                if (i, k) in r and (k, j) in r:
                    r.add((i, j))
    return frozenset(r)


def extensions(n, relations):
    result = []
    for w in permutations(range(n)):
        pos = {x: i for i, x in enumerate(w)}
        if all(pos[a] < pos[b] for a, b in relations):
            result.append(w)
    return result


def partition_denominator(n, degree):
    p = [1] + [0] * degree
    for i in range(1, n + 1):
        for j in range(i, degree + 1):
            p[j] += p[j - i]
    return p


def shuffles(u, v):
    if not u:
        yield v
    elif not v:
        yield u
    else:
        for tail in shuffles(u[1:], v):
            yield u[:1] + tail
        for tail in shuffles(u, v[1:]):
            yield v[:1] + tail


def run_checks():
    counts = Counter()
    for n in range(8):
        rows = [Counter() for _ in range(max(n, 1))]
        images = set()
        for w in permutations(range(1, n + 1)):
            v = foata(w)
            assert inverse_foata(v) == w
            assert foata(inverse_foata(w)) == w
            assert inversions(v) == major(w)
            assert v[-1:] == w[-1:]
            assert inverse_descents(v) == inverse_descents(w)
            images.add(v)
            rows[len(descents(w))][major(w)] += 1
            counts['permutation_map_cases'] += 1
        assert len(images) == sum(sum(c.values()) for c in rows)
        expected = eulerian_row(n)
        assert [[c[i] for i in range(max(c, default=0) + 1)] for c in rows] == expected
        counts['q_eulerian_rows'] += 1

    for n in range(8):
        grouped = {}
        for w in product(range(3), repeat=n):
            v = foata(w)
            assert inverse_foata(v) == w and foata(inverse_foata(w)) == w
            assert Counter(v) == Counter(w) and inversions(v) == major(w)
            key = tuple(w.count(i) for i in range(3))
            if key not in grouped:
                grouped[key] = (Counter(), Counter())
            grouped[key][0][inversions(w)] += 1
            grouped[key][1][major(w)] += 1
            counts['repeated_word_map_cases'] += 1
        for key, (a, b) in grouped.items():
            assert a == b
            assert [a[i] for i in range(max(a, default=0) + 1)] == qmultinomial(key)
            counts['multiset_distribution_cases'] += 1

    for n in range(10):
        for k in range(n + 1):
            words = []
            for selected in combinations(range(n), k):
                w = tuple(0 if i in selected else 1 for i in range(n))
                b = [sum(y == 1 for y in w[:i]) for i in selected]
                assert b == sorted(b) and sum(b) == inversions(w)
                words.append(inversions(w))
                counts['binary_partition_bijection_cases'] += 1
            p = list(gaussian(n, k))
            assert polynomial_from_values(words) == p
            assert p == p[::-1] == list(gaussian(n, n - k))
            counts['gaussian_polynomial_cases'] += 1
        # Finite q-binomial theorem, represent coefficient of z^k as a q list.
        product_coefficients = [[1]]
        for i in range(n):
            new = [[0] for _ in range(len(product_coefficients) + 1)]
            for k, p in enumerate(product_coefficients):
                new[k] = add(new[k], p)
                new[k + 1] = add(new[k + 1], shift(p, i))
            product_coefficients = new
        assert product_coefficients == [shift(gaussian(n, k), k * (k - 1) // 2) for k in range(n + 1)]
        counts['finite_q_binomial_theorem_cases'] += 1

    # Coefficients of the weighted Worpitzky identity, no numerical q substitution.
    for n in range(7):
        row = eulerian_row(n)
        for m in range(7):
            left = [1]
            for _ in range(n):
                left = multiply(left, [1] * (m + 1))
            den = [[0] for _ in range(m + 1)]
            den[0] = [1]
            for j in range(n + 1):
                for power in range(1, m + 1):
                    den[power] = add(den[power], shift(den[power - 1], j))
            right = [0]
            for d, p in enumerate(row):
                if d <= m:
                    right = add(right, multiply(p, den[m - d]))
            assert left == right
            counts['weighted_worpitzky_coefficients'] += 1

    for n in range(7):
        seen = set()
        for w in permutations(range(1, n + 1)):
            for r in range(n + 1):
                u, v = w[:r], w[r:]
                if (u, v) in seen:
                    continue
                seen.add((u, v))
                actual = polynomial_from_values(major(x) for x in shuffles(u, v))
                expected = shift(gaussian(n, r), major(u) + major(v))
                assert actual == expected
                counts['arbitrary_label_shuffle_cases'] += 1

    # Every naturally ordered poset through four vertices; all complementary
    # labels are handled by the proof, natural and reversed labels tested here.
    for n in range(5):
        possible = list(combinations(range(n), 2))
        posets = {closure(n, [e for i, e in enumerate(possible) if mask >> i & 1])
                  for mask in range(1 << len(possible))}
        for rel in posets:
            ext = extensions(n, rel)
            ds = [len(descents(w)) for w in ext]
            for m in range(5):
                weak = strict = 0
                for f in product(range(1, m + 1), repeat=n):
                    weak += all(f[a] <= f[b] for a, b in rel)
                    strict += all(f[a] < f[b] for a, b in rel)
                wp = sum(binomial_polynomial(m + n - 1 - d, n) for d in ds)
                sp = sum(binomial_polynomial(m + d, n) for d in ds)
                negative = sum(binomial_polynomial(-m + n - 1 - d, n) for d in ds)
                assert weak == wp and strict == sp == (-1) ** n * negative
                counts['order_polynomial_parameter_cases'] += 1
            if n <= 3:
                for labels in permutations(range(1, n + 1)):
                    bound = 7
                    actual = [0] * (bound + 1)
                    for f in product(range(bound + 1), repeat=n):
                        if sum(f) > bound:
                            continue
                        if all(f[a] >= f[b] and (labels[a] < labels[b] or f[a] > f[b]) for a, b in rel):
                            actual[sum(f)] += 1
                    denominator = partition_denominator(n, bound)
                    expected = [0] * (bound + 1)
                    for w in ext:
                        m = major(tuple(labels[i] for i in w))
                        for j in range(m, bound + 1):
                            expected[j] += denominator[j - m]
                    assert actual == expected
                    counts['labeled_poset_series_cases'] += 1

    # Weighted descent-set inclusion/exclusion in the deepened existing page.
    for n in range(1, 7):
        allperms = list(permutations(range(n)))
        for mask in range(1 << (n - 1)):
            S = {i + 1 for i in range(n - 1) if mask >> i & 1}
            cuts = [0] + sorted(S) + [n]
            sizes = [b - a for a, b in zip(cuts, cuts[1:])]
            alpha = polynomial_from_values(inversions(w) for w in allperms if set(descents(w)) <= S)
            assert alpha == qmultinomial(sizes)
            counts['weighted_descent_subset_cases'] += 1

    w = (3, 1, 4, 2, 6, 5)
    assert foata(foata(w)) != w
    assert list(gaussian(4, 2)) == [1, 1, 2, 1, 1]
    assert sum(a * (-1) ** i for i, a in enumerate(gaussian(4, 2))) == 2
    return {
        'status': 'passed',
        'scope': 'Exact finite tests only; unrestricted theorems require the written proofs',
        'counts': dict(counts),
        'endpoint': {
            'input': list(w), 'image': list(foata(w)), 'inverse': list(inverse_foata(foata(w))),
            'input_major': major(w), 'input_inversions': inversions(w),
            'image_inversions': inversions(foata(w)),
            'binary_2_2': list(gaussian(4, 2)), 'multiset_2_1_1': qmultinomial((2, 1, 1)),
            'q_eulerian_4': eulerian_row(4),
            'shuffle_31_42': [[list(x), major(x), inversions(x)] for x in shuffles((3, 1), (4, 2))],
            'diamond_weak_2': 6, 'diamond_strict_4': 6, 'V_strict_3': 5
        }
    }


if __name__ == '__main__':
    print(json.dumps(run_checks(), ensure_ascii=False, indent=2))
