"""A BtrBlocks-style cascading scheme selector (size model, not a file format).

Follows Kuschewski et al., SIGMOD 2023, Sections 3.1-3.2: 64,000-value blocks,
statistics-based pruning (no RLE if average run < 2, no Frequency if >= 50% of
values are unique), a sample of 10 runs x 64 values taken from 10 equal parts
of the block, at most 3 cascade levels, uncompressed once the depth is used up.
Simplifications: FOR + bit-packing per 128 values stands in for
SIMD-FastPFOR/FastBP128, the Frequency exception bitmap is a plain bitmap
instead of Roaring, and there is no FSST or Pseudodecimal.
"""
import numpy as np

BLOCK, RUNS, RUNLEN, DEPTH = 64000, 10, 64, 3


def _bits(x):
    return np.where(x > 0, np.floor(np.log2(np.maximum(x, 1))).astype(np.int64) + 1, 0)


def for_bp128(a, w):
    n = len(a)
    pad = (-n) % 128
    b = np.concatenate([a, np.repeat(a[-1:], pad)]).reshape(-1, 128)
    rng = (b.max(axis=1) - b.min(axis=1)).astype(np.uint64)
    widths = _bits(rng.astype(np.float64))
    return 1 + 4 + int(np.sum(w + 1 + (128 * widths + 7) // 8))


def _runs(a):
    starts = np.flatnonzero(np.concatenate([[True], a[1:] != a[:-1]]))
    lengths = np.diff(np.append(starts, len(a)))
    return a[starts], lengths


def int_schemes(a, w, stats):
    """Yield (name, fixed_bytes, [(child_array, child_width), ...])."""
    n, distinct, avg_run, _ = stats     # viability; sizes use len(a)
    m = len(a)
    yield "UNCOMP", 1 + w * m, None
    if distinct == 1:
        yield "ONE", 1 + w, []
        return
    yield "FOR_BP128", for_bp128(a, w), []
    if avg_run >= 2:
        vals, lens = _runs(a)
        yield "RLE", 1 + 4, [(vals, w), (lens, 4)]
    if distinct < n:
        uniq, codes = np.unique(a, return_inverse=True)
        yield "DICT", 1 + 4 + w * len(uniq), [(codes.astype(np.int64), 4)]
    if distinct / n < 0.5:
        vals, counts = np.unique(a, return_counts=True)
        top = vals[np.argmax(counts)]
        exc = a[a != top]
        yield "FREQ", 1 + w + 4 + (m + 7) // 8, [(exc, w)] if len(exc) else []


def int_stats(a):
    n = len(a)
    vals, counts = np.unique(a, return_counts=True)
    runs = 1 + int(np.count_nonzero(a[1:] != a[:-1]))
    return n, len(vals), n / runs, counts.max() / n


def best(a, w, depth):
    """Exhaustive: smallest cascade within `depth` levels. Returns (bytes, tree)."""
    if len(a) == 0:
        return 0, "EMPTY"
    if depth == 0:
        return 1 + w * len(a), "UNCOMP"
    out = None
    for name, fixed, children in int_schemes(a, w, int_stats(a)):
        size, sub = fixed, []
        for c, cw in children or []:
            s, t = best(c, cw, depth - 1)
            size += s
            sub.append(t)
        tree = name + ("(" + ",".join(sub) + ")" if sub else "")
        if out is None or size < out[0]:
            out = (size, tree)
    return out


def sample(a, rng):
    n = len(a)
    if n <= RUNS * RUNLEN:
        return a
    part = n // RUNS
    idx = [np.arange(s, s + RUNLEN) for s in
           (p * part + rng.integers(0, part - RUNLEN + 1) for p in range(RUNS))]
    return a[np.concatenate(idx)]


def sampled(a, w, depth, rng):
    """BtrBlocks-style: estimate every viable scheme on a sample, keep the best."""
    if len(a) == 0:
        return 0, "EMPTY"
    if depth == 0:
        return 1 + w * len(a), "UNCOMP"
    stats = int_stats(a)          # pruning uses statistics of the whole block
    s = sample(a, rng)            # ratios are estimated on the sample
    best_ratio, choice = -1.0, None
    for name, fixed, children in int_schemes(s, w, stats):
        est = fixed + sum(best(c, cw, depth - 1)[0] for c, cw in children or [])
        ratio = (w * len(s)) / est
        if ratio > best_ratio:
            best_ratio, choice = ratio, name
    for name, fixed, children in int_schemes(a, w, stats):
        if name != choice:
            continue
        size, sub = fixed, []
        for c, cw in children or []:
            sz, t = sampled(c, cw, depth - 1, rng)
            size += sz
            sub.append(t)
        return size, name + ("(" + ",".join(sub) + ")" if sub else "")
    return 1 + w * len(a), "UNCOMP"   # chosen scheme not viable on the full block


def str_column(vals, rng, use_sampling):
    """Strings: UNCOMP, ONE or DICT with integer codes cascaded."""
    arr = np.array(vals, dtype=object)
    n = len(arr)
    raw = 1 + 4 * (n + 1) + sum(len(v) for v in vals)
    uniq, codes = np.unique(arr, return_inverse=True)
    if len(uniq) == 1:
        return 1 + 4 + len(uniq[0]), "ONE"
    if len(uniq) == n:
        return raw, "UNCOMP"
    pick = sampled if use_sampling else (lambda c, w, d, r: best(c, w, d))
    cs, ct = pick(codes.astype(np.int64), 4, DEPTH - 1, rng)
    dsize = 1 + 4 * (len(uniq) + 1) + sum(len(v) for v in uniq) + cs
    return (dsize, "DICT(%s)" % ct) if dsize < raw else (raw, "UNCOMP")


def column(vals, w, seed=86):
    """Per-block sizes for a column. w = 4/8 for ints, 0 for byte strings."""
    rng = np.random.default_rng(seed)
    rows = []
    for b in range(0, len(vals), BLOCK):
        chunk = vals[b:b + BLOCK]
        if w:
            a = np.asarray(chunk, dtype=np.int64)
            opt = best(a, w, DEPTH)
            smp = sampled(a, w, DEPTH, rng)
        else:
            opt = str_column(chunk, rng, False)
            smp = str_column(chunk, rng, True)
        rows.append((opt, smp))
    return rows
