/**
 * @file lut.c
 * @brief 查找表默认数据 + CW→RGBWC 插值算法实现
 *
 * 算法一参考（JavaScript 原版）：
 *   function runCalc(inputW, inputC)  — 选择插值层和亮度系数 T
 *   function interp(group, w1_target, T) — 线性插值
 *
 * 定点化规则：所有比例系数 × FIXED_SCALE(10000) 后用 int32_t 运算，
 *             最终结果 / FIXED_SCALE 并 clamp 到 [0,255]
 */

#include "lut.h"
#include <string.h>

/* =========================================================
 * 内置默认查找表（Flash 无效时使用）
 * 节点按 w1 降序存储（w1=255 在前，w1=0 在后）
 * ========================================================= */

/** Layer 0：100% 亮度组，13 个有效点，剩余填零 */
static const LutPoint_t s_default_layer0[LUT_POINTS_PER_LAYER] = {
    {255,   0,   0,   0,   0, 255,   0},
    {245,  10, 111,  25,   4, 115,   0},
    {232,  23,  81,  52,   7, 115,   0},
    {224,  31,  93,  64,  14,  94,   0},
    {200,  55,  70,  98,  31,  57,   0},
    {175,  80,  55, 116,  46,  38,   0},
    {157,  98,  66, 117,  55,  17,   0},
    {139, 116,  46, 137,  63,  11,   0},
    {125, 130,  36, 123,  72,  24,   0},
    {100, 155,  26, 125,  81,  23,   0},
    { 75, 180,  20, 126,  89,  20,   0},
    { 50, 205,  15, 127,  95,  19,   0},
    {  0, 255,   5, 126, 105,  19,   0},
    {0, 0, 0, 0, 0, 0, 0},
    {0, 0, 0, 0, 0, 0, 0},
    {0, 0, 0, 0, 0, 0, 0},
    {0, 0, 0, 0, 0, 0, 0},
    {0, 0, 0, 0, 0, 0, 0},
    {0, 0, 0, 0, 0, 0, 0},
    {0, 0, 0, 0, 0, 0, 0},
};

/** Layer 1：70% 亮度组，13 个有效点，剩余填零 */
static const LutPoint_t s_default_layer1[LUT_POINTS_PER_LAYER] = {
    {186,   0, 102,   6,   1,  77,   0},
    {179,   7,  81,  18,   3,  84,   0},
    {169,  17,  59,  38,   5,  84,   0},
    {163,  23,  68,  47,  10,  69,   0},
    {146,  40,  51,  71,  23,  42,   0},
    {128,  58,  40,  85,  34,  28,   0},
    {115,  71,  48,  85,  40,  12,   0},
    {101,  85,  34, 100,  46,   8,   0},
    { 91,  95,  26,  90,  53,  18,   0},
    { 73, 113,  19,  91,  59,  17,   0},
    { 55, 131,  15,  92,  65,  15,   0},
    { 36, 150,  11,  93,  69,  14,   0},
    {  0, 186,   4,  92,  77,  14,   0},
    {0, 0, 0, 0, 0, 0, 0},
    {0, 0, 0, 0, 0, 0, 0},
    {0, 0, 0, 0, 0, 0, 0},
    {0, 0, 0, 0, 0, 0, 0},
    {0, 0, 0, 0, 0, 0, 0},
    {0, 0, 0, 0, 0, 0, 0},
    {0, 0, 0, 0, 0, 0, 0},
};

/* =========================================================
 * 全局查找表实例
 * ========================================================= */
LutTable_t g_lut;

/* =========================================================
 * 内部函数：统计某层有效节点数（w1 或 c1 非零视为有效）
 * ========================================================= */
static int LUT_CountValid(const LutPoint_t *layer)
{
    int cnt = 0;
    for (int i = 0; i < LUT_POINTS_PER_LAYER; i++) {
        if (layer[i].w1 != 0 || layer[i].c1 != 0
            || layer[i].r != 0 || layer[i].g != 0
            || layer[i].b != 0 || layer[i].w != 0) {
            cnt++;
        } else {
            break;  /* 遇到全零节点则停止（降序表末尾填充） */
        }
    }
    /* 至少保留 1 个节点 */
    return cnt > 0 ? cnt : 1;
}

/* =========================================================
 * LUT_LoadDefault：加载内置默认查找表
 * ========================================================= */
void LUT_LoadDefault(void)
{
    memcpy(g_lut.layer[0], s_default_layer0, sizeof(s_default_layer0));
    memcpy(g_lut.layer[1], s_default_layer1, sizeof(s_default_layer1));
    g_lut.crc = 0;  /* 默认表不写 Flash，CRC 留 0 */
}

/* =========================================================
 * 内部函数：整数定点线性插值
 *
 * 原 JS interp(group, w1_target, T)：
 *   - 查找表按 w1 降序，找最后一个 w1 >= w1_target 的点 p1
 *   - p2 是 p1 的下一个点（w1 < w1_target）
 *   - ratio = (w1_target - p1.w1) / (p2.w1 - p1.w1)
 *     注意分母为负（p2.w1 < p1.w1），比例结果在 [0,1]
 *   - out_ch = clamp(round((p1.ch + ratio*(p2.ch-p1.ch)) * T), 0, 255)
 *
 * 定点化：T_fp = T * FIXED_SCALE（整数），ratio 用 int32_t 分子/分母表示
 *
 * @param layer    查找表层指针（降序，有效节点数 n_valid）
 * @param n_valid  有效节点数
 * @param w1_target 目标 w1 值（定点乘 FIXED_SCALE 前的原始值，但这里传整数）
 * @param T_fp     亮度系数 T 放大 FIXED_SCALE 倍后的整数（如 T=1 → T_fp=10000）
 * @param out      输出结果
 * ========================================================= */
static void LUT_Interp(const LutPoint_t *layer, int n_valid,
                        int32_t w1_target_fp,  /* w1_target * FIXED_SCALE */
                        int32_t T_fp,
                        RgbwcOut_t *out)
{
    /*
     * 说明：w1_target_fp 是放大后的目标值，layer[i].w1 是原始 0-255 整数。
     * 比较时统一乘 FIXED_SCALE：layer[i].w1 * FIXED_SCALE vs w1_target_fp
     */

    int p1_idx = -1;
    int p2_idx = -1;

    /* 1. 找最后一个 w1*FIXED_SCALE >= w1_target_fp 的节点作为 p1 */
    for (int i = 0; i < n_valid; i++) {
        int32_t w1_fp = (int32_t)layer[i].w1 * FIXED_SCALE;
        if (w1_fp >= w1_target_fp) {
            p1_idx = i;  /* 降序遍历，持续更新取最后一个（最小满足条件的 w1） */
        }
    }

    if (p1_idx < 0) {
        /* w1_target 比所有节点都大：使用第一个节点（w1 最大） */
        p1_idx = 0;
    }

    /* 2. p2 是 p1 的下一个节点 */
    if (p1_idx + 1 < n_valid) {
        p2_idx = p1_idx + 1;
    } else {
        /* p1 已是最后一个节点，无 p2，直接用 p1 值乘 T */
        p2_idx = -1;
    }

    /* 3. 对每个通道做插值 */
    const uint8_t *p1_vals = &layer[p1_idx].r;  /* r,g,b,w,c 连续 5 字节 */
    const uint8_t *p2_vals = (p2_idx >= 0) ? &layer[p2_idx].r : &layer[p1_idx].r;

    int32_t w1_p1_fp = (int32_t)layer[p1_idx].w1 * FIXED_SCALE;
    int32_t w1_p2_fp = (p2_idx >= 0) ? (int32_t)layer[p2_idx].w1 * FIXED_SCALE
                                      : w1_p1_fp;

    /*
     * ratio = (w1_target - p1.w1) / (p2.w1 - p1.w1)
     * 分子：w1_target_fp - w1_p1_fp  （通常 <= 0，因为 p1.w1 >= w1_target）
     * 分母：w1_p2_fp - w1_p1_fp      （< 0，因为降序 p2.w1 < p1.w1）
     * ratio 结果在 [0, 1]，ratio_fp = ratio * FIXED_SCALE（整数）
     */

    int32_t ratio_fp = 0;
    int32_t denom = w1_p2_fp - w1_p1_fp;  /* 分母 */
    if (denom != 0) {
        int32_t numer = w1_target_fp - w1_p1_fp;  /* 分子 */
        /* 放大 FIXED_SCALE 倍避免浮点，结果 ratio_fp in [0, FIXED_SCALE] */
        ratio_fp = (numer * FIXED_SCALE) / denom;
        /* clamp ratio 到 [0, FIXED_SCALE] 避免越界 */
        if (ratio_fp < 0) ratio_fp = 0;
        if (ratio_fp > FIXED_SCALE) ratio_fp = FIXED_SCALE;
    }

    /* 通道指针映射：r=0, g=1, b=2, w=3, c=4 */
    uint8_t *out_vals[5] = {&out->r, &out->g, &out->b, &out->w, &out->c};

    for (int ch = 0; ch < 5; ch++) {
        /*
         * 插值公式（JS 原版）：
         *   val = p1.ch + ratio * (p2.ch - p1.ch)
         *   out = clamp(round(val * T), 0, 255)
         *
         * 定点化：
         *   val_fp = p1.ch * FIXED_SCALE + ratio_fp * (p2.ch - p1.ch)
         *   (ratio_fp 已是 ratio * FIXED_SCALE，所以 val_fp = val * FIXED_SCALE)
         *   after_T_fp = val_fp * T_fp  →  单位 FIXED_SCALE^2
         *   result = round(after_T_fp / (FIXED_SCALE * FIXED_SCALE))
         */
        int32_t p1_ch = (int32_t)p1_vals[ch];
        int32_t p2_ch = (int32_t)p2_vals[ch];

        /* val * FIXED_SCALE */
        int32_t val_fp = p1_ch * FIXED_SCALE + ratio_fp * (p2_ch - p1_ch);

        /* val * T，单位 FIXED_SCALE^2，用 int64_t 避免溢出 */
        int64_t after_T = (int64_t)val_fp * T_fp;

        /* 四舍五入并缩回原始量纲 */
        int32_t divisor = (int32_t)FIXED_SCALE * FIXED_SCALE;
        int32_t result = (int32_t)((after_T + divisor / 2) / divisor);

        *out_vals[ch] = (uint8_t)CLAMP(result, 0, 255);
    }
}

/* =========================================================
 * LUT_CalcCW：算法一主函数
 *
 * 移植自 JS runCalc(inputW, inputC)：
 *   - total >= 253          → layer[0], T=1
 *   - total > 187           → layer[0], T=total/255, W1=inputW/T
 *   - 185 <= total <= 187   → layer[1], T=1
 *   - total < 185           → layer[1], T=total/186, W1=inputW/T
 * ========================================================= */
void LUT_CalcCW(uint8_t inputW, uint8_t inputC, RgbwcOut_t *out)
{
    int total = (int)inputW + (int)inputC;

    /* 选择插值层和计算 T_fp, w1_target_fp */
    int    layer_idx;
    int32_t T_fp;          /* T * FIXED_SCALE */
    int32_t w1_target_fp;  /* w1_target * FIXED_SCALE */

    if (total >= 253) {
        /* JS: interp(layer[0], inputW, 1.0) — 直接用 layer[0]，T=1 */
        layer_idx    = 0;
        T_fp         = FIXED_SCALE;                       /* T = 1.0 */
        w1_target_fp = (int32_t)inputW * FIXED_SCALE;    /* W1 = inputW */
    }
    else if (total > 187) {
        /*
         * JS: T = total/255, W1 = inputW/T → W1 = inputW*255/total
         * 定点：T_fp = total * FIXED_SCALE / 255
         *        w1_target_fp = inputW * FIXED_SCALE * 255 / total
         *        （等价于 inputW * FIXED_SCALE / T）
         */
        layer_idx    = 0;
        T_fp         = (int32_t)total * FIXED_SCALE / 255;
        /* W1 = inputW / T = inputW * 255 / total */
        w1_target_fp = (int32_t)inputW * 255 * FIXED_SCALE / total;
    }
    else if (total >= 185) {
        /* JS: interp(layer[1], inputW, 1.0) */
        layer_idx    = 1;
        T_fp         = FIXED_SCALE;
        w1_target_fp = (int32_t)inputW * FIXED_SCALE;
    }
    else {
        /*
         * JS: T = total/186, W1 = inputW/T = inputW*186/total
         * 防止 total==0 时除零
         */
        layer_idx = 1;
        if (total == 0) {
            /* 输入全零 → 输出全零 */
            out->r = out->g = out->b = out->w = out->c = 0;
            return;
        }
        T_fp         = (int32_t)total * FIXED_SCALE / 186;
        w1_target_fp = (int32_t)inputW * 186 * FIXED_SCALE / total;
    }

    /* 统计有效节点数 */
    int n_valid = LUT_CountValid(g_lut.layer[layer_idx]);

    /* 执行插值 */
    LUT_Interp(g_lut.layer[layer_idx], n_valid, w1_target_fp, T_fp, out);
}
