/* SPDX-License-Identifier: MIT. Copyright (c) 2026 Joerg Burbach.
 * Memory-only FormA-26 geometry + SCN attribute streams.
 * Link shared/packbits.c and rfxl.c. Reference: module_format_forma.pbi.
 */
#include "forma.h"
#include "../../shared/packbits.h"
#include "../../rfxl/codecs/rfxl.h"
#include <stdlib.h>
#include <string.h>
#include <math.h>
#define LIMIT ((size_t)67108864)
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 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 n, pos; int error; } reader;
static uint32_t next32(reader *r) { uint32_t value; if (r->pos > r->n || r->n - r->pos < 4) { r->error = 1; return 0; } value = u32(r->s + r->pos); r->pos += 4; return value; }
static void skip(reader *r, size_t n) { if (r->pos > r->n || n > r->n - r->pos) r->error = 1; else r->pos += n; }
static void string(reader *r) { size_t n = next32(r); skip(r, n); }
static int64_t varint(reader *r) {
    uint64_t z = 0; unsigned i, b;
    for (i = 0; i < 5; i++) { if (r->pos >= r->n) break; b = r->s[r->pos++]; z |= (uint64_t)(b & 127) << (i * 7); if (!(b & 128)) return z & 1 ? -(int64_t)((z >> 1) + 1) : (int64_t)(z >> 1); }
    r->error = 1; return 0;
}
static void floats(reader *r, size_t count) { size_t i; for (i = 0; i < count && !r->error; i++) { uint32_t raw = next32(r); float f; memcpy(&f, &raw, 4); if (!isfinite(f)) r->error = 1; } }
static int scene_valid(const uint8_t *s, size_t n, size_t vertices, size_t faces, const uint32_t *indices) {
    reader r; uint32_t counts[5], v, normals, uvs, f, flags, bytes; size_t i, j, vertex_base = 0, face_base = 0; int64_t material, cn[3], cu[3], value; int second;
    if (!s || n < 4 || n > LIMIT || (memcmp(s, "SCN1", 4) && memcmp(s, "SCN2", 4))) return 0;
    second = s[3] == '2'; r.s = s; r.n = n; r.pos = 4; r.error = 0;
    string(&r); skip(&r, 8); floats(&r, 8); floats(&r, 6); skip(&r, 4); floats(&r, 1); skip(&r, 8); floats(&r, 3); floats(&r, 1); skip(&r, second ? 28 : 24);
    for (i = 0; i < 5; i++) { counts[i] = next32(&r); if (counts[i] > n) return 0; }
    for (i = 0; i < counts[0] && !r.error; i++) { string(&r); skip(&r, 4); floats(&r, 8); skip(&r, 4); string(&r); skip(&r, 4); floats(&r, 2); skip(&r, 8); floats(&r, 1); skip(&r, 4); }
    for (i = 0; i < counts[1] && !r.error; i++) {
        rfxl_image image;
        string(&r); skip(&r, 24); bytes = next32(&r); if (r.error || bytes > n - r.pos) return 0;
        if (bytes) { if (rfxl_decode(s + r.pos, bytes, &image)) return 0; free(image.rgba); } skip(&r, bytes);
    }
    for (i = 0; i < counts[2] && !r.error; i++) { floats(&r, 6); skip(&r, 4); floats(&r, 1); skip(&r, 8); floats(&r, 3); }
    for (i = 0; i < counts[3] && !r.error; i++) {
        string(&r); v = next32(&r); normals = next32(&r); uvs = next32(&r); f = next32(&r);
        value = varint(&r); if (value < 0 || value > 31) return 0; flags = (uint32_t)value; material = varint(&r);
        if ((flags & 1 && flags & 8) || (flags & 2 && flags & 16) || v > vertices - vertex_base || f > faces - face_base || normals > n / 12 || uvs > n / 8) return 0;
        for (j = 0; j < 3; j++) { cn[j] = -1; cu[j] = -1; }
        if (flags & 8) for (j = 0; j < 3; j++) cn[j] = varint(&r);
        if (flags & 16) for (j = 0; j < 3; j++) cu[j] = varint(&r);
        floats(&r, (size_t)normals * 3); floats(&r, (size_t)uvs * 2);
        for (j = 0; j < f && !r.error; j++) {
            unsigned axis;
            for (axis = 0; axis < 3; axis++) if (indices[(face_base + j) * 3 + axis] < vertex_base || indices[(face_base + j) * 3 + axis] >= vertex_base + v) return 0;
            for (axis = 0; axis < 3; axis++) { value = flags & 1 ? varint(&r) : cn[axis]; if (value < -1 || value >= normals) return 0; }
            for (axis = 0; axis < 3; axis++) { value = flags & 2 ? varint(&r) : cu[axis]; if (value < -1 || value >= uvs) return 0; }
            value = flags & 4 ? varint(&r) : material; if (value < -1 || value >= counts[0]) return 0;
        }
        vertex_base += v; face_base += f;
    }
    for (i = 0; i < counts[4] && !r.error; i++) { string(&r); skip(&r, 8); string(&r); string(&r); string(&r); skip(&r, 4); floats(&r, 6); skip(&r, 20); }
    return !r.error && r.pos == n && vertex_base == vertices && face_base == faces;
}
static uint8_t *unpack(const uint8_t *s, size_t n, size_t raw, unsigned method) {
    uint8_t *out;
    if (method == 1) return codec_packbits_decode(s, n, raw);
    if (method || n != raw) return NULL;
    out = (uint8_t *)malloc(raw ? raw : 1); if (out) memcpy(out, s, raw); return out;
}
static int decode_values(const uint8_t *s, size_t n, size_t count, unsigned axes, unsigned method, int32_t *out) {
    size_t axis, i, start; uint8_t *raw; int64_t previous, previous2, value; reader r;
    if (method >= 2) {
        r.s = s; r.n = n; r.pos = 0; r.error = 0;
        for (axis = 0; axis < axes; axis++) {
            previous = previous2 = 0;
            for (i = 0; i < count; i++) { value = (method == 3 && i >= 2 ? 2 * previous - previous2 : previous) + varint(&r); out[i * axes + axis] = (int32_t)(uint32_t)value; previous2 = previous; previous = out[i * axes + axis]; }
        }
        return !r.error && r.pos == n;
    }
    raw = unpack(s, n, count * axes * 4, method); if (!raw) return 0;
    for (axis = 0; axis < axes; axis++) {
        uint32_t sum = 0; start = axis * count * 4;
        for (i = 0; i < count; i++) { sum += raw[start + i] | (uint32_t)raw[start + count + i] << 8 | (uint32_t)raw[start + count * 2 + i] << 16 | (uint32_t)raw[start + count * 3 + i] << 24; out[i * axes + axis] = (int32_t)sum; }
    }
    free(raw); return 1;
}
void forma_free(forma_geometry *g) { if (g) { free(g->xyz); free(g->indices); free(g->scene); memset(g, 0, sizeof(*g)); } }
int forma_decode(const uint8_t *s, size_t n, forma_geometry *g) {
    size_t coord, index, offset, i, scene_size; unsigned cm, im, method;
    if (!g) return 1;
    memset(g, 0, sizeof(*g));
    if (!s || n < 24 || n > LIMIT || memcmp(s, "FMA1", 4) || s[4] != 26 || s[7] || s[5] > 3 || s[6] > 2) return 1;
    cm = s[5]; im = s[6]; g->vertices = u32(s + 8); g->faces = u32(s + 12); coord = u32(s + 16); index = u32(s + 20);
    if (!g->vertices || g->vertices > 4194304 || g->faces > 4194304 || !coord || coord > g->vertices * 12 || index > g->faces * 12 || (!g->faces && index) || coord > n - 24 || index > n - 24 - coord) goto error;
    g->xyz = (int32_t *)malloc(g->vertices * 12); g->indices = (uint32_t *)malloc(g->faces ? g->faces * 12 : 1); if (!g->xyz || !g->indices) goto error;
    if (!decode_values(s + 24, coord, g->vertices, 3, cm, g->xyz) || (g->faces && !decode_values(s + 24 + coord, index, g->faces * 3, 1, im, (int32_t *)g->indices))) goto error;
    for (i = 0; i < g->faces * 3; i++) if (g->indices[i] >= g->vertices) goto error;
    offset = 24 + coord + index;
    if (offset < n) {
        if (n - offset < 20 || memcmp(s + offset, "FMAS", 4) || u32(s + offset + 4) != 1) goto error;
        scene_size = u32(s + offset + 8); method = u32(s + offset + 12);
        if (!scene_size || scene_size > LIMIT || u32(s + offset + 16) != n - offset - 20 || method > 1) goto error;
        g->scene = unpack(s + offset + 20, n - offset - 20, scene_size, method); g->scene_size = scene_size;
        if (!g->scene || !scene_valid(g->scene, scene_size, g->vertices, g->faces, g->indices)) goto error;
    }
    return 0;
error:
    forma_free(g); return 1;
}
static size_t putvar(uint8_t *out, int32_t delta) {
    uint32_t z = ((uint32_t)delta << 1) ^ (uint32_t)-(delta < 0); size_t n = 0;
    do { out[n++] = (uint8_t)((z & 127) | (z >= 128 ? 128 : 0)); z >>= 7; } while (z); return n;
}
static uint8_t *encode_values(const int32_t *values, size_t count, unsigned axes, size_t *size, unsigned *method) {
    size_t i, axis, plane, n = count * axes * 4, packed_size, pos; uint32_t previous, value, delta; int64_t p1, p2, prediction; unsigned candidate_method;
    uint8_t *raw = NULL, *best = NULL, *candidate = NULL;
    raw = (uint8_t *)malloc(n); candidate = (uint8_t *)malloc(count * axes * 5); if (!raw || !candidate) goto error;
    for (axis = 0; axis < axes; axis++) {
        previous = 0;
        for (i = 0; i < count; i++) { value = (uint32_t)values[i * axes + axis]; delta = value - previous; previous = value; for (plane = 0; plane < 4; plane++) raw[axis * count * 4 + plane * count + i] = (uint8_t)(delta >> (plane * 8)); }
    }
    best = codec_packbits_encode(raw, n, &packed_size); if (!best) goto error;
    if (packed_size < n) { *size = packed_size; *method = 1; free(raw); raw = NULL; }
    else { free(best); best = raw; raw = NULL; *size = n; *method = 0; }
    for (candidate_method = 2; candidate_method <= (axes == 3 ? 3u : 2u); candidate_method++) {
        pos = 0;
        for (axis = 0; axis < axes; axis++) {
            p1 = p2 = 0;
            for (i = 0; i < count; i++) { int32_t v = values[i * axes + axis]; prediction = candidate_method == 3 && i >= 2 ? 2 * p1 - p2 : p1; pos += putvar(candidate + pos, (int32_t)(uint32_t)((int64_t)v - prediction)); p2 = p1; p1 = v; }
        }
        if (pos < *size) { uint8_t *replacement = (uint8_t *)malloc(pos); if (!replacement) goto error; memcpy(replacement, candidate, pos); free(best); best = replacement; *size = pos; *method = candidate_method; }
    }
    free(candidate); return best;
error:
    free(raw); free(best); free(candidate); return NULL;
}
int forma_encode(const forma_geometry *g, uint8_t **data, size_t *size) {
    uint8_t *coord = NULL, *index = NULL, *scene = NULL, *out = NULL; size_t cs, is = 0, ss = 0, i, n, offset; unsigned cm, im = 0, sm = 0;
    if (!data || !size) return 1;
    *data = NULL; *size = 0;
    if (!g || !g->xyz || !g->vertices || g->vertices > 4194304 || g->faces > 4194304 || (g->faces && !g->indices)) return 1;
    for (i = 0; i < g->faces * 3; i++) if (g->indices[i] >= g->vertices) return 1;
    if (g->scene_size && !scene_valid(g->scene, g->scene_size, g->vertices, g->faces, g->indices)) return 1;
    coord = encode_values(g->xyz, g->vertices, 3, &cs, &cm); if (!coord) goto error;
    if (g->faces) { index = encode_values((const int32_t *)g->indices, g->faces * 3, 1, &is, &im); if (!index) goto error; }
    if (g->scene_size) {
        scene = codec_packbits_encode(g->scene, g->scene_size, &ss); if (!scene) goto error;
        if (ss < g->scene_size) sm = 1; else { free(scene); ss = g->scene_size; scene = (uint8_t *)malloc(ss); if (!scene) goto error; memcpy(scene, g->scene, ss); }
    }
    n = 24 + cs + is + (ss ? 20 + ss : 0); if (n > LIMIT) goto error;
    out = (uint8_t *)calloc(n, 1); if (!out) goto error;
    memcpy(out, "FMA1", 4); out[4] = 26; out[5] = (uint8_t)cm; out[6] = (uint8_t)im; put32(out + 8, g->vertices); put32(out + 12, g->faces); put32(out + 16, cs); put32(out + 20, is); memcpy(out + 24, coord, cs); if (is) memcpy(out + 24 + cs, index, is);
    if (ss) { offset = 24 + cs + is; memcpy(out + offset, "FMAS", 4); put32(out + offset + 4, 1); put32(out + offset + 8, g->scene_size); put32(out + offset + 12, sm); put32(out + offset + 16, ss); memcpy(out + offset + 20, scene, ss); }
    free(coord); free(index); free(scene); *data = out; *size = n; return 0;
error:
    free(coord); free(index); free(scene); free(out); return 1;
}
