#include "./node_modules/node-addon-api/napi.h"
#include <cstdint>
#include <vector>

using namespace Napi;

////////// Intel ADPCM step variation table ////////////
static int indexTable[16] = {
    -1, -1, -1, -1, 2, 4, 6, 8,
    -1, -1, -1, -1, 2, 4, 6, 8
};

static int stepsizeTable[89] = {
    7,     8,     9,    10,    11,    12,    13,    14,    16,   17,
    19,    21,    23,    25,    28,    31,    34,    37,    41,    45,
    50,    55,    60,    66,    73,    80,    88,    97,   107,   118,
    130,   143,   157,   173,   190,   209,   230,   253,   279,   307,
    337,   371,   408,   449,   494,   544,   598,   658,   724,   796,
    876,   963,  1060,  1166,  1282,  1411,  1552,  1707,  1878,  2066,
    2272,  2499,  2749,  3024,  3327,  3660,  4026,  4428,  4871,  5358,
    5894,  6484,  7132,  7845,  8630,  9493, 10442, 11487, 12635, 13899,
    15289, 16818, 18500, 20350, 22385, 24623, 27086, 29794, 32767
};

struct adpcm_state {
    int valprev;
    int index;
};

void adpcm_coder(const short indata[], char outdata[], int len, adpcm_state *state) {
    // short *inp = const_cast<short*>(indata);
	int len123 = sizeof(indata);
	short *inp = new short[len123 + 1];
	for (int i = 0; i < len123; ++i) {
		inp[i] = indata[i];
	}

    signed char *outp = reinterpret_cast<signed char*>(outdata);

    int valpred = state->valprev;
    int index = state->index;
    int step = stepsizeTable[index];

    int bufferstep = 1;
    int outputbuffer = 0;

    for (; len > 0; len--) {
        int val = *inp++;
        int diff = val - valpred;
        int sign = (diff < 0) ? 8 : 0;
        if (sign) diff = -diff;

        int delta = 0;
        int vpdiff = (step >> 3);

        if (diff >= step) {
            delta = 4;
            diff -= step;
            vpdiff += step;
        }
        step >>= 1;
        if (diff >= step) {
            delta |= 2;
            diff -= step;
            vpdiff += step;
        }
        step >>= 1;
        if (diff >= step) {
            delta |= 1;
            vpdiff += step;
        }

        if (sign) valpred -= vpdiff;
        else valpred += vpdiff;

        if (valpred > 32767) valpred = 32767;
        else if (valpred < -32768) valpred = -32768;

        delta |= sign;

        index += indexTable[delta];
        if (index < 0) index = 0;
        if (index > 88) index = 88;
        step = stepsizeTable[index];

        if (bufferstep) outputbuffer = (delta << 4) & 0xf0;
        else *outp++ = (delta & 0x0f) | outputbuffer;

        bufferstep = !bufferstep;
    }

    if (!bufferstep) *outp++ = outputbuffer;

    state->valprev = valpred;
    state->index = index;
}

void adpcm_decoder(const char indata[], short outdata[], int len, adpcm_state *state) {
    // signed char *inp = const_cast<signed char*>(indata);
	size_t len123 = strlen(indata);
	signed char* inp = new signed char[len123 + 1];
	strcpy(reinterpret_cast<char*>(inp), indata);

    short *outp = outdata;

    int valpred = state->valprev;
    int index = state->index;
    int step = stepsizeTable[index];

    int bufferstep = 0;
    int inputbuffer = 0;

    for (; len > 0; len--) {
        int delta;
        if (bufferstep) delta = inputbuffer & 0xf;
        else {
            inputbuffer = *inp++;
            delta = (inputbuffer >> 4) & 0xf;
        }
        bufferstep = !bufferstep;

        index += indexTable[delta];
        if (index < 0) index = 0;
        if (index > 88) index = 88;

        int sign = delta & 8;
        delta = delta & 7;

        int vpdiff = step >> 3;
        if (delta & 4) vpdiff += step;
        if (delta & 2) vpdiff += step >> 1;
        if (delta & 1) vpdiff += step >> 2;

        if (sign) valpred -= vpdiff;
        else valpred += vpdiff;

        if (valpred > 32767) valpred = 32767;
        else if (valpred < -32768) valpred = -32768;

        step = stepsizeTable[index];
        *outp++ = valpred;
    }

    state->valprev = valpred;
    state->index = index;
}

// Wrapper for adpcm_coder
Napi::Value Encode(const Napi::CallbackInfo &info) {
    Napi::Env env = info.Env();

    if (info.Length() < 3 || !info[0].IsArray() || !info[1].IsArray() || !info[2].IsObject()) {
        Napi::TypeError::New(env, "Expected: (indata: number[], outdata: number[], state: { valprev: number, index: number })").ThrowAsJavaScriptException();
        return env.Null();
    }

    Napi::Array indata = info[0].As<Napi::Array>();
    Napi::Array outdata = info[1].As<Napi::Array>();
    Napi::Object stateObj = info[2].As<Napi::Object>();

    int len = indata.Length();
    std::vector<short> input(len);
    std::vector<char> output(len / 2);

    for (int i = 0; i < len; i++) {
        input[i] = indata.Get(i).ToNumber().Int32Value();
    }

    adpcm_state state;
    state.valprev = stateObj.Get("valprev").ToNumber().Int32Value();
    state.index = stateObj.Get("index").ToNumber().Int32Value();

    adpcm_coder(input.data(), output.data(), len, &state);

    for (int i = 0; i < len / 2; i++) {
        outdata.Set(i, Napi::Number::New(env, output[i]));
    }

    stateObj.Set("valprev", Napi::Number::New(env, state.valprev));
    stateObj.Set("index", Napi::Number::New(env, state.index));

    return env.Null();
}

// Wrapper for adpcm_decoder
Napi::Value Decode(const Napi::CallbackInfo &info) {
    Napi::Env env = info.Env();

    // printf("========================= %d\n", info[0].IsArray());

    if (info.Length() < 3 || !info[0].IsArray() || !info[1].IsArray() || !info[2].IsObject()) {
        Napi::TypeError::New(env, "Expected: (indata: number[], outdata: number[], state: { valprev: number, index: number })").ThrowAsJavaScriptException();
        return env.Null();
    }

    Napi::Array indata = info[0].As<Napi::Array>();
    Napi::Array outdata = info[1].As<Napi::Array>();
    Napi::Object stateObj = info[2].As<Napi::Object>();

    int len = indata.Length();
    std::vector<char> input(len);
    std::vector<short> output(len * 2);

    for (int i = 0; i < len; i++) {
        input[i] = indata.Get(i).ToNumber().Int32Value();
    }

    adpcm_state state;
    state.valprev = stateObj.Get("valprev").ToNumber().Int32Value();
    state.index = stateObj.Get("index").ToNumber().Int32Value();

    // adpcm_decoder(input.data(), output.data(), len * 2, &state);
    adpcm_decoder(input.data(), output.data(), len, &state);

    for (int i = 0; i < len * 2; i++) {
        outdata.Set(i, Napi::Number::New(env, output[i]));
    }

    stateObj.Set("valprev", Napi::Number::New(env, state.valprev));
    stateObj.Set("index", Napi::Number::New(env, state.index));

    return env.Null();
}

// Initialize the addon
Napi::Object Init(Napi::Env env, Napi::Object exports) {
    exports.Set("encode", Napi::Function::New(env, Encode));
    exports.Set("decode", Napi::Function::New(env, Decode));
    return exports;
}

NODE_API_MODULE(adpcm, Init)