/**
 * @file rgb_power.c
 * @brief RGB 恒功率修正算法实现
 *
 * 定点化规则：全程使用 int32_t，精度放大 FIXED_SCALE(10000) 倍
 * 除零保护：A==0 时（三路全黑）直接输出全零
 */

#include "rgb_power.h"
#include "hc32f003_config.h"

void RGB_PowerCorrect(uint8_t r, uint8_t g, uint8_t b,
                      uint8_t *rt, uint8_t *gt, uint8_t *bt)
{
    /* --- 步骤 1：A = max(R, G, B) --- */
    int32_t A = (int32_t)MAX(r, MAX(g, b));

    /* --- 除零保护：A == 0 → 三路全黑直接输出零 --- */
    if (A == 0) {
        *rt = 0;
        *gt = 0;
        *bt = 0;
        return;
    }

    /*
     * 步骤 2：P = A / 255.0
     * 定点：P_fp = A * FIXED_SCALE / 255
     */
    int32_t P_fp = A * FIXED_SCALE / 255;   /* P * FIXED_SCALE */

    /*
     * 步骤 3：B1 = (R + G + B) - 255 * P
     * 定点：B1_fp = (R+G+B)*FIXED_SCALE - 255 * P_fp
     *           = B1 * FIXED_SCALE
     */
    int32_t sum_rgb = (int32_t)r + g + b;
    int32_t B1_fp = sum_rgb * FIXED_SCALE - 255 * P_fp;   /* B1 * FIXED_SCALE */

    /*
     * 步骤 4：C_sum = (R+G+B) / (255*P)
     * 定点：C_sum_fp = sum_rgb * FIXED_SCALE / (255 * P_fp / FIXED_SCALE)
     *               = sum_rgb * FIXED_SCALE * FIXED_SCALE / (255 * P_fp)
     * 使用 int64_t 避免溢出（sum_rgb <= 765, FIXED_SCALE^2 = 1e8）
     */
    int32_t denom_C = 255 * P_fp;   /* 255 * P * FIXED_SCALE */
    if (denom_C == 0) {
        /* 理论上 A>0 时不会到这，保险起见 */
        *rt = r; *gt = g; *bt = b;
        return;
    }
    /* C_sum_fp = C_sum * FIXED_SCALE */
    int64_t C_sum_fp = (int64_t)sum_rgb * FIXED_SCALE * FIXED_SCALE / denom_C;

    /*
     * 步骤 5：D = B1 / C_sum
     * 定点：D_fp = B1_fp * FIXED_SCALE / C_sum_fp
     *           = D * FIXED_SCALE
     * 注意：B1_fp 单位 FIXED_SCALE，C_sum_fp 单位 FIXED_SCALE
     * D_fp = B1_fp * FIXED_SCALE / C_sum_fp → D * FIXED_SCALE
     */
    if (C_sum_fp == 0) {
        *rt = r; *gt = g; *bt = b;
        return;
    }
    int64_t D_fp = (int64_t)B1_fp * FIXED_SCALE / C_sum_fp;   /* D * FIXED_SCALE */

    /*
     * 步骤 6-8：RT = round(R - R/(255*P) * D)
     *   = round(R - R * D / (255*P))
     *
     * 定点：
     *   R_over_255P_fp = R * FIXED_SCALE / (255 * P_fp / FIXED_SCALE)
     *                  = R * FIXED_SCALE^2 / (255 * P_fp)
     *   correction_fp  = R_over_255P_fp * D_fp / FIXED_SCALE
     *                  = R/(255*P) * D * FIXED_SCALE
     *   result         = round((R*FIXED_SCALE - correction_fp) / FIXED_SCALE)
     */

    /* 辅助 Lambda：计算单通道修正后值，channel 为 0-255 原始值 */
    /* 使用 int64_t 避免溢出 */
    {
        /* 公因子：1 / (255*P) * D ，以 FIXED_SCALE 为单位 */
        /* factor_fp = D / (255*P) * FIXED_SCALE^2 / FIXED_SCALE = D/(255*P)*FIXED_SCALE */
        /* 即 factor_fp = D_fp / P_fp / 255 * FIXED_SCALE */
        /* 用 int64_t: factor_fp = D_fp * FIXED_SCALE / (255 * P_fp) */
        int64_t factor_fp = D_fp * FIXED_SCALE / denom_C;   /* D/(255*P) * FIXED_SCALE */

        /* RT = round(R - R*factor) = round(R*(1 - factor)) */
        /* (1-factor) * FIXED_SCALE = FIXED_SCALE - factor_fp */
        int64_t one_minus_factor_fp = (int64_t)FIXED_SCALE - factor_fp;

        /* RT_fp = R * one_minus_factor_fp（单位 FIXED_SCALE） */
        int64_t RT_fp = (int64_t)r * one_minus_factor_fp;
        int64_t GT_fp = (int64_t)g * one_minus_factor_fp;
        int64_t BT_fp = (int64_t)b * one_minus_factor_fp;

        /* 四舍五入并缩回 [0,255] */
        int32_t rt_val = (int32_t)((RT_fp + FIXED_SCALE / 2) / FIXED_SCALE);
        int32_t gt_val = (int32_t)((GT_fp + FIXED_SCALE / 2) / FIXED_SCALE);
        int32_t bt_val = (int32_t)((BT_fp + FIXED_SCALE / 2) / FIXED_SCALE);

        *rt = (uint8_t)CLAMP(rt_val, 0, 255);
        *gt = (uint8_t)CLAMP(gt_val, 0, 255);
        *bt = (uint8_t)CLAMP(bt_val, 0, 255);
    }
}
