/* fjp.c - Fast JPEG-lite lossy image codec
 *
 * Constraints:
 *   - <1000 lines of C
 *   - Integer arithmetic only (no float/double)
 *   - Faster encode/decode than JPEG in practice
 *   - Target: >=30 dB PSNR at >=4x compression
 *
 * Compile:  gcc -O2 -std=c99 -o fjp fjp.c
 * Encode:   ./fjp e input.ppm output.fjp [quality]
 * Decode:   ./fjp d input.fjp output.ppm
 */

#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <stdint.h>

/* =========================================================================
 * Bit stream
 * ========================================================================= */

typedef struct { uint8_t *buf; size_t cap, pos; } BWr;

static void bw_init(BWr *b, size_t cap) {
    b->buf = (uint8_t *)calloc(cap, 1);
    if (!b->buf) { fprintf(stderr, "OOM\n"); exit(1); }
    b->cap = cap; b->pos = 0;
}
static void bw_put(BWr *b, int bit) {
    size_t by = b->pos >> 3;
    if (by < b->cap && bit) b->buf[by] |= (uint8_t)(1u << (7 - (b->pos & 7)));
    b->pos++;
}
static void bw_bits(BWr *b, uint32_t v, int n) {
    for (int i = n - 1; i >= 0; i--) bw_put(b, (int)((v >> i) & 1));
}
/* Signed value: n zeros, a 1, then n-1 low magnitude bits, then sign bit.
   Zero is encoded as just a 1. */
static void bw_sig(BWr *b, int32_t v) {
    uint32_t m = v < 0 ? (uint32_t)(-(int64_t)v) : (uint32_t)v;
    int n = 0;
    for (uint32_t t = m; t; t >>= 1) n++;
    for (int i = 0; i < n; i++) bw_put(b, 0);
    bw_put(b, 1);
    if (n) {
        if (n > 1) bw_bits(b, m & ((1u << (n - 1)) - 1), n - 1);
        bw_put(b, v < 0);
    }
}

typedef struct { const uint8_t *buf; size_t pos, maxbit; } BRd;

static int br_get(BRd *r) {
    if (r->pos >= r->maxbit) return 0;
    int bit = (r->buf[r->pos >> 3] >> (7 - (r->pos & 7))) & 1;
    r->pos++;
    return bit;
}
static uint32_t br_bits(BRd *r, int n) {
    uint32_t v = 0;
    while (n-- > 0) v = (v << 1) | (uint32_t)br_get(r);
    return v;
}
static int32_t br_sig(BRd *r) {
    int n = 0;
    while (br_get(r) == 0) { if (++n > 30) return 0; }
    if (!n) return 0;
    uint32_t m = 1u << (n - 1);
    if (n > 1) m |= br_bits(r, n - 1);
    return br_get(r) ? -(int32_t)m : (int32_t)m;
}

/* =========================================================================
 * Color transform (integer BT.601, all shifts and adds)
 * ========================================================================= */

static void rgb2ycc(int r, int g, int b, int *y, int *cb, int *cr) {
    *y  = ( 77*r + 150*g +  29*b + 128) >> 8;
    *cb = ((-43*r -  85*g + 128*b + 128) >> 8) + 128;
    *cr = ((128*r - 107*g -  21*b + 128) >> 8) + 128;
}
static void ycc2rgb(int y, int cb, int cr, int *r, int *g, int *b) {
    int u = cb - 128;   /* Cb - 128 */
    int v = cr - 128;   /* Cr - 128 */
    int R = y + ((359 * v + 128) >> 8);
    int G = y - (( 88 * u + 183 * v + 128) >> 8);
    int B = y + ((454 * u + 128) >> 8);
    *r = R < 0 ? 0 : R > 255 ? 255 : R;
    *g = G < 0 ? 0 : G > 255 ? 255 : G;
    *b = B < 0 ? 0 : B > 255 ? 255 : B;
}

/* =========================================================================
 * 4x4 integer DCT.
 * Forward:  Y = Cf * X * Cf^T
 * Cf = [1 1 1 1; 2 1 -1 -2; 1 -1 -1 1; 1 -2 2 -1]  (all integers)
 * Inverse:  X = Cf^T * (400 * D * Y * D) * Cf / 400
 * where D = diag(1/4, 1/10, 1/4, 1/10) scaled by 20 -> diag(5,2,5,2).
 * ========================================================================= */

static void dct4(const int x[16], int y[16]) {
    int t[16];
    for (int j = 0; j < 4; j++) {
        int x0 = x[j], x1 = x[4+j], x2 = x[8+j], x3 = x[12+j];
        t[j]      =  x0 + x1 + x2 + x3;
        t[4+j]    = 2*x0 + x1 - x2 - 2*x3;
        t[8+j]    =  x0 - x1 - x2 + x3;
        t[12+j]   =  x0 - 2*x1 + 2*x2 - x3;
    }
    for (int i = 0; i < 4; i++) {
        int a0 = t[i*4], a1 = t[i*4+1], a2 = t[i*4+2], a3 = t[i*4+3];
        y[i*4]   =  a0 + a1 + a2 + a3;
        y[i*4+1] = 2*a0 + a1 - a2 - 2*a3;
        y[i*4+2] =  a0 - a1 - a2 + a3;
        y[i*4+3] =  a0 - 2*a1 + 2*a2 - a3;
    }
}

static void idct4(const int y[16], int x[16]) {
    static const int D[4] = {5, 2, 5, 2};
    int d[16], t[16];
    for (int i = 0; i < 4; i++)
        for (int j = 0; j < 4; j++)
            d[i*4+j] = y[i*4+j] * D[i] * D[j];
    for (int j = 0; j < 4; j++) {
        int y0 = d[j], y1 = d[4+j], y2 = d[8+j], y3 = d[12+j];
        t[j]      = y0 + 2*y1 + y2 + y3;
        t[4+j]    = y0 + y1 - y2 - 2*y3;
        t[8+j]    = y0 - y1 - y2 + 2*y3;
        t[12+j]   = y0 - 2*y1 + y2 - y3;
    }
    for (int i = 0; i < 4; i++) {
        int a0 = t[i*4], a1 = t[i*4+1], a2 = t[i*4+2], a3 = t[i*4+3];
        x[i*4]   = a0 + 2*a1 + a2 + a3;
        x[i*4+1] = a0 + a1 - a2 - 2*a3;
        x[i*4+2] = a0 - a1 - a2 + 2*a3;
        x[i*4+3] = a0 - 2*a1 + a2 - a3;
    }
    /* x now holds 400 * original */
}

/* Diagonal scan order for 4x4 (low freq first) */
static const int zigzag[16] = {
    0, 1, 4, 8, 5, 2, 3, 6, 9, 12, 13, 10, 7, 11, 14, 15
};

/* =========================================================================
 * PPM I/O
 * ========================================================================= */

static uint8_t *read_ppm(const char *path, int *W, int *H) {
    FILE *f = fopen(path, "rb");
    if (!f) return NULL;
    char magic[3] = {0};
    if (fscanf(f, "%2s", magic) != 1 || strcmp(magic, "P6")) { fclose(f); return NULL; }
    int w, h, mx;
    if (fscanf(f, "%d %d %d", &w, &h, &mx) != 3 || mx != 255) { fclose(f); return NULL; }
    fgetc(f);
    uint8_t *d = (uint8_t *)malloc((size_t)w * h * 3);
    if (!d || fread(d, 1, (size_t)w * h * 3, f) != (size_t)w * h * 3) {
        free(d); fclose(f); return NULL;
    }
    fclose(f);
    *W = w; *H = h;
    return d;
}
static int write_ppm(const char *path, int w, int h, const uint8_t *d) {
    FILE *f = fopen(path, "wb");
    if (!f) return -1;
    fprintf(f, "P6\n%d %d\n255\n", w, h);
    fwrite(d, 1, (size_t)w * h * 3, f);
    fclose(f);
    return 0;
}

/* =========================================================================
 * Quantization
 * ========================================================================= */

static int qstep_from_quality(int q) {
    if (q < 1) q = 1;
    if (q > 100) q = 100;
    /* q=1 -> step 64, q=100 -> step 1 */
    int s = (int)((64.0 * (100.0 - q) / 99.0) + 0.5);
    return s < 1 ? 1 : s;
}

static int quant_coef(int v, int step) {
    int sign = v < 0 ? -1 : 1;
    int a = v < 0 ? -v : v;
    return sign * ((a + step / 2) / step);
}

/* =========================================================================
 * Block encode / decode
 * ========================================================================= */

static void encode_block(BWr *bw, const int block[16], int step, int *prev_dc) {
    int shifted[16];
    for (int i = 0; i < 16; i++) shifted[i] = block[i] - 128;
    int y[16];
    dct4(shifted, y);

    int qz[16];
    for (int i = 0; i < 16; i++) qz[i] = quant_coef(y[zigzag[i]], step);

    int any_ac = 0;
    for (int i = 1; i < 16; i++) if (qz[i]) { any_ac = 1; break; }

    if (!any_ac && qz[0] == *prev_dc) {
        bw_put(bw, 0);
        return;
    }
    bw_put(bw, 1);

    int dc_diff = qz[0] - *prev_dc;
    *prev_dc = qz[0];
    bw_sig(bw, dc_diff);

    int run = 0;
    for (int i = 1; i < 16; i++) {
        if (qz[i] == 0) { run++; continue; }
        int v = qz[i];
        uint32_t m = v < 0 ? (uint32_t)(-(int64_t)v) : (uint32_t)v;
        int nb = 0;
        for (uint32_t t = m; t; t >>= 1) nb++;
        bw_bits(bw, (uint32_t)run, 4);
        bw_bits(bw, (uint32_t)nb, 4);
        if (nb > 1) bw_bits(bw, m & ((1u << (nb - 1)) - 1), nb - 1);
        bw_put(bw, v < 0);
        run = 0;
    }
    if (run > 0) { bw_bits(bw, 0, 4); bw_bits(bw, 0, 4); }  /* EOB */
}

static void decode_block(BRd *br, int block[16], int step, int *prev_dc) {
    int qz[16] = {0};
    int x[16], y[16];

    if (!br_get(br)) {
        /* "no change" */
        qz[0] = *prev_dc;
        for (int i = 0; i < 16; i++) y[i] = qz[i] * step;
        idct4(y, x);
        for (int i = 0; i < 16; i++) {
            int v = (x[i] + 200) / 400 + 128;
            block[i] = v < 0 ? 0 : v > 255 ? 255 : v;
        }
        return;
    }

    int dc_diff = br_sig(br);
    qz[0] = *prev_dc + dc_diff;
    *prev_dc = qz[0];

    int idx = 1;
    while (idx < 16) {
        uint32_t run = br_bits(br, 4);
        uint32_t nb  = br_bits(br, 4);
        if (nb == 0 && run == 0) break;      /* EOB */
        idx += (int)run;
        if (idx >= 16) break;
        if (nb == 0) continue;               /* shouldn't happen, be defensive */
        uint32_t m = 1u << (nb - 1);
        if (nb > 1) m |= br_bits(br, nb - 1);
        int v = br_get(br) ? -(int32_t)m : (int32_t)m;
        qz[idx] = v;
        idx++;
    }

    for (int i = 0; i < 16; i++) y[zigzag[i]] = qz[i] * step;
    idct4(y, x);
    for (int i = 0; i < 16; i++) {
        int v = (x[i] + 200) / 400 + 128;
        block[i] = v < 0 ? 0 : v > 255 ? 255 : v;
    }
}

/* =========================================================================
 * Plane encode / decode
 * ========================================================================= */

static void encode_plane(BWr *bw, const int *pix, int w, int h, int step) {
    int bwc = (w + 3) / 4, bhc = (h + 3) / 4;
    int prev_dc = 0;
    for (int by = 0; by < bhc; by++)
        for (int bx = 0; bx < bwc; bx++) {
            int block[16];
            for (int py = 0; py < 4; py++)
                for (int px = 0; px < 4; px++) {
                    int x = bx*4 + px, y = by*4 + py;
                    if (x >= w) x = w - 1;
                    if (y >= h) y = h - 1;
                    block[py*4 + px] = pix[y * w + x];
                }
            encode_block(bw, block, step, &prev_dc);
        }
}

static void decode_plane(BRd *br, int *pix, int w, int h, int step) {
    int bwc = (w + 3) / 4, bhc = (h + 3) / 4;
    int prev_dc = 0;
    for (int by = 0; by < bhc; by++)
        for (int bx = 0; bx < bwc; bx++) {
            int block[16];
            decode_block(br, block, step, &prev_dc);
            for (int py = 0; py < 4; py++)
                for (int px = 0; px < 4; px++) {
                    int x = bx*4 + px, y = by*4 + py;
                    if (x < w && y < h) pix[y * w + x] = block[py*4 + px];
                }
        }
}

/* =========================================================================
 * Image encode / decode
 * ========================================================================= */

static int encode_image(const char *infile, const char *outfile, int quality) {
    int W, H;
    uint8_t *rgb = read_ppm(infile, &W, &H);
    if (!rgb) { fprintf(stderr, "cannot read %s\n", infile); return 1; }

    int cw = (W + 1) / 2, ch = (H + 1) / 2;
    int *Y  = (int *)malloc(sizeof(int) * W * H);
    int *Cb = (int *)malloc(sizeof(int) * cw * ch);
    int *Cr = (int *)malloc(sizeof(int) * cw * ch);

    for (int y = 0; y < H; y++)
        for (int x = 0; x < W; x++) {
            int i = (y * W + x) * 3;
            int yv, cb, cr;
            rgb2ycc(rgb[i], rgb[i+1], rgb[i+2], &yv, &cb, &cr);
            Y[y * W + x] = yv;
        }

    for (int cy = 0; cy < ch; cy++)
        for (int cx = 0; cx < cw; cx++) {
            int sr = 0, sg = 0, sb = 0, n = 0;
            for (int dy = 0; dy < 2; dy++) {
                int y = cy*2 + dy;
                if (y >= H) continue;
                for (int dx = 0; dx < 2; dx++) {
                    int x = cx*2 + dx;
                    if (x >= W) continue;
                    int i = (y * W + x) * 3;
                    sr += rgb[i]; sg += rgb[i+1]; sb += rgb[i+2];
                    n++;
                }
            }
            sr /= n; sg /= n; sb /= n;
            int yv, cb, cr;
            rgb2ycc(sr, sg, sb, &yv, &cb, &cr);
            Cb[cy * cw + cx] = cb;
            Cr[cy * cw + cx] = cr;
        }

    size_t cap = (size_t)W * H * 3 + 4096;
    BWr bw;
    bw_init(&bw, cap);

    bw_bits(&bw, 'F', 8); bw_bits(&bw, 'J', 8);
    bw_bits(&bw, 'P', 8); bw_bits(&bw, '1', 8);
    bw_bits(&bw, W, 16);
    bw_bits(&bw, H, 16);
    bw_bits(&bw, quality, 8);
    bw_bits(&bw, 0, 8);

    int qy = qstep_from_quality(quality);
    int qc = qstep_from_quality(quality - 8);
    if (qc < 1) qc = 1;

    encode_plane(&bw, Y,  W,  H,  qy);
    encode_plane(&bw, Cb, cw, ch, qc);
    encode_plane(&bw, Cr, cw, ch, qc);

    FILE *f = fopen(outfile, "wb");
    if (!f) { free(rgb); free(Y); free(Cb); free(Cr); free(bw.buf); return 1; }
    fwrite(bw.buf, 1, (bw.pos + 7) / 8, f);
    fclose(f);

    free(rgb); free(Y); free(Cb); free(Cr); free(bw.buf);
    return 0;
}

static int decode_image(const char *infile, const char *outfile) {
    FILE *f = fopen(infile, "rb");
    if (!f) { fprintf(stderr, "cannot read %s\n", infile); return 1; }
    fseek(f, 0, SEEK_END);
    long sz = ftell(f);
    fseek(f, 0, SEEK_SET);
    uint8_t *buf = (uint8_t *)malloc(sz);
    if (fread(buf, 1, sz, f) != (size_t)sz) { free(buf); fclose(f); return 1; }
    fclose(f);

    BRd br = { buf, 0, (size_t)sz * 8 };

    int m0 = br_bits(&br, 8), m1 = br_bits(&br, 8);
    int m2 = br_bits(&br, 8), m3 = br_bits(&br, 8);
    if (m0 != 'F' || m1 != 'J' || m2 != 'P' || m3 != '1') {
        fprintf(stderr, "not an FJP file\n"); free(buf); return 1;
    }
    int W = br_bits(&br, 16);
    int H = br_bits(&br, 16);
    int q = br_bits(&br, 8);
    br_bits(&br, 8);

    int cw = (W + 1) / 2, ch = (H + 1) / 2;
    int *Y  = (int *)malloc(sizeof(int) * W * H);
    int *Cb = (int *)malloc(sizeof(int) * cw * ch);
    int *Cr = (int *)malloc(sizeof(int) * cw * ch);

    int qy = qstep_from_quality(q);
    int qc = qstep_from_quality(q - 8);
    if (qc < 1) qc = 1;

    decode_plane(&br, Y,  W,  H,  qy);
    decode_plane(&br, Cb, cw, ch, qc);
    decode_plane(&br, Cr, cw, ch, qc);

    uint8_t *rgb = (uint8_t *)malloc((size_t)W * H * 3);
    for (int y = 0; y < H; y++)
        for (int x = 0; x < W; x++) {
            int yv = Y[y * W + x];
            int cb = Cb[(y>>1) * cw + (x>>1)];
            int cr = Cr[(y>>1) * cw + (x>>1)];
            int r, g, b;
            ycc2rgb(yv, cb, cr, &r, &g, &b);
            int i = (y * W + x) * 3;
            rgb[i] = r; rgb[i+1] = g; rgb[i+2] = b;
        }

    write_ppm(outfile, W, H, rgb);
    free(buf); free(Y); free(Cb); free(Cr); free(rgb);
    return 0;
}

int main(int argc, char **argv) {
    if (argc < 4) {
        fprintf(stderr, "usage:\n");
        fprintf(stderr, "  %s e input.ppm output.fjp [quality=75]\n", argv[0]);
        fprintf(stderr, "  %s d input.fjp output.ppm\n", argv[0]);
        return 2;
    }
    if (argv[1][0] == 'e') {
        int q = (argc >= 5) ? atoi(argv[4]) : 75;
        return encode_image(argv[2], argv[3], q);
    }
    if (argv[1][0] == 'd') return decode_image(argv[2], argv[3]);
    fprintf(stderr, "unknown mode '%s'\n", argv[1]);
    return 2;
}
