#!/usr/bin/env python3
"""Byte- and bit-level experiments for the Huffman/DEFLATE article.

    BUILD_DIR=/tmp/hd python3 experiments.py

Builds huff_test and inflate_stats with gcc, downloads the Canterbury corpus
(sha256 checked) into BUILD_DIR, compresses every file with the zlib linked
into this Python and, if the `zopfli` package is importable, with zopfli.
Every stream is decoded by inflate_stats and compared with the original.
Writes results/*.txt.  Nothing here measures time.
"""
import hashlib
import os
import platform
import random
import subprocess
import sys
import tarfile
import urllib.request
import zlib

HERE = os.path.dirname(os.path.abspath(__file__))
RES = os.path.join(HERE, "results")
BUILD = os.environ.get("BUILD_DIR", "/tmp/huffman-deflate-build")
CORPUS_URL = "https://corpus.canterbury.ac.nz/resources/cantrbry.tar.gz"
CORPUS_SHA = "f140e8a5b73d3f53198555a63bfb827889394a42f20825df33c810c3d5e3f8fb"
FILES = ["alice29.txt", "asyoulik.txt", "cp.html", "fields.c", "grammar.lsp", "kennedy.xls",
         "lcet10.txt", "plrabn12.txt", "ptt5", "sum", "xargs.1"]
CFLAGS = ["-std=c11", "-O2", "-Wall", "-Wextra", "-Wno-unused-function"]

try:
    import zopfli.zopfli as zop
except ImportError:
    zop = None


def sh(*args, **kw):
    return subprocess.run(args, check=True, capture_output=True, text=True, **kw).stdout


def build():
    os.makedirs(BUILD, exist_ok=True)
    for prog in ("huff_test", "inflate_stats"):
        sh("gcc", *CFLAGS, "-o", os.path.join(BUILD, prog), os.path.join(HERE, prog + ".c"), "-lm")


def corpus():
    tgz = os.path.join(BUILD, "cantrbry.tar.gz")
    if not os.path.exists(tgz):
        urllib.request.urlretrieve(CORPUS_URL, tgz)
    if hashlib.sha256(open(tgz, "rb").read()).hexdigest() != CORPUS_SHA:
        sys.exit("corpus checksum mismatch")
    d = os.path.join(BUILD, "cantrbry")
    if not os.path.isdir(d):
        with tarfile.open(tgz) as t:
            t.extractall(d)
    return {f: open(os.path.join(d, f), "rb").read() for f in FILES}


def deflate(data, level=9, strategy=zlib.Z_DEFAULT_STRATEGY):
    c = zlib.compressobj(level, zlib.DEFLATED, -15, 8, strategy)
    return c.compress(data) + c.flush()


def zopfli_raw(data):
    z = zop.compress(data, numiterations=15)       # zlib container, library defaults
    assert z[1] & 0x20 == 0                        # no preset dictionary
    return z[2:-4]


def stats(stream, data, tag):
    """decode with inflate_stats, check the output, return the summary dict"""
    src = os.path.join(BUILD, "s.deflate")
    out = os.path.join(BUILD, "s.out")
    open(src, "wb").write(stream)
    line = sh(os.path.join(BUILD, "inflate_stats"), "-o", out, src).strip().splitlines()[-1]
    if open(out, "rb").read() != data:
        sys.exit("round trip failed: " + tag)
    d = {}
    for kv in line.split():
        k, v = kv.split("=")
        d[k] = float(v) if "." in v else int(v)
    return d


METHODS = [("huffman-only", lambda b: deflate(b, 9, zlib.Z_HUFFMAN_ONLY)),
           ("fixed", lambda b: deflate(b, 9, zlib.Z_FIXED)),
           ("zlib-1", lambda b: deflate(b, 1)),
           ("zlib-6", lambda b: deflate(b, 6)),
           ("zlib-9", lambda b: deflate(b, 9))]


def pct(a, b):
    return 100.0 * a / b if b else 0.0


def env():
    lines = [f"python {platform.python_version()}",
             f"zlib header {zlib.ZLIB_VERSION}, runtime {zlib.ZLIB_RUNTIME_VERSION}",
             "zopfli " + (sh(sys.executable, "-m", "pip", "show", "zopfli").split("\n")[1] if zop else "not installed"),
             sh("gcc", "--version").splitlines()[0],
             f"kernel {platform.release()} {platform.machine()}"]
    return "\n".join(lines) + "\n"


def example():
    data = b"abcabcabcabc"
    raw = deflate(data, 6)
    src = os.path.join(BUILD, "ex.deflate")
    open(src, "wb").write(raw)
    out = [f"input {data!r}", "raw deflate: " + raw.hex(" ")]
    out.append(sh(os.path.join(BUILD, "inflate_stats"), "-t", src))
    for name, wbits in (("zlib", 15), ("gzip", 31)):
        c = zlib.compressobj(6, zlib.DEFLATED, wbits)
        s = c.compress(data) + c.flush()
        out.append(f"{name}: {len(s)} bytes: " + s.hex(" "))
    for s in (b"AABCAAB",):
        r = deflate(s, 9)
        open(src, "wb").write(r)
        t = sh(os.path.join(BUILD, "inflate_stats"), "-t", src)
        out.append(f"input {s!r} level 9: {len(r)} bytes, {t.count('copy')} copies, {t.count('lit ')} literals")
    return "\n".join(out) + "\n"


def main():
    build()
    os.makedirs(RES, exist_ok=True)
    open(os.path.join(RES, "env.txt"), "w").write(env())
    open(os.path.join(RES, "huff_test.txt"), "w").write(sh(os.path.join(BUILD, "huff_test")))
    open(os.path.join(RES, "example.txt"), "w").write(example())

    files = corpus()
    methods = list(METHODS) + ([("zopfli", zopfli_raw)] if zop else [])
    st = {}
    for f, data in files.items():
        for m, fn in methods:
            st[f, m] = stats(fn(data), data, f"{f}/{m}")
    names = [m for m, _ in methods]

    # sizes
    L = ["# raw DEFLATE bytes (no container); every stream decoded and compared",
         "file".ljust(13) + "original".rjust(10) + "".join(n.rjust(14) for n in names)]
    tot = {n: 0 for n in names}
    for f, data in files.items():
        row = f.ljust(13) + str(len(data)).rjust(10)
        for n in names:
            row += str(st[f, n]["in_bytes"]).rjust(14)
            tot[n] += st[f, n]["in_bytes"]
        L.append(row)
    orig = sum(len(d) for d in files.values())
    L.append("total".ljust(13) + str(orig).rjust(10) + "".join(str(tot[n]).rjust(14) for n in names))
    L.append("ratio".ljust(13) + "1.0000".rjust(10) + "".join(f"{tot[n] / orig:14.4f}" for n in names))
    L.append("")
    L.append("# blocks per method, summed over files: stored/fixed/dynamic")
    for n in names:
        b = [sum(st[f, n][k] for f in files) for k in ("stored_blocks", "fixed_blocks", "dyn_blocks")]
        L.append(f"{n}: {b[0]}/{b[1]}/{b[2]}")
    open(os.path.join(RES, "sizes.txt"), "w").write("\n".join(L) + "\n")

    # where the bits go
    L = ["# share of stream bits; tree = dynamic block headers after the 3-bit block header",
         "# lit = literal codes; lenc/lenx = length codes / extra bits; distc/distx = distance codes / extra bits"]
    for n in [x for x in ("zlib-9", "zopfli") if x in names]:
        L.append(f"\n[{n}]")
        L.append("file".ljust(13) + "blocks".rjust(7) + "tree%".rjust(7) + "lit%".rjust(7) + "lenc%".rjust(7)
                 + "lenx%".rjust(7) + "distc%".rjust(7) + "distx%".rjust(7) + "other%".rjust(7)
                 + "literals".rjust(9) + "matches".rjust(8) + "B/match".rjust(8))
        agg = {}
        for f in files:
            d = st[f, n]
            for k, v in d.items():
                agg[k] = agg.get(k, 0) + v
            L.append(anat_row(f, d))
        L.append(anat_row("all", agg))
        L.append("all, bits: " + " ".join(f"{k}={agg[k]}" for k in
                 ("used_bits", "hdr", "tree", "lit", "len", "lenx", "dist", "distx", "eob", "stored")))
    open(os.path.join(RES, "anatomy.txt"), "w").write("\n".join(L) + "\n")

    # entropy-coder efficiency inside dynamic blocks
    L = ["# dynamic blocks only; symbol-code bits (literal/length + distance codes, no extra bits)",
         "# each column re-codes the same per-block symbol counts; values are bits",
         "# entropy = sum c*log2(N/c) per block and alphabet"]
    for n in [x for x in ("zlib-9", "zopfli") if x in names]:
        L.append(f"\n[{n}]")
        L.append("file".ljust(13) + "entropy".rjust(10) + "huffman".rjust(10) + "pm15".rjust(10)
                 + "zlib15port".rjust(11) + "actual".rjust(10) + "fixed".rjust(10)
                 + "act/ent-1%".rjust(11) + "act/huf-1%".rjust(11) + "tree".rjust(8) + "flat4".rjust(8))
        agg = {}
        for f in files:
            d = st[f, n]
            for k, v in d.items():
                agg[k] = agg.get(k, 0) + v
            L.append(red_row(f, d))
        L.append(red_row("all", agg))
        L.append(f"code-length code: actual {agg['dyn_actcl']} bits, package-merge(7) {agg['dyn_pm7cl']} bits; "
                 f"mean HCLEN+4 = {agg['cl_sent'] / agg['dyn_blocks']:.2f}")
        L.append(f"literal/length symbols with code length <= 9: {pct(agg['dyn_root9'], agg['dyn_syms']):.2f}%; "
                 f"distance symbols with length <= 6: {pct(agg['dyn_droot6'], agg['dyn_dsyms']):.2f}%")
        L.append(f"dynamic blocks using at most one distance code: {agg['dist_used1']} of {agg['dyn_blocks']}")
    open(os.path.join(RES, "redundancy.txt"), "w").write("\n".join(L) + "\n")

    # length limiting: inputs whose unrestricted Huffman code is deeper than 15
    L = ["# Z_HUFFMAN_ONLY, one dynamic block each; literal/length code bits for the block's own counts",
         "# fib-k: k-1 byte values with counts 1,2,3,5,... plus end-of-block (count 1), shuffled; geo-r: 16000 bytes, P(i) ~ r^i over 256 values",
         "input".ljust(10) + "symbols".rjust(8) + "maxlen".rjust(7) + "huffman".rjust(9) + "pm15".rjust(9)
         + "zlib15port".rjust(11) + "zlib".rjust(9) + "zlib_max".rjust(9) + "zlib/pm15-1%".rjust(13)]
    rnd = random.Random(1)
    cases = []
    for k in (16, 17, 18, 19):
        fib = [1, 1]
        while len(fib) < k:
            fib.append(fib[-1] + fib[-2])
        buf = bytearray()
        for i, c in enumerate(fib[1:]):           # end-of-block supplies the first 1
            buf += bytes([65 + i]) * c
        rnd.shuffle(buf)
        cases.append((f"fib-{k}", bytes(buf)))
    for r in (0.5, 0.6, 0.7):
        w = [r ** i for i in range(256)]
        cases.append((f"geo-{r}", bytes(rnd.choices(range(256), weights=w, k=16000))))
    for name, data in cases:
        d = stats(deflate(data, 9, zlib.Z_HUFFMAN_ONLY), data, name)
        assert d["dyn_blocks"] == 1
        L.append(name.ljust(10) + str(len(set(data)) + 1).rjust(8) + str(d["dyn_maxhuff"]).rjust(7)
                 + str(d["dyn_huff"]).rjust(9) + str(d["dyn_pm15"]).rjust(9) + str(d["dyn_zlib15"]).rjust(11)
                 + str(d["dyn_actual"]).rjust(9) + str(d["dyn_maxact"]).rjust(9)
                 + f"{pct(d['dyn_actual'] - d['dyn_pm15'], d['dyn_pm15']):13.4f}")
    L.append("# symbols includes end-of-block; maxlen = deepest unrestricted Huffman code")
    open(os.path.join(RES, "limit.txt"), "w").write("\n".join(L) + "\n")

    # small inputs: which block type wins
    a = files["alice29.txt"]
    L = ["# prefixes of alice29.txt; raw DEFLATE bytes; type = block types zlib -9 chose (S/F/D)",
         "n".rjust(6) + "stored".rjust(8) + "huff-only".rjust(10) + "fixed".rjust(8) + "zlib-9".rjust(8)
         + "type".rjust(6) + "tree".rjust(6) + "lits".rjust(6) + "matches".rjust(8) + ("zopfli".rjust(8) if zop else "")]
    for k in (16, 32, 64, 128, 256, 512, 1024, 2048, 4096, 8192):
        p = a[:k]
        d9 = stats(deflate(p, 9), p, "small")
        t = "S" * d9["stored_blocks"] + "F" * d9["fixed_blocks"] + "D" * d9["dyn_blocks"]
        row = (str(k).rjust(6) + str(k + 5).rjust(8) + str(len(deflate(p, 9, zlib.Z_HUFFMAN_ONLY))).rjust(10)
               + str(len(deflate(p, 9, zlib.Z_FIXED))).rjust(8) + str(d9["in_bytes"]).rjust(8) + t.rjust(6)
               + str(d9["tree"]).rjust(6) + str(d9["nlit"]).rjust(6) + str(d9["nmatch"]).rjust(8))
        if zop:
            row += str(len(zopfli_raw(p))).rjust(8)
        L.append(row)
    L.append("# stored = n + 5: 3 header bits padded to a byte, then LEN and NLEN")
    L.append("# tree = dynamic-header bits of the zlib -9 stream; lits/matches = its LZ77 symbols")
    open(os.path.join(RES, "small.txt"), "w").write("\n".join(L) + "\n")


def anat_row(name, d):
    tot = d["used_bits"]
    other = d["hdr"] + d["eob"] + d["stored"]
    bpm = d["matchbytes"] / d["nmatch"] if d["nmatch"] else 0
    return (name.ljust(13) + str(d["stored_blocks"] + d["fixed_blocks"] + d["dyn_blocks"]).rjust(7)
            + "".join(f"{pct(d[k], tot):7.2f}" for k in ("tree", "lit", "len", "lenx", "dist", "distx"))
            + f"{pct(other, tot):7.2f}" + str(d["nlit"]).rjust(9) + str(d["nmatch"]).rjust(8) + f"{bpm:8.2f}")


def red_row(name, d):
    return (name.ljust(13) + f"{d['dyn_entropy']:10.0f}" + str(d["dyn_huff"]).rjust(10) + str(d["dyn_pm15"]).rjust(10)
            + str(d["dyn_zlib15"]).rjust(11) + str(d["dyn_actual"]).rjust(10) + str(d["dyn_fixed"]).rjust(10)
            + f"{pct(d['dyn_actual'] - d['dyn_entropy'], d['dyn_entropy']):11.3f}"
            + f"{pct(d['dyn_actual'] - d['dyn_huff'], d['dyn_huff']):11.4f}"
            + str(d["dyn_tree"]).rjust(8) + str(d["dyn_tree_flat"]).rjust(8))


if __name__ == "__main__":
    main()
