/* SPDX-License-Identifier: MIT
 * Copyright (c) 2026 Joerg Burbach.
 * RFXL-26 memory codec, based on Formats/module_format_rfxl.pbi.
 */
#include "rfxl.h"
#include <stdlib.h>
#include <string.h>

static uint16_t u16(const uint8_t *s) { return (uint16_t)(s[0] | (uint16_t)s[1] << 8); }
static uint32_t u32(const uint8_t *s) { return s[0] | (uint32_t)s[1] << 8 | (uint32_t)s[2] << 16 | (uint32_t)s[3] << 24; }
static int dimensions(uint32_t w, uint32_t h) { return w && h && w <= 8192 && h <= 8192 && w <= 16777216 / h; }
int rfxl_inspect(const uint8_t *s, size_t n, rfxl_info *info) {
    if (!s || !info || n < 12 || n > 67112960 || memcmp(s, "RFXL", 4) || s[4] != 26) return 1;
    if (!dimensions(u16(s + 5), u16(s + 7)) || (s[9] != 3 && s[9] != 4) || (s[10] != 0 && s[10] != 4)) return 1;
    info->width = u16(s + 5); info->height = u16(s + 7); info->channels = s[9]; info->mode = s[10];
    return 0;
}
typedef struct { const uint8_t *data; size_t size, bit; int error; } bit_reader;
static unsigned bits(bit_reader *r, unsigned count) {
    unsigned value = 0, i;
    for (i = 0; i < count; i++) {
        if (r->bit / 8 >= r->size) { r->error = 1; return 0; }
        value |= (unsigned)(r->data[r->bit / 8] >> (r->bit % 8) & 1) << i;
        r->bit++;
    }
    return value;
}
uint8_t *rfxl_unpack(const uint8_t *s, size_t n, unsigned method, size_t expected) {
    size_t p = 0, d = 0, i, count;
    unsigned token, kind, flags, a, b, nib;
    uint8_t *out;
    if (!s || !n || !expected || expected > 67108864 || method > 4) return NULL;
    out = (uint8_t *)calloc(expected, 1);
    if (!out) return NULL;
    if (method == 0) { if (n != expected) goto error; memcpy(out, s, n); return out; }
    if (method == 3) {
        bit_reader r = {s, n, 8, 0};
        size_t block = (size_t)16 << (s[0] & 3), end;
        while (d < expected) {
            unsigned k = bits(&r, 4), q, z;
            end = expected - d < block ? expected : d + block;
            while (d < end) {
                q = 0;
                while (bits(&r, 1)) { if (++q > 4) goto error; }
                z = q == 4 ? bits(&r, 8) : (q << k) | bits(&r, k);
                if (r.error) goto error;
                out[d++] = (uint8_t)(z & 1 ? -(int)((z >> 1) + 1) : (int)(z >> 1));
            }
        }
        return out;
    }
    while (d < expected) {
        if (p >= n) goto error;
        token = s[p++];
        if (method == 4) {
            flags = token;
            for (i = 0; i < 8 && d < expected; i++) {
                if (flags & (1u << i)) {
                    size_t offset, j;
                    if (n - p < 2) goto error;
                    a = s[p++]; b = s[p++]; nib = a >> 4; count = nib + 4;
                    if (nib == 15) {
                        unsigned extra;
                        do { if (p >= n) goto error; extra = s[p++]; count += extra; if (count > expected - d) goto error; } while (extra == 255);
                    }
                    offset = ((a & 15) << 8 | b) + 1;
                    if (offset > d || count > expected - d) goto error;
                    for (j = 0; j < count; j++) { out[d] = out[d - offset]; d++; }
                } else { if (p >= n) goto error; out[d++] = s[p++]; }
            }
            continue;
        }
        kind = method == 1 ? token & 128 : token & 192;
        count = method == 1 ? token & 127 : (token & 63) + 1;
        if (!count || count > expected - d) goto error;
        if (!kind) { if (count > n - p) goto error; memcpy(out + d, s + p, count); p += count; d += count; }
        else if (method == 2 && kind == 64) d += count;
        else if (method == 1 || kind == 128) { if (p >= n) goto error; memset(out + d, s[p++], count); d += count; }
        else {
            for (i = 0; i < count; i += 2) {
                if (p >= n) goto error;
                a = s[p] & 15; b = s[p++] >> 4;
                out[d++] = (uint8_t)(a >= 8 ? (int)a - 16 : (int)a);
                if (i + 1 < count) out[d++] = (uint8_t)(b >= 8 ? (int)b - 16 : (int)b);
            }
        }
    }
    return out;
error:
    free(out); return NULL;
}
static void predict(const uint8_t *s, uint8_t *out, size_t n, size_t width, unsigned bpp, unsigned mode, int inverse) {
    size_t i, row = width * bpp;
    const uint8_t *history = inverse ? out : s;
    for (i = 0; i < n; i++) {
        size_t x = i % row;
        int l = x >= bpp ? history[i - bpp] : 0, u = i >= row ? history[i - row] : 0;
        int ul = x >= bpp && i >= row ? history[i - row - bpp] : 0, p = 0;
        if (mode == 1) p = l;
        else if (mode == 2) p = u;
        else if (mode == 3) p = l + u - ul;
        else if (mode == 4) p = ul >= l && ul >= u ? (l < u ? l : u) : ul <= l && ul <= u ? (l > u ? l : u) : l + u - ul;
        out[i] = (uint8_t)(inverse ? s[i] + p : s[i] - p);
    }
}
int rfxl_decode(const uint8_t *data, size_t size, rfxl_image *image) {
    rfxl_info info;
    const uint8_t *s, *palette;
    uint8_t *rgba = NULL, *stream = NULL, *predicted = NULL, *pal = NULL, *mask = NULL;
    size_t n, pixels, length, i, offset, count, j, palette_length, changed, mask_length, index_length;
    unsigned format, mode, method, bpp, palbits, delta;
    if (!image) return 1;
    memset(image, 0, sizeof(*image));
    if (rfxl_inspect(data, size, &info)) return 1;
    s = data + 11; n = size - 11; pixels = (size_t)info.width * info.height;
    rgba = (uint8_t *)malloc(pixels * 4);
    if (!rgba) goto error;
    if (info.mode == 4) {
        size_t foreground = 0;
        if (n < 23) goto error;
        method = s[4]; changed = u32(s + 9); count = u16(s + 13); mask_length = u32(s + 15); index_length = u32(s + 19); offset = 23 + count * 4;
        if (u32(s + 5) != pixels || !changed || changed > pixels || !count || count > 256 || !mask_length || !index_length || offset > n || mask_length > n - offset || index_length != n - offset - mask_length || (method != 2 && method != 4)) goto error;
        palette = s + 23;
        mask = rfxl_unpack(s + offset, mask_length, method, (pixels + 7) / 8);
        stream = rfxl_unpack(s + offset + mask_length, index_length, method, changed);
        if (!mask || !stream) goto error;
        for (i = 0; i < pixels; i++) {
            if (mask[i / 8] >> (i % 8) & 1) {
                if (foreground >= changed || stream[foreground] >= count) goto error;
                memcpy(rgba + i * 4, palette + (size_t)stream[foreground++] * 4, 4);
            } else memcpy(rgba + i * 4, s, 4);
        }
        if (foreground != changed) goto error;
    } else {
        if (n < 7) goto error;
        format = s[0]; mode = s[1]; method = s[2]; length = u32(s + 3);
        if (format > 10 || mode > 4 || method > 4) goto error;
        if (format == 1 || format == 4 || format == 5 || format == 6) {
            uint8_t mtf[256];
            if (n < 10) goto error;
            palbits = format == 1 ? 2 : format == 4 ? 8 : format == 5 ? 4 : 1;
            count = u16(s + 7); delta = s[9];
            if (length != pixels || !count || count > (1u << palbits) || delta > 1 || (palbits != 8 && mode > 1)) goto error;
            if (delta) {
                if (n < 12) goto error;
                palette_length = u16(s + 10); offset = 12 + palette_length;
                if (offset > n) goto error;
                predicted = rfxl_unpack(s + 12, palette_length, 3, count * 4);
                pal = (uint8_t *)malloc(count * 4);
                if (!predicted || !pal) goto error;
                predict(predicted, pal, count * 4, count, 4, 1, 1); free(predicted); predicted = NULL; palette = pal;
            } else { offset = 10 + count * 4; if (offset > n) goto error; palette = s + 10; }
            length = (pixels * palbits + 7) / 8;
            stream = rfxl_unpack(s + offset, n - offset, method, length);
            if (!stream) goto error;
            if (palbits == 8 && mode) {
                predicted = stream; stream = (uint8_t *)malloc(length); if (!stream) goto error;
                predict(predicted, stream, length, info.width, 1, mode, 1);
                free(predicted); predicted = NULL;
            }
            for (i = 0; i < (1u << palbits); i++) mtf[i] = (uint8_t)i;
            for (i = 0; i < pixels; i++) {
                size_t bit = i * palbits;
                unsigned index = stream[bit / 8] >> (bit % 8) & ((1u << palbits) - 1);
                if (palbits != 8 && mode == 1) { unsigned value = mtf[index]; for (j = index; j > 0; j--) mtf[j] = mtf[j - 1]; mtf[0] = (uint8_t)value; index = value; }
                if (index >= count) goto error;
                memcpy(rgba + i * 4, palette + index * 4, 4);
            }
        } else {
            bpp = format == 7 ? 1 : format == 2 || format == 3 ? 2 : format == 9 || format == 10 ? 3 : 4;
            if (length != pixels * bpp) goto error;
            predicted = rfxl_unpack(s + 7, n - 7, method, length);
            stream = (uint8_t *)malloc(length);
            if (!predicted || !stream) goto error;
            predict(predicted, stream, length, info.width, format == 8 || format == 10 ? 1 : bpp, mode, 1);
            for (i = 0; i < pixels; i++) {
                size_t d = i * 4; unsigned v;
                j = i * bpp; rgba[d + 3] = 255;
                if (format == 7) { v = stream[i]; rgba[d] = (uint8_t)((v >> 5) * 255 / 7); rgba[d + 1] = (uint8_t)((v >> 2 & 7) * 255 / 7); rgba[d + 2] = (uint8_t)((v & 3) * 85); }
                else if (format == 2) { rgba[d] = stream[j] & 240; rgba[d + 1] = (uint8_t)(stream[j] << 4 & 240); rgba[d + 2] = stream[j + 1] & 240; v = stream[j + 1] << 4 & 240; rgba[d + 3] = (uint8_t)(v == 240 ? 255 : v); }
                else if (format == 3) { v = u16(stream + j); rgba[d] = (uint8_t)(v >> 8 & 248); rgba[d + 1] = (uint8_t)(v >> 3 & 252); rgba[d + 2] = (uint8_t)(v << 3 & 248); }
                else if (format == 8 || format == 10) { unsigned c; for (c = 0; c < bpp; c++) rgba[d + c] = stream[(size_t)c * pixels + i]; }
                else memcpy(rgba + d, stream + j, bpp);
            }
        }
    }
    free(stream); free(predicted); free(pal); free(mask);
    image->width = info.width; image->height = info.height; image->rgba = rgba;
    return 0;
error:
    free(rgba); free(stream); free(predicted); free(pal); free(mask); return 1;
}
static size_t pack(const uint8_t *s, size_t n, uint8_t *out) {
    size_t i = 0, d = 0;
    while (i < n) {
        size_t run = 1, start, j;
        while (run < 64 && i + run < n && s[i + run] == s[i]) run++;
        if (run >= 3 || !s[i]) { out[d++] = (uint8_t)((s[i] ? 128 : 64) | (run - 1)); if (s[i]) out[d++] = s[i]; i += run; }
        else {
            start = i++;
            while (i - start < 64 && i < n && s[i] && !(i + 2 < n && s[i] == s[i + 1] && s[i] == s[i + 2])) i++;
            out[d++] = (uint8_t)(i - start - 1); for (j = start; j < i; j++) out[d++] = s[j];
        }
    }
    return d;
}
int rfxl_encode(const rfxl_image *image, enum rfxl_quality quality, uint8_t **data, size_t *size) {
    uint8_t *stream = NULL, *prediction = NULL, *candidate = NULL, *out = NULL;
    size_t pixels, length, best_length, i, j, packed_length;
    unsigned alpha = 0, shift, bpp, mode, best_mode = 0, method = 0;
    if (!data || !size) return 1;
    *data = NULL; *size = 0;
    if (!image || !image->rgba || !dimensions(image->width, image->height) || quality < RFXL_CRAPPY || quality > RFXL_LOSSLESS) return 1;
    pixels = (size_t)image->width * image->height;
    for (i = 0; i < pixels; i++) if (image->rgba[i * 4 + 3] != 255) { alpha = 1; break; }
    bpp = alpha ? 4 : 3; length = pixels * bpp;
    shift = image->width < 64 || image->height < 64 || quality == RFXL_LOSSLESS ? 0 : 6 - (unsigned)quality;
    stream = (uint8_t *)malloc(length); prediction = (uint8_t *)malloc(length); candidate = (uint8_t *)malloc(length * 2);
    out = (uint8_t *)malloc(18 + length);
    if (!stream || !prediction || !candidate || !out) goto error;
    for (i = 0; i < pixels; i++) for (j = 0; j < bpp; j++) stream[i * bpp + j] = j < 3 ? (uint8_t)(image->rgba[i * 4 + j] >> shift << shift) : image->rgba[i * 4 + j];
    best_length = length; memcpy(out + 18, stream, length);
    for (mode = 0; mode <= 4; mode++) {
        predict(stream, prediction, length, image->width, bpp, mode, 0);
        packed_length = pack(prediction, length, candidate);
        if (packed_length < best_length) { best_length = packed_length; best_mode = mode; method = 2; memcpy(out + 18, candidate, packed_length); }
    }
    memcpy(out, "RFXL", 4); out[4] = 26; out[5] = (uint8_t)image->width; out[6] = (uint8_t)(image->width >> 8); out[7] = (uint8_t)image->height; out[8] = (uint8_t)(image->height >> 8);
    out[9] = (uint8_t)bpp; out[10] = 0; out[11] = alpha ? 0 : 9; out[12] = (uint8_t)best_mode; out[13] = (uint8_t)method;
    for (i = 0; i < 4; i++) out[14 + i] = (uint8_t)(length >> (i * 8));
    free(stream); free(prediction); free(candidate); *data = out; *size = 18 + best_length; return 0;
error:
    free(stream); free(prediction); free(candidate); free(out); return 1;
}
