#include "example.h"
#include <jsi/jsi.h>
#include <iostream>

using namespace facebook;

extern "C" {
    #define STB_IMAGE_IMPLEMENTATION
    #include "stb_image.h"
}

bool load_image(std::vector<unsigned char>& image, const std::string& filename, int& x, int& y);
std::tuple<int, int, int> rgb2lab(int r, int g, int b);
float sRGBCompand(float x);
float labCompand(float x);
double deg2rad(double deg);
double rad2deg(double rad);

namespace example {

void install(jsi::Runtime &jsiRuntime) {
    
    auto getPixelColor = jsi::Function::createFromHostFunction(jsiRuntime,
                                                          jsi::PropNameID::forAscii(jsiRuntime,
                                                                               "getPixelColor"),
                                                          3,
                                                          [](jsi::Runtime &runtime,
                                                             const jsi::Value &thisValue,
                                                             const jsi::Value *arguments,
                                                             size_t count) -> jsi::String {
        if (count != 3) {
            throw jsi::JSError(runtime, "JSIColor.getPixelColor(...) expects 3 arguments (imageFile, x, y)!");
        }

        std::string imageFile = arguments[0].getString(runtime).utf8(runtime);
        int x = arguments[1].getNumber();
        int y = arguments[2].getNumber();

        int width, height;
        
        std::vector<unsigned char> image;
        bool success = load_image(image, imageFile, width, height);

        if (!success) {
            throw jsi::JSError(runtime, "Failed to load image!");
        }

        const size_t RGBA = 4;
        size_t index = RGBA * (y * width + x);

        int r = static_cast<int>(image[index + 0]);
        int g = static_cast<int>(image[index + 1]);
        int b = static_cast<int>(image[index + 2]);

        std::string color = "rgb(" + std::to_string(r) + "," + std::to_string(g) + "," + std::to_string(b) + ")";

        return jsi::String::createFromUtf8(runtime, color);
    });
    
    jsiRuntime.global().setProperty(jsiRuntime, "getPixelColor", std::move(getPixelColor));
    
    auto computeColorDistance = jsi::Function::createFromHostFunction(jsiRuntime,
                                                          jsi::PropNameID::forAscii(jsiRuntime,
                                                                               "getPixelColor"),
                                                          6,
                                                          [](jsi::Runtime &runtime,
                                                             const jsi::Value &thisValue,
                                                             const jsi::Value *arguments,
                                                             size_t count) -> jsi::Value {
        if (count != 6) {
            throw jsi::JSError(runtime, "JSIColor.convertRGBtoLAB(...) expects 6 arguments (r1, g1, b1, r2, g2, b2)!");
        }
        
        const double kL = 1.0, kC = 1.0, kH = 1.0;
        const double deg360InRad = deg2rad(360.0);
        const double deg180InRad = deg2rad(180.0);
        const double pow25To7 = 6103515625.0; /* pow(25, 7) */

        int r1 = arguments[0].getNumber();
        int g1 = arguments[1].getNumber();
        int b1 = arguments[2].getNumber();

        int r2 = arguments[3].getNumber();
        int g2 = arguments[4].getNumber();
        int b2 = arguments[5].getNumber();

        std::tuple<int, int, int> lab1 = rgb2lab(r1, g1, b1);
        std::tuple<int, int, int> lab2 = rgb2lab(r2, g2, b2);

        int L1 = std::get<0>(lab1);
        int a1 = std::get<1>(lab1);
        int b1_ = std::get<2>(lab1);

        int L2 = std::get<0>(lab2);
        int a2 = std::get<1>(lab2);
        int b2_ = std::get<2>(lab2);

        // Step 1
        double C1 = sqrt((a1 * a1) + (b1_ * b1_));
        double C2 = sqrt((a2 * a2) + (b2_ * b2_));

        double barC = (C1 + C2) / 2.0;
        double G = 0.5 * (1 - sqrt(pow(barC, 7.0) / (pow(barC, 7.0) + pow25To7)));
        
        double a1Prime = (1.0 + G) * a1;
        double a2Prime = (1.0 + G) * a2;

        double CPrime1 = sqrt((a1Prime * a1Prime) + (b1_ * b1_));
        double CPrime2 = sqrt((a2Prime * a2Prime) + (b2_ * b2_));

        double hPrime1;
        if (b1_ == 0 && a1Prime == 0) {
            hPrime1 = 0.0;
        } else {
            hPrime1 = atan2(b1_, a1Prime);
            
            if (hPrime1 < 0) {
                hPrime1 += deg360InRad;
            }
        }

        double hPrime2;
        if (b2_ == 0 && a2Prime == 0) {
            hPrime2 = 0.0;
        } else {
            hPrime2 = atan2(b2_, a2Prime);
            
            if (hPrime2 < 0) {
                hPrime2 += deg360InRad;
            }
        }
        
        // Step 2
        double deltaLPrime = L2 - L1;
        double deltaCPrime = CPrime2 - CPrime1;
        double deltahPrime;
        double CPrimeProduct = CPrime1 * CPrime2;
        if (CPrimeProduct == 0) {
            deltahPrime = 0;
        } else {
            deltahPrime = hPrime2 - hPrime1;
            
            if (deltahPrime < -deg180InRad)
                deltahPrime += deg360InRad;
            else if (deltahPrime > deg180InRad)
                deltahPrime -= deg360InRad;
        }
        
        double deltaHPrime = 2.0 * sqrt(CPrimeProduct) * sin(deltahPrime / 2.0);
        
        // Step 3
        double barLPrime = (L1 + L2) / 2.0;
        double barCPrime = (CPrime1 + CPrime2) / 2.0;
        double barhPrime, hPrimeSum = hPrime1 + hPrime2;
        if (CPrime1 * CPrime2 == 0) {
            barhPrime = hPrimeSum;
        } else {
            if (fabs(hPrime1 - hPrime2) <= deg180InRad)
                barhPrime = hPrimeSum / 2.0;
            else {
                if (hPrimeSum < deg360InRad)
                    barhPrime = (hPrimeSum + deg360InRad) / 2.0;
                else
                    barhPrime = (hPrimeSum - deg360InRad) / 2.0;
            }
        }
        
        double T = 1.0 - (0.17 * cos(barhPrime - deg2rad(30.0))) + 
            (0.24 * cos(2.0 * barhPrime)) +
            (0.32 * cos(3.0 * barhPrime + deg2rad(6.0))) - 
            (0.20 * cos(4.0 * barhPrime - deg2rad(63.0)));
        
        double deltaTheta = deg2rad(30.0) * exp(-pow((barhPrime - deg2rad(275.0)) / deg2rad(25.0), 2.0));
        double RC = 2.0 * sqrt(pow(barCPrime, 7.0) / (pow(barCPrime, 7.0) + pow25To7));
        double SL = 1.0 + ((0.015 * pow(barLPrime - 50.0, 2.0)) / sqrt(20.0 + pow(barLPrime - 50.0, 2.0)));
        double SC = 1.0 + 0.045 * barCPrime;
        double SH = 1.0 + 0.015 * barCPrime * T;
        double RT = -sin(2.0 * deltaTheta) * RC;

        int deltaE = sqrt(
            pow(deltaLPrime / (kL * SL), 2.0) +
            pow(deltaCPrime / (kC * SC), 2.0) +
            pow(deltaHPrime / (kH * SH), 2.0) +
            RT * (deltaCPrime / (kC * SC)) * (deltaHPrime / (kH * SH))
        );
            
        return jsi::Value(deltaE);
    });

    jsiRuntime.global().setProperty(jsiRuntime, "computeColorDistance", std::move(computeColorDistance));
}

}

bool load_image(std::vector<unsigned char>& image, const std::string& filename, int& x, int& y) {
    int n;
    unsigned char* data = stbi_load(filename.c_str(), &x, &y, &n, 4);
    if (data != nullptr) {
        image = std::vector<unsigned char>(data, data + x * y * 4);
    }

    stbi_image_free(data);
    return (data != nullptr);
}

std::tuple<int, int, int> rgb2lab(int r, int g, int b) {
    float r_ = sRGBCompand(r);
    float g_ = sRGBCompand(g);
    float b_ = sRGBCompand(b);

    float x = 0.412453 * r_ + 0.357580 * g_ + 0.180423 * b_;
    float y = 0.212671 * r_ + 0.715160 * g_ + 0.072169 * b_;
    float z = 0.019334 * r_ + 0.119193 * g_ + 0.950227 * b_;

    float fx = labCompand(x / 0.950456);
    float fy = labCompand(y);
    float fz = labCompand(z / 1.088754);

    float L = 116.0 * fy - 16.0;
    float a = 500.0 * (fx - fy);
    float b__ = 200.0 * (fy - fz);

    return std::make_tuple(L, a, b__);
}

float sRGBCompand(float x) {
    float absX = std::abs(x);
    float out = absX > 0.04045 ? pow((absX + 0.055) / 1.055, 2.4) : absX / 12.92;

    return x > 0 ? out : -out;
}

float labCompand(float x) {
    return x > 0.008856 ? pow(x, 1.0 / 3.0) : (7.787036 * x) + 0.1379310;
}

double deg2rad(double deg) {
    return deg * M_PI / 180.0;
}

double rad2deg(double rad) {
    return rad * 180.0 / M_PI;
}
