/* SPDX-License-Identifier: MIT. Copyright (c) 2026 Joerg Burbach.
 * Memory-only LeXA-26; reference: Formats/module_format_lexa.pbi.
 */
#include "lexa.h"
#include <stdlib.h>
#include <string.h>
#define LIMIT ((size_t)67108864)
static unsigned u16(const uint8_t *s) { return s[0] | (unsigned)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 void put16(uint8_t *s, size_t v) { s[0] = (uint8_t)v; s[1] = (uint8_t)(v >> 8); }
static void put32(uint8_t *s, size_t v) { size_t i; for (i = 0; i < 4; i++) s[i] = (uint8_t)(v >> (i * 8)); }
typedef struct { const uint8_t *s; size_t size, pos; unsigned tag, left; int error; } reader;
static unsigned getbit(reader *r) {
    unsigned value;
    if (!r->left) { if (r->pos > r->size || r->size - r->pos < 2) { r->error = 1; return 0; } r->tag = u16(r->s + r->pos); r->pos += 2; r->left = 16; }
    r->left--; value = r->tag >> 15 & 1; r->tag = r->tag << 1 & 65535; return value;
}
static uint32_t getgamma(reader *r) {
    uint32_t value = 1;
    do { if (value >= UINT32_C(0x40000000)) { r->error = 1; return 0; } value = value * 2 + getbit(r); } while (getbit(r) && !r->error);
    return value;
}
uint8_t *lexa_brief_lz_decode(const uint8_t *s, size_t n, size_t expected, size_t *consumed) {
    reader r; uint8_t *out; size_t d = 0, count, offset, i; uint32_t high;
    if (consumed) *consumed = 0;
    if (!s || !n || expected > LIMIT) return NULL;
    out = (uint8_t *)calloc(expected + 1, 1); if (!out) return NULL;
    r.s = s; r.size = n; r.pos = 0; r.tag = 0; r.left = 1; r.error = 0;
    while (d < expected) {
        if (getbit(&r)) {
            count = (size_t)getgamma(&r) + 2; high = getgamma(&r);
            if (r.error || high < 2 || high - 2 >= UINT32_C(0xffffff) || r.pos >= n) goto error;
            offset = (size_t)(high - 2) * 256 + s[r.pos++] + 1;
            if (offset > d || count > expected - d) goto error;
            for (i = 0; i < count; i++) { out[d] = out[d - offset]; d++; }
        } else { if (r.error || r.pos >= n) goto error; out[d++] = s[r.pos++]; }
    }
    if (consumed) *consumed = r.pos;
    return out;
error:
    free(out); return NULL;
}
typedef struct { uint8_t *out; size_t pos, tag_pos; unsigned tag, left; } writer;
static void putbit(writer *w, unsigned bit) {
    if (!w->left) { put16(w->out + w->tag_pos, w->tag); w->tag_pos = w->pos; w->pos += 2; w->left = 15; }
    else w->left--;
    w->tag = (w->tag << 1 | bit) & 65535;
}
static void putgamma(writer *w, size_t value) {
    size_t mask = 1;
    while (mask * 2 <= value) mask *= 2;
    mask /= 2; putbit(w, (value & mask) != 0); mask /= 2;
    while (mask) { putbit(w, 1); putbit(w, (value & mask) != 0); mask /= 2; }
    putbit(w, 0);
}
static unsigned hash4(const uint8_t *s) { return (u32(s) * UINT32_C(2654435761)) >> 15; }
uint8_t *lexa_brief_lz_encode(const uint8_t *s, size_t n, size_t *size) {
    uint8_t *out; int32_t *lookup; size_t cur = 1, hash_pos = 0, i, length, offset; int32_t previous; writer w;
    if (!size) return NULL;
    *size = 0; if (!s || !n || n > LIMIT) return NULL;
    out = (uint8_t *)calloc(n + n / 8 + 64, 1); if (!out) return NULL;
    out[0] = s[0]; if (n == 1) { *size = 1; return out; }
    lookup = (int32_t *)malloc(131072 * sizeof(*lookup)); if (!lookup) { free(out); return NULL; }
    for (i = 0; i < 131072; i++) lookup[i] = -1;
    w.out = out; w.pos = 3; w.tag_pos = 1; w.tag = 0; w.left = 16;
    while (n >= 4 && cur <= n - 4) {
        while (hash_pos < cur) { lookup[hash4(s + hash_pos)] = (int32_t)hash_pos; hash_pos++; }
        previous = lookup[hash4(s + cur)]; length = 0;
        if (previous >= 0) while (length < n - cur && s[(size_t)previous + length] == s[cur + length]) length++;
        if (length > 4 || (length == 4 && cur - (size_t)previous - 1 < 0x7e00)) {
            offset = cur - (size_t)previous - 1; putbit(&w, 1); putgamma(&w, length - 2); putgamma(&w, offset / 256 + 2); out[w.pos++] = (uint8_t)offset; cur += length;
        } else { putbit(&w, 0); out[w.pos++] = s[cur++]; }
    }
    while (cur < n) { putbit(&w, 0); out[w.pos++] = s[cur++]; }
    putbit(&w, 1); w.tag = w.tag << w.left & 65535; put16(out + w.tag_pos, w.tag); *size = w.pos; free(lookup); return out;
}
static uint8_t *lzss_decode(const uint8_t *s, size_t n, size_t expected) {
    size_t p = 0, d = 0, i, count, offset; unsigned flags, bit, a, b; uint8_t *out;
    if (!s || !n || !expected || expected > LIMIT) return NULL;
    out = (uint8_t *)malloc(expected); if (!out) return NULL;
    while (d < expected) {
        if (p >= n) goto error;
        flags = s[p++];
        for (bit = 0; bit < 8 && d < expected; bit++) {
            if (flags & (1u << bit)) {
                if (n - p < 2) goto error;
                a = s[p++]; b = s[p++]; count = (a >> 4) + 4; offset = ((a & 15) << 8 | b) + 1;
                if (offset > d || count > expected - d) goto error;
                for (i = 0; i < count; i++) { out[d] = out[d - offset]; d++; }
            } else { if (p >= n) goto error; out[d++] = s[p++]; }
        }
    }
    return out;
error:
    free(out); return NULL;
}
static unsigned hash3(const uint8_t *s) { return ((s[0] * 251u + s[1]) * 251u + s[2]) & 65535; }
static uint8_t *lzss_encode(const uint8_t *s, size_t n, size_t *size) {
    int32_t *head = NULL, *previous = NULL, candidate; uint8_t *out = NULL; size_t p = 0, d = 0, i, j, best, offset, count, flags; unsigned bit, visits, key;
    if (!n || n > LIMIT) return NULL;
    head = (int32_t *)malloc(65536 * sizeof(*head)); previous = (int32_t *)malloc(n * sizeof(*previous)); out = (uint8_t *)calloc(n + (n + 7) / 8, 1);
    if (!head || !previous || !out) goto error;
    for (i = 0; i < 65536; i++) head[i] = -1;
    for (i = 0; i < n; i++) previous[i] = -1;
    while (p < n) {
        flags = d++;
        for (bit = 0; bit < 8 && p < n; bit++) {
            best = offset = 0; visits = 64; candidate = p + 2 < n ? head[hash3(s + p)] : -1;
            while (candidate >= 0 && p - (size_t)candidate <= 4096 && visits--) {
                count = 0; while (count < 19 && p + count < n && s[(size_t)candidate + count] == s[p + count]) count++;
                if (count > best) { best = count; offset = p - (size_t)candidate; } candidate = previous[candidate];
            }
            if (best >= 4) { out[flags] |= (uint8_t)(1u << bit); out[d++] = (uint8_t)((best - 4) << 4 | (offset - 1) >> 8); out[d++] = (uint8_t)(offset - 1); }
            else { best = 1; out[d++] = s[p]; }
            for (j = 0; j < best; j++) { if (p + 2 < n) { key = hash3(s + p); previous[p] = head[key]; head[key] = (int32_t)p; } p++; }
        }
    }
    free(head); free(previous); *size = d; return out;
error:
    free(head); free(previous); free(out); return NULL;
}
int lexa_inspect(const uint8_t *s, size_t n, lexa_info *info) {
    size_t count, pos, i, length, raw_size, row, col, glyph; unsigned method, width, height;
    if (!s || !info || n < 8 || n > LIMIT || memcmp(s, "LEXA", 4) || s[4] != 26) return 1;
    memset(info, 0, sizeof(*info)); info->kind = s[5];
    if (s[5] == 1) {
        count = u16(s + 6); if (!count || count > (n - 8) / 6) return 1;
        pos = 8 + count * 6;
        for (i = 0; i < count; i++) { length = u16(s + 8 + i * 6); if (length > n - pos) return 1; pos += length; }
        for (i = 0; i < count; i++) {
            length = u32(s + 10 + i * 6); if (length < 5 || length > n - pos || s[pos] > 1) return 1;
            raw_size = u32(s + pos + 1); if (raw_size > LIMIT || (!s[pos] && raw_size != length - 5)) return 1; pos += length;
        }
        if (pos != n) return 1;
        info->text_count = (uint16_t)count; return 0;
    }
    if (s[5] != 2 || n <= 20 || s[6] || s[7] || s[9] || s[14] || s[15] || s[8] > 2) return 1;
    method = s[8]; width = u16(s + 10); height = u16(s + 12); raw_size = u32(s + 16);
    if (!raw_size || raw_size > LIMIT || (!method && n - 20 != raw_size)) return 1;
    if (method == 2) {
        row = (width + 7) / 8; col = (height + 7) / 8; glyph = row * height;
        if (!width || !height || !glyph || raw_size % glyph || width * col != glyph) return 1;
    } else if (width || height) return 1;
    info->method = (uint8_t)method; info->width = (uint16_t)width; info->height = (uint16_t)height; info->raw_size = (uint32_t)raw_size; return 0;
}
static uint8_t *transpose(const uint8_t *s, size_t n, unsigned w, unsigned h, int inverse) {
    size_t row = (w + 7) / 8, col = (h + 7) / 8, glyph = row * h, base, x, y, r, c; uint8_t *out;
    if (!w || !h || !glyph || n % glyph || w * col != glyph) return NULL;
    out = (uint8_t *)calloc(n, 1); if (!out) return NULL;
    for (base = 0; base < n; base += glyph) for (y = 0; y < h; y++) for (x = 0; x < w; x++) {
        r = base + y * row + x / 8; c = base + x * col + y / 8;
        if (inverse ? (s[c] >> (7 - y % 8) & 1) : (s[r] >> (7 - x % 8) & 1)) out[inverse ? r : c] |= (uint8_t)(1u << (7 - (inverse ? x % 8 : y % 8)));
    }
    return out;
}
void lexa_free(lexa_document *doc) {
    size_t i; if (!doc) return;
    if (doc->entries) for (i = 0; i < doc->count; i++) { free((void *)doc->entries[i].id); free((void *)doc->entries[i].bytes); }
    free(doc->entries); free(doc->font); memset(doc, 0, sizeof(*doc));
}
int lexa_decode(const uint8_t *s, size_t n, lexa_document *doc) {
    lexa_info info; size_t ids, pos, i, length, raw; uint8_t *font;
    if (!doc) return 1;
    memset(doc, 0, sizeof(*doc)); if (lexa_inspect(s, n, &info)) return 1; doc->kind = info.kind;
    if (info.kind == 1) {
        doc->count = info.text_count; doc->entries = (lexa_text *)calloc(doc->count, sizeof(*doc->entries)); if (!doc->entries) goto error;
        ids = pos = 8 + doc->count * 6; for (i = 0; i < doc->count; i++) pos += u16(s + 8 + i * 6);
        for (i = 0; i < doc->count; i++) {
            doc->entries[i].id_size = u16(s + 8 + i * 6); length = u32(s + 10 + i * 6); raw = u32(s + pos + 1);
            doc->entries[i].id = (uint8_t *)calloc(doc->entries[i].id_size + 1, 1); if (!doc->entries[i].id) goto error;
            memcpy((void *)doc->entries[i].id, s + ids, doc->entries[i].id_size); ids += doc->entries[i].id_size;
            if (s[pos]) doc->entries[i].bytes = lexa_brief_lz_decode(s + pos + 5, length - 5, raw, NULL);
            else { doc->entries[i].bytes = (uint8_t *)calloc(raw + 1, 1); if (doc->entries[i].bytes) memcpy((void *)doc->entries[i].bytes, s + pos + 5, raw); }
            if (!doc->entries[i].bytes) goto error;
            doc->entries[i].size = raw; pos += length;
        }
    } else {
        font = info.method ? lzss_decode(s + 20, n - 20, info.raw_size) : (uint8_t *)malloc(info.raw_size); if (!font) goto error;
        if (!info.method) memcpy(font, s + 20, info.raw_size);
        if (info.method == 2) { doc->font = transpose(font, info.raw_size, info.width, info.height, 1); free(font); } else doc->font = font;
        if (!doc->font) goto error;
        doc->font_size = info.raw_size; doc->width = info.width; doc->height = info.height;
    }
    return 0;
error:
    lexa_free(doc); return 1;
}
int lexa_encode_text(const lexa_text *entries, size_t count, uint8_t **data, size_t *size) {
    uint8_t **payload = NULL, *out = NULL; size_t *lengths = NULL, total = 8, i, packed_size, p; unsigned compressed;
    if (!data || !size) return 1;
    *data = NULL; *size = 0; if (!entries || !count || count > 65535) return 1;
    total += count * 6; payload = (uint8_t **)calloc(count, sizeof(*payload)); lengths = (size_t *)calloc(count, sizeof(*lengths)); if (!payload || !lengths) goto error;
    for (i = 0; i < count; i++) {
        if (entries[i].id_size > 65535 || entries[i].size > LIMIT || (!entries[i].id && entries[i].id_size) || (!entries[i].bytes && entries[i].size)) goto error;
        payload[i] = lexa_brief_lz_encode(entries[i].bytes, entries[i].size, &packed_size);
        compressed = payload[i] && packed_size + 5 < entries[i].size;
        if (!compressed) { free(payload[i]); payload[i] = NULL; packed_size = entries[i].size; }
        lengths[i] = packed_size; if (packed_size + entries[i].id_size + 5 > LIMIT - total) goto error; total += packed_size + entries[i].id_size + 5;
    }
    out = (uint8_t *)calloc(total, 1); if (!out) goto error;
    memcpy(out, "LEXA", 4); out[4] = 26; out[5] = 1; put16(out + 6, count); p = 8 + count * 6;
    for (i = 0; i < count; i++) { put16(out + 8 + i * 6, entries[i].id_size); put32(out + 10 + i * 6, lengths[i] + 5); if (entries[i].id_size) memcpy(out + p, entries[i].id, entries[i].id_size); p += entries[i].id_size; }
    for (i = 0; i < count; i++) { out[p] = payload[i] ? 1 : 0; put32(out + p + 1, entries[i].size); if (lengths[i]) memcpy(out + p + 5, payload[i] ? payload[i] : entries[i].bytes, lengths[i]); p += lengths[i] + 5; }
    for (i = 0; i < count; i++) free(payload[i]);
    free(payload); free(lengths); *data = out; *size = total; return 0;
error:
    if (payload) for (i = 0; i < count; i++) free(payload[i]);
    free(payload); free(lengths); free(out); return 1;
}
int lexa_encode_font(const uint8_t *raw, size_t n, uint16_t w, uint16_t h, uint8_t **data, size_t *size) {
    uint8_t *packed = NULL, *columns = NULL, *column_pack = NULL, *restored = NULL, *out = NULL; const uint8_t *payload = raw; size_t length = n, packed_size, column_size; unsigned method = 0;
    if (!data || !size) return 1;
    *data = NULL; *size = 0; if (!raw || !n || n > LIMIT - 20) return 1;
    packed = lzss_encode(raw, n, &packed_size); if (!packed) goto error;
    if (packed_size < n) { payload = packed; length = packed_size; method = 1; }
    if (w || h) {
        columns = transpose(raw, n, w, h, 0); if (!columns) goto error;
        restored = transpose(columns, n, w, h, 1); if (!restored) goto error;
        column_pack = lzss_encode(columns, n, &column_size); if (!column_pack) goto error;
        if (column_size < length && memcmp(restored, raw, n) == 0) { payload = column_pack; length = column_size; method = 2; }
    }
    out = (uint8_t *)calloc(20 + length, 1); if (!out) goto error;
    memcpy(out, "LEXA", 4); out[4] = 26; out[5] = 2; out[8] = (uint8_t)method; put16(out + 10, method == 2 ? w : 0); put16(out + 12, method == 2 ? h : 0); put32(out + 16, n); memcpy(out + 20, payload, length);
    free(packed); free(columns); free(column_pack); free(restored); *data = out; *size = 20 + length; return 0;
error:
    free(packed); free(columns); free(column_pack); free(restored); free(out); return 1;
}
