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

#ifdef _WIN32
#define EXPORT __declspec(dllexport)
#else
#define EXPORT __attribute__((visibility("default")))
#endif

// Fast 4D Multilinear Interpolator for HM44 (16x16x16x16x4 uint8)
// HM44 is indexed: lut[k, y, m, c, 4]
static inline void interp_hm44_single(
    uint8_t c, uint8_t m, uint8_t y, uint8_t k,
    const uint8_t *hm44_lut,
    uint8_t out[4]
) {
    const float scale = 15.0f / 255.0f;
    float cf = (float)c * scale;
    float mf = (float)m * scale;
    float yf = (float)y * scale;
    float kf = (float)k * scale;

    int c0 = (int)cf; if (c0 > 14) c0 = 14;
    int c1 = c0 + 1;
    float wc = cf - (float)c0;

    int m0 = (int)mf; if (m0 > 14) m0 = 14;
    int m1 = m0 + 1;
    float wm = mf - (float)m0;

    int y0 = (int)yf; if (y0 > 14) y0 = 14;
    int y1 = y0 + 1;
    float wy = yf - (float)y0;

    int k0 = (int)kf; if (k0 > 14) k0 = 14;
    int k1 = k0 + 1;
    float wk = kf - (float)k0;

    float w0000 = (1.0f - wk) * (1.0f - wy) * (1.0f - wm) * (1.0f - wc);
    float w0001 = (1.0f - wk) * (1.0f - wy) * (1.0f - wm) * wc;
    float w0010 = (1.0f - wk) * (1.0f - wy) * wm * (1.0f - wc);
    float w0011 = (1.0f - wk) * (1.0f - wy) * wm * wc;
    float w0100 = (1.0f - wk) * wy * (1.0f - wm) * (1.0f - wc);
    float w0101 = (1.0f - wk) * wy * (1.0f - wm) * wc;
    float w0110 = (1.0f - wk) * wy * wm * (1.0f - wc);
    float w0111 = (1.0f - wk) * wy * wm * wc;

    float w1000 = wk * (1.0f - wy) * (1.0f - wm) * (1.0f - wc);
    float w1001 = wk * (1.0f - wy) * (1.0f - wm) * wc;
    float w1010 = wk * (1.0f - wy) * wm * (1.0f - wc);
    float w1011 = wk * (1.0f - wy) * wm * wc;
    float w1100 = wk * wy * (1.0f - wm) * (1.0f - wc);
    float w1101 = wk * wy * (1.0f - wm) * wc;
    float w1110 = wk * wy * wm * (1.0f - wc);
    float w1111 = wk * wy * wm * wc;

    #define HM44_IDX(k_i, y_i, m_i, c_i) ((((k_i) * 16 + (y_i)) * 16 + (m_i)) * 16 + (c_i)) * 4

    const uint8_t *p0000 = &hm44_lut[HM44_IDX(k0, y0, m0, c0)];
    const uint8_t *p0001 = &hm44_lut[HM44_IDX(k0, y0, m0, c1)];
    const uint8_t *p0010 = &hm44_lut[HM44_IDX(k0, y0, m1, c0)];
    const uint8_t *p0011 = &hm44_lut[HM44_IDX(k0, y0, m1, c1)];
    const uint8_t *p0100 = &hm44_lut[HM44_IDX(k0, y1, m0, c0)];
    const uint8_t *p0101 = &hm44_lut[HM44_IDX(k0, y1, m0, c1)];
    const uint8_t *p0110 = &hm44_lut[HM44_IDX(k0, y1, m1, c0)];
    const uint8_t *p0111 = &hm44_lut[HM44_IDX(k0, y1, m1, c1)];

    const uint8_t *p1000 = &hm44_lut[HM44_IDX(k1, y0, m0, c0)];
    const uint8_t *p1001 = &hm44_lut[HM44_IDX(k1, y0, m0, c1)];
    const uint8_t *p1010 = &hm44_lut[HM44_IDX(k1, y0, m1, c0)];
    const uint8_t *p1011 = &hm44_lut[HM44_IDX(k1, y0, m1, c1)];
    const uint8_t *p1100 = &hm44_lut[HM44_IDX(k1, y1, m0, c0)];
    const uint8_t *p1101 = &hm44_lut[HM44_IDX(k1, y1, m0, c1)];
    const uint8_t *p1110 = &hm44_lut[HM44_IDX(k1, y1, m1, c0)];
    const uint8_t *p1111 = &hm44_lut[HM44_IDX(k1, y1, m1, c1)];

    for (int ch = 0; ch < 4; ch++) {
        float val =
            (float)p0000[ch] * w0000 + (float)p0001[ch] * w0001 +
            (float)p0010[ch] * w0010 + (float)p0011[ch] * w0011 +
            (float)p0100[ch] * w0100 + (float)p0101[ch] * w0101 +
            (float)p0110[ch] * w0110 + (float)p0111[ch] * w0111 +
            (float)p1000[ch] * w1000 + (float)p1001[ch] * w1001 +
            (float)p1010[ch] * w1010 + (float)p1011[ch] * w1011 +
            (float)p1100[ch] * w1100 + (float)p1101[ch] * w1101 +
            (float)p1110[ch] * w1110 + (float)p1111[ch] * w1111;

        int ival = (int)(val + 0.5f);
        if (ival < 0) ival = 0;
        if (ival > 255) ival = 255;
        out[ch] = (uint8_t)ival;
    }
}

// Fast 4D Multilinear Interpolator for ICC mft1 (21x21x21x21x4)
// CLUT is indexed: clut[c, m, y, k, 4]
static inline void interp_icc_single(
    uint8_t c, uint8_t m, uint8_t y, uint8_t k,
    const uint8_t in_tables[4][256],
    const uint8_t *clut,
    const uint8_t out_tables[4][256],
    int grid_pts,
    uint8_t out[4]
) {
    uint8_t c_in = in_tables[0][c];
    uint8_t m_in = in_tables[1][m];
    uint8_t y_in = in_tables[2][y];
    uint8_t k_in = in_tables[3][k];

    float scale = (float)(grid_pts - 1) / 255.0f;
    float cf = (float)c_in * scale;
    float mf = (float)m_in * scale;
    float yf = (float)y_in * scale;
    float kf = (float)k_in * scale;

    int max_idx = grid_pts - 2;
    int c0 = (int)cf; if (c0 > max_idx) c0 = max_idx;
    int c1 = c0 + 1;
    float wc = cf - (float)c0;

    int m0 = (int)mf; if (m0 > max_idx) m0 = max_idx;
    int m1 = m0 + 1;
    float wm = mf - (float)m0;

    int y0 = (int)yf; if (y0 > max_idx) y0 = max_idx;
    int y1 = y0 + 1;
    float wy = yf - (float)y0;

    int k0 = (int)kf; if (k0 > max_idx) k0 = max_idx;
    int k1 = k0 + 1;
    float wk = kf - (float)k0;

    float w0000 = (1.0f - wk) * (1.0f - wy) * (1.0f - wm) * (1.0f - wc);
    float w0001 = (1.0f - wk) * (1.0f - wy) * (1.0f - wm) * wc;
    float w0010 = (1.0f - wk) * (1.0f - wy) * wm * (1.0f - wc);
    float w0011 = (1.0f - wk) * (1.0f - wy) * wm * wc;
    float w0100 = (1.0f - wk) * wy * (1.0f - wm) * (1.0f - wc);
    float w0101 = (1.0f - wk) * wy * (1.0f - wm) * wc;
    float w0110 = (1.0f - wk) * wy * wm * (1.0f - wc);
    float w0111 = (1.0f - wk) * wy * wm * wc;

    float w1000 = wk * (1.0f - wy) * (1.0f - wm) * (1.0f - wc);
    float w1001 = wk * (1.0f - wy) * (1.0f - wm) * wc;
    float w1010 = wk * (1.0f - wy) * wm * (1.0f - wc);
    float w1011 = wk * (1.0f - wy) * wm * wc;
    float w1100 = wk * wy * (1.0f - wm) * (1.0f - wc);
    float w1101 = wk * wy * (1.0f - wm) * wc;
    float w1110 = wk * wy * wm * (1.0f - wc);
    float w1111 = wk * wy * wm * wc;

    #define ICC_IDX(c_i, m_i, y_i, k_i) ((((c_i) * grid_pts + (m_i)) * grid_pts + (y_i)) * grid_pts + (k_i)) * 4

    const uint8_t *p0000 = &clut[ICC_IDX(c0, m0, y0, k0)];
    const uint8_t *p0001 = &clut[ICC_IDX(c1, m0, y0, k0)];
    const uint8_t *p0010 = &clut[ICC_IDX(c0, m1, y0, k0)];
    const uint8_t *p0011 = &clut[ICC_IDX(c1, m1, y0, k0)];
    const uint8_t *p0100 = &clut[ICC_IDX(c0, m0, y1, k0)];
    const uint8_t *p0101 = &clut[ICC_IDX(c1, m0, y1, k0)];
    const uint8_t *p0110 = &clut[ICC_IDX(c0, m1, y1, k0)];
    const uint8_t *p0111 = &clut[ICC_IDX(c1, m1, y1, k0)];

    const uint8_t *p1000 = &clut[ICC_IDX(c0, m0, y0, k1)];
    const uint8_t *p1001 = &clut[ICC_IDX(c1, m0, y0, k1)];
    const uint8_t *p1010 = &clut[ICC_IDX(c0, m1, y0, k1)];
    const uint8_t *p1011 = &clut[ICC_IDX(c1, m1, y0, k1)];
    const uint8_t *p1100 = &clut[ICC_IDX(c0, m0, y1, k1)];
    const uint8_t *p1101 = &clut[ICC_IDX(c1, m0, y1, k1)];
    const uint8_t *p1110 = &clut[ICC_IDX(c0, m1, y1, k1)];
    const uint8_t *p1111 = &clut[ICC_IDX(c1, m1, y1, k1)];

    for (int ch = 0; ch < 4; ch++) {
        float val =
            (float)p0000[ch] * w0000 + (float)p0001[ch] * w0001 +
            (float)p0010[ch] * w0010 + (float)p0011[ch] * w0011 +
            (float)p0100[ch] * w0100 + (float)p0101[ch] * w0101 +
            (float)p0110[ch] * w0110 + (float)p0111[ch] * w0111 +
            (float)p1000[ch] * w1000 + (float)p1001[ch] * w1001 +
            (float)p1010[ch] * w1010 + (float)p1011[ch] * w1011 +
            (float)p1100[ch] * w1100 + (float)p1101[ch] * w1101 +
            (float)p1110[ch] * w1110 + (float)p1111[ch] * w1111;

        int ival = (int)(val + 0.5f);
        if (ival < 0) ival = 0;
        if (ival > 255) ival = 255;
        out[ch] = out_tables[ch][ival];
    }
}

// Full Pipeline Processor: Processes CMYK buffer directly, updates 256-bin histograms
// Optionally writes out transformed CMYK buffer for preview/TAC generation
EXPORT void process_cmyk_pipeline(
    const uint8_t *input_cmyk,          // (num_pixels * 4)
    size_t num_pixels,
    int use_icc,
    const uint8_t icc_in_tables[4][256],
    const uint8_t *icc_clut,
    const uint8_t icc_out_tables[4][256],
    int icc_grid_pts,
    const uint8_t *hm44_lut,
    uint8_t *output_cmyk,               // optional, if not NULL, writes transformed (num_pixels * 4)
    uint64_t hist_c[256],
    uint64_t hist_m[256],
    uint64_t hist_y[256],
    uint64_t hist_k[256],
    uint64_t *tac_histogram            // optional, if not NULL, 401 bins for TAC 0..400%
) {
    uint8_t icc_out[4];
    uint8_t final_out[4];

    for (size_t i = 0; i < num_pixels; i++) {
        size_t idx = i * 4;
        uint8_t c = input_cmyk[idx];
        uint8_t m = input_cmyk[idx + 1];
        uint8_t y = input_cmyk[idx + 2];
        uint8_t k = input_cmyk[idx + 3];

        if (use_icc) {
            interp_icc_single(c, m, y, k, icc_in_tables, icc_clut, icc_out_tables, icc_grid_pts, icc_out);
        } else {
            icc_out[0] = c;
            icc_out[1] = m;
            icc_out[2] = y;
            icc_out[3] = k;
        }

        interp_hm44_single(icc_out[0], icc_out[1], icc_out[2], icc_out[3], hm44_lut, final_out);

        hist_c[final_out[0]]++;
        hist_m[final_out[1]]++;
        hist_y[final_out[2]]++;
        hist_k[final_out[3]]++;

        if (output_cmyk) {
            output_cmyk[idx] = final_out[0];
            output_cmyk[idx + 1] = final_out[1];
            output_cmyk[idx + 2] = final_out[2];
            output_cmyk[idx + 3] = final_out[3];
        }

        if (tac_histogram) {
            int tac_pct = (int)(((int)final_out[0] + (int)final_out[1] + (int)final_out[2] + (int)final_out[3]) * 100 / 255);
            if (tac_pct > 400) tac_pct = 400;
            tac_histogram[tac_pct]++;
        }
    }
}
