/* inflate_stats.c -- a raw DEFLATE (RFC 1951) decoder that accounts for
 * every bit of the stream.
 *
 *   ./inflate_stats [-t] [-b] [-o out] file.deflate
 *     -t  trace every header field and symbol (for tiny inputs)
 *     -b  one line per block
 *     -o  write the decompressed bytes to out
 *
 * The last line is a key=value summary consumed by experiments.py.  For
 * dynamic blocks it also re-codes the block's own symbol counts with
 * unrestricted Huffman, package-merge (15 / 7 bits), the zlib overflow
 * repair, and the fixed code, and reports the empirical entropy.
 */
#include <stdio.h>
#include "huff.h"

static const uint16_t len_base[29] = {3, 4, 5, 6, 7, 8, 9, 10, 11, 13, 15, 17, 19, 23, 27, 31,
                                      35, 43, 51, 59, 67, 83, 99, 115, 131, 163, 195, 227, 258};
static const uint8_t len_extra[29] = {0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2,
                                      3, 3, 3, 3, 4, 4, 4, 4, 5, 5, 5, 5, 0};
static const uint16_t dist_base[30] = {1, 2, 3, 4, 5, 7, 9, 13, 17, 25, 33, 49, 65, 97, 129, 193,
                                       257, 385, 513, 769, 1025, 1537, 2049, 3073, 4097, 6145,
                                       8193, 12289, 16385, 24577};
static const uint8_t dist_extra[30] = {0, 0, 0, 0, 1, 1, 2, 2, 3, 3, 4, 4, 5, 5, 6, 6,
                                       7, 7, 8, 8, 9, 9, 10, 10, 11, 11, 12, 12, 13, 13};
static const uint8_t cl_order[19] = {16, 17, 18, 0, 8, 7, 9, 6, 10, 5, 11, 4, 12, 3, 13, 2, 14, 1, 15};

typedef struct {
    const uint8_t *p;
    size_t n;
    uint64_t pos;                   /* bit position */
} bitin;

static void die(const char *m)
{
    fprintf(stderr, "inflate_stats: %s\n", m);
    exit(2);
}

/* data elements other than Huffman codes: LSB first (RFC 1951 3.1.1) */
static uint32_t getbits(bitin *b, int n)
{
    uint32_t v = 0;
    for (int i = 0; i < n; i++) {
        if (b->pos >= 8 * (uint64_t)b->n) die("unexpected end of input");
        v |= (uint32_t)((b->p[b->pos >> 3] >> (b->pos & 7)) & 1) << i;
        b->pos++;
    }
    return v;
}

/* canonical decoding table: count[len] and symbols sorted by (len, value) */
typedef struct {
    uint16_t count[16];
    uint16_t sym[320];
    uint8_t len[320];
    int n;
} huffdec;

static int build(huffdec *h, const uint8_t *len, int n)
{
    uint16_t offs[16];
    memset(h->count, 0, sizeof h->count);
    h->n = n;
    for (int i = 0; i < n; i++) h->count[len[i]]++, h->len[i] = len[i];
    int left = 1;
    for (int l = 1; l < 16; l++) {
        left = 2 * left - h->count[l];
        if (left < 0) return -1;    /* oversubscribed */
    }
    offs[1] = 0;
    for (int l = 1; l < 15; l++) offs[l + 1] = (uint16_t)(offs[l] + h->count[l]);
    for (int i = 0; i < n; i++)
        if (len[i]) h->sym[offs[len[i]]++] = (uint16_t)i;
    return left;                    /* 0 = complete code */
}

/* Huffman codes: MSB of the code first, one bit at a time.  The
 * count/first/index loop is the method of zlib contrib/puff/puff.c. */
static int decode(bitin *b, const huffdec *h)
{
    int code = 0, first = 0, index = 0;
    for (int l = 1; l < 16; l++) {
        code |= (int)getbits(b, 1);
        int c = h->count[l];
        if (code - c < first) return h->sym[index + (code - first)];
        index += c;
        first = (first + c) << 1;
        code <<= 1;
    }
    die("invalid Huffman code");
    return -1;
}

typedef struct {
    uint64_t blocks[3], hdr, tree, stored, lit, len, lenx, dist, distx, eob;
    uint64_t nlit, nmatch, matchbytes, out;
    /* dynamic blocks only: symbol-code bits under alternative lengths */
    uint64_t dyn_actual, dyn_huff, dyn_pm15, dyn_zlib15, dyn_fixed, dyn_pm7cl, dyn_actcl;
    double dyn_entropy;
    uint64_t dyn_tree, dyn_tree_flat, dyn_syms, dyn_root9, dyn_dsyms, dyn_droot6, cl_sent;
    uint64_t dist_used1;            /* dynamic blocks using <= 1 distance code */
    int dyn_maxhuff, dyn_maxact;    /* longest lit/len code: unrestricted, actual */
} stats;

static stats S;
static int trace, perblock;
static uint8_t *outbuf;
static size_t outcap;

static void emit(uint8_t c)
{
    if (S.out == outcap) {
        outcap = outcap ? 2 * outcap : 1 << 16;
        outbuf = realloc(outbuf, outcap);
        if (!outbuf) die("out of memory");
    }
    outbuf[S.out++] = c;
}

static void fixed_lengths(uint8_t *ll, uint8_t *dl)
{
    int i = 0;
    for (; i < 144; i++) ll[i] = 8;
    for (; i < 256; i++) ll[i] = 9;
    for (; i < 280; i++) ll[i] = 7;
    for (; i < 288; i++) ll[i] = 8;
    for (i = 0; i < 30; i++) dl[i] = 5;
}

/* bits spent on literal/length and distance codes (no extra bits) */
static uint64_t sym_cost(const uint64_t *lc, const uint8_t *ll, const uint64_t *dc, const uint8_t *dl)
{
    return cost(lc, ll, 288) + cost(dc, dl, 30);
}

/* decode one block's symbols; count them and charge bits to categories */
static void codes(bitin *b, const huffdec *lh, const huffdec *dh, uint64_t *lc, uint64_t *dc)
{
    for (;;) {
        uint64_t p0 = b->pos;
        int sym = decode(b, lh);
        lc[sym]++;
        uint64_t cb = b->pos - p0;
        if (sym < 256) {
            S.lit += cb, S.nlit++;
            emit((uint8_t)sym);
            if (trace) printf("  lit  %3d %-4s %2llu bits\n", sym,
                              sym >= 32 && sym < 127 ? (char[]){'\'', (char)sym, '\'', 0} : "",
                              (unsigned long long)cb);
            continue;
        }
        if (sym == 256) {
            S.eob += cb;
            if (trace) printf("  eob  256      %2llu bits\n", (unsigned long long)cb);
            return;
        }
        if (sym > 285) die("invalid length symbol");
        S.len += cb;
        int li = sym - 257;
        uint64_t p1 = b->pos;
        int len = len_base[li] + (int)getbits(b, len_extra[li]);
        S.lenx += b->pos - p1;
        uint64_t p2 = b->pos;
        int ds = decode(b, dh);
        if (ds > 29) die("invalid distance symbol");
        dc[ds]++;
        S.dist += b->pos - p2;
        uint64_t p3 = b->pos;
        uint64_t dist = dist_base[ds] + getbits(b, dist_extra[ds]);
        S.distx += b->pos - p3;
        if (dist > S.out) die("distance too far back");
        if (trace)
            printf("  copy len %d (sym %d, %llu+%d bits) dist %llu (sym %d, %llu+%d bits)\n", len, sym,
                   (unsigned long long)cb, len_extra[li], (unsigned long long)dist, ds,
                   (unsigned long long)(p3 - p2), dist_extra[ds]);
        S.nmatch++, S.matchbytes += (uint64_t)len;
        for (int i = 0; i < len; i++) emit(outbuf[S.out - dist]);
    }
}

static void analyse_dynamic(const uint64_t *lc, const uint8_t *ll, const uint64_t *dc, const uint8_t *dl)
{
    uint8_t a[288], d[30], fl[288], fd[30];
    S.dyn_actual += sym_cost(lc, ll, dc, dl);
    huff_lengths(lc, 288, a), huff_lengths(dc, 30, d);
    S.dyn_huff += sym_cost(lc, a, dc, d);
    for (int i = 0; i < 288; i++) {
        if (a[i] > S.dyn_maxhuff) S.dyn_maxhuff = a[i];
        if (ll[i] > S.dyn_maxact) S.dyn_maxact = ll[i];
    }
    pm_lengths(lc, 288, 15, a), pm_lengths(dc, 30, 15, d);
    S.dyn_pm15 += sym_cost(lc, a, dc, d);
    zlib_lengths(lc, 288, 15, a), zlib_lengths(dc, 30, 15, d);
    S.dyn_zlib15 += sym_cost(lc, a, dc, d);
    fixed_lengths(fl, fd);
    S.dyn_fixed += sym_cost(lc, fl, dc, fd);
    S.dyn_entropy += entropy_bits(lc, 288) + entropy_bits(dc, 30);
    int used = 0;
    for (int i = 0; i < 30; i++) used += dc[i] != 0;
    S.dist_used1 += used <= 1;
    for (int i = 0; i < 288; i++) {
        S.dyn_syms += lc[i];
        if (ll[i] && ll[i] <= 9) S.dyn_root9 += lc[i];
    }
    for (int i = 0; i < 30; i++) {
        S.dyn_dsyms += dc[i];
        if (dl[i] && dl[i] <= 6) S.dyn_droot6 += dc[i];
    }
}

static void block_dynamic(bitin *b)
{
    uint64_t p0 = b->pos;
    int hlit = (int)getbits(b, 5) + 257, hdist = (int)getbits(b, 5) + 1, hclen = (int)getbits(b, 4) + 4;
    if (hlit > 286 || hdist > 30) die("bad HLIT/HDIST");
    uint8_t cll[19] = {0}, lens[320] = {0}, ll[288] = {0}, dl[30] = {0};
    for (int i = 0; i < hclen; i++) cll[cl_order[i]] = (uint8_t)getbits(b, 3);
    huffdec ch, lh, dh;
    if (build(&ch, cll, 19) != 0) die("incomplete code-length code");
    uint64_t ccount[19] = {0};
    int n = 0;
    while (n < hlit + hdist) {
        int sym = decode(b, &ch), rep = 0, val = 0;
        ccount[sym]++;
        if (sym < 16) { lens[n++] = (uint8_t)sym; continue; }
        if (sym == 16) {
            if (n == 0) die("repeat with no previous length");
            val = lens[n - 1], rep = 3 + (int)getbits(b, 2);
        } else if (sym == 17) rep = 3 + (int)getbits(b, 3);
        else rep = 11 + (int)getbits(b, 7);
        if (n + rep > hlit + hdist) die("too many lengths");
        while (rep--) lens[n++] = (uint8_t)val;
    }
    if (lens[256] == 0) die("no end-of-block code");
    memcpy(ll, lens, (size_t)hlit);
    memcpy(dl, lens + hlit, (size_t)hdist);
    int lleft = build(&lh, ll, 288), dleft = build(&dh, dl, 30);
    if (lleft < 0 || dleft < 0) die("oversubscribed code");
    uint64_t tree = b->pos - p0;
    S.tree += tree, S.dyn_tree += tree, S.cl_sent += (uint64_t)hclen;
    S.dyn_tree_flat += 14 + 4ull * (uint64_t)(hlit + hdist);
    uint8_t pm7[19];
    pm_lengths(ccount, 19, 7, pm7);
    S.dyn_actcl += cost(ccount, cll, 19), S.dyn_pm7cl += cost(ccount, pm7, 19);
    if (trace) {
        printf("  HLIT %d HDIST %d HCLEN %d, code-length lengths:", hlit, hdist, hclen);
        for (int i = 0; i < 19; i++) printf(" %d:%d", i, cll[i]);
        printf("\n  tree %llu bits\n", (unsigned long long)tree);
    }
    uint64_t lc[288] = {0}, dc[30] = {0};
    codes(b, &lh, &dh, lc, dc);
    analyse_dynamic(lc, ll, dc, dl);
}

static void block_stored(bitin *b)
{
    uint64_t p0 = b->pos;
    b->pos = (b->pos + 7) & ~7ull;
    if (b->pos + 32 > 8 * (uint64_t)b->n) die("truncated stored block");
    uint32_t len = getbits(b, 16), nlen = getbits(b, 16);
    if ((len ^ 0xffff) != nlen) die("stored LEN/NLEN mismatch");
    for (uint32_t i = 0; i < len; i++) emit((uint8_t)getbits(b, 8));
    S.stored += b->pos - p0;
    if (trace) printf("  stored %u bytes\n", len);
}

static void block_fixed(bitin *b)
{
    uint8_t ll[288], dl[30];
    huffdec lh, dh;
    uint64_t lc[288] = {0}, dc[30] = {0};
    fixed_lengths(ll, dl);
    build(&lh, ll, 288), build(&dh, dl, 30);
    codes(b, &lh, &dh, lc, dc);
}

static uint64_t cat_bits(void)
{
    return S.hdr + S.tree + S.stored + S.lit + S.len + S.lenx + S.dist + S.distx + S.eob;
}

int main(int argc, char **argv)
{
    const char *outname = NULL, *inname = NULL;
    for (int i = 1; i < argc; i++) {
        if (!strcmp(argv[i], "-t")) trace = 1;
        else if (!strcmp(argv[i], "-b")) perblock = 1;
        else if (!strcmp(argv[i], "-o") && i + 1 < argc) outname = argv[++i];
        else inname = argv[i];
    }
    if (!inname) die("usage: inflate_stats [-t] [-b] [-o out] file.deflate");
    FILE *f = fopen(inname, "rb");
    if (!f) die("cannot open input");
    fseek(f, 0, SEEK_END);
    long sz = ftell(f);
    fseek(f, 0, SEEK_SET);
    uint8_t *buf = malloc((size_t)sz + 1);
    if (!buf || fread(buf, 1, (size_t)sz, f) != (size_t)sz) die("read error");
    fclose(f);

    bitin b = {buf, (size_t)sz, 0};
    int last;
    do {
        uint64_t p0 = b.pos, out0 = S.out, t0 = S.tree;
        last = (int)getbits(&b, 1);
        int type = (int)getbits(&b, 2);
        S.hdr += 3;
        if (trace) printf("block BFINAL=%d BTYPE=%d\n", last, type);
        if (type == 3) die("reserved block type");
        S.blocks[type]++;
        if (type == 0) block_stored(&b);
        else if (type == 1) block_fixed(&b);
        else block_dynamic(&b);
        if (perblock)
            printf("block type=%d out=%llu bits=%llu tree=%llu\n", type,
                   (unsigned long long)(S.out - out0), (unsigned long long)(b.pos - p0),
                   (unsigned long long)(S.tree - t0));
    } while (!last);

    uint64_t used = b.pos, pad = 8 * ((used + 7) / 8) - used;
    if (cat_bits() != used) die("bit accounting mismatch");
    if (outname) {
        FILE *o = fopen(outname, "wb");
        if (!o || fwrite(outbuf, 1, S.out, o) != S.out || fclose(o)) die("write error");
    }
    printf("in_bytes=%ld used_bits=%llu pad=%llu trailing_bytes=%llu out_bytes=%llu "
           "stored_blocks=%llu fixed_blocks=%llu dyn_blocks=%llu "
           "hdr=%llu tree=%llu stored=%llu lit=%llu len=%llu lenx=%llu dist=%llu distx=%llu eob=%llu "
           "nlit=%llu nmatch=%llu matchbytes=%llu "
           "dyn_actual=%llu dyn_huff=%llu dyn_pm15=%llu dyn_zlib15=%llu dyn_fixed=%llu dyn_entropy=%.1f "
           "dyn_tree=%llu dyn_tree_flat=%llu dyn_actcl=%llu dyn_pm7cl=%llu cl_sent=%llu "
           "dyn_syms=%llu dyn_root9=%llu dyn_dsyms=%llu dyn_droot6=%llu dist_used1=%llu "
           "dyn_maxhuff=%d dyn_maxact=%d\n",
           sz, (unsigned long long)used, (unsigned long long)pad,
           (unsigned long long)((uint64_t)sz - (used + 7) / 8), (unsigned long long)S.out,
           (unsigned long long)S.blocks[0], (unsigned long long)S.blocks[1], (unsigned long long)S.blocks[2],
           (unsigned long long)S.hdr, (unsigned long long)S.tree, (unsigned long long)S.stored,
           (unsigned long long)S.lit, (unsigned long long)S.len, (unsigned long long)S.lenx,
           (unsigned long long)S.dist, (unsigned long long)S.distx, (unsigned long long)S.eob,
           (unsigned long long)S.nlit, (unsigned long long)S.nmatch, (unsigned long long)S.matchbytes,
           (unsigned long long)S.dyn_actual, (unsigned long long)S.dyn_huff, (unsigned long long)S.dyn_pm15,
           (unsigned long long)S.dyn_zlib15, (unsigned long long)S.dyn_fixed, S.dyn_entropy,
           (unsigned long long)S.dyn_tree, (unsigned long long)S.dyn_tree_flat,
           (unsigned long long)S.dyn_actcl, (unsigned long long)S.dyn_pm7cl, (unsigned long long)S.cl_sent,
           (unsigned long long)S.dyn_syms, (unsigned long long)S.dyn_root9,
           (unsigned long long)S.dyn_dsyms, (unsigned long long)S.dyn_droot6,
           (unsigned long long)S.dist_used1, S.dyn_maxhuff, S.dyn_maxact);
    free(buf);
    free(outbuf);
    return 0;
}
