#include <stdlib.h>
#include <stdio.h>
#include <limits.h>

#include <lib/def.h>
#include <lib/xmem.h>
#include <lib/error.h>

#include <vocoder/dft.h>
#include <vocoder/window_function.h>
#include <vocoder/filter.h>

#include <context.h>


static positive long gcd(positive long value1, positive long value2) {
    assert(value1 > 0 && value2 > 0);

    if (value1 < value2) {
        swap(value1, value2, long);
    };

    while (value2 > 0) {
        value1 %= value2;
        swap(value1, value2, long);
    };

    return value1;
}

status_t resample(const float * restrict input_samples, positive long input_sample_count, float * restrict output_samples, positive long output_sample_count, positive long input_sample_rate, positive long output_sample_rate, float fade_bandwidth) {
    assert(input_sample_rate > 0 && output_sample_rate > 0);

    float *prototype_filter = NULL;
    float *_polyphase_filterbank = NULL;

    {
        char message[256];
        sprintf(message, "Resampling %ld samples @ %ldHz to %ld samples @ %ldHz", input_sample_count, input_sample_rate, output_sample_count, output_sample_rate);
        error_push(message);
    };

    positive long sr_gcd = gcd(input_sample_rate, output_sample_rate);
    positive long interpolation_factor = output_sample_rate / sr_gcd;
    positive long decimation_factor = input_sample_rate / sr_gcd;

    if (unlikely(interpolation_factor > INT_MAX / SP_VECTOR_BANDWIDTH)) {
        error_set("Resampling requires interpolation factor that would exceed the integer limit");
        goto fail;
    } else if (unlikely(decimation_factor > INT_MAX)) {
        error_set("Resampling requires decimation factor that would exceed the integer limit");
        goto fail;
    };

    float transition_length;
    if (output_sample_rate > input_sample_rate) {
        transition_length = (fade_bandwidth * 2.0f + (float)(output_sample_rate - input_sample_rate)) / (float)output_sample_rate;
    } else {
        transition_length = (fade_bandwidth * 2.0f) / (float)output_sample_rate;
    };

    positive int tap_count = get_lowpass_kaiser_fir_tap_count(RESAMPLE_STOPBAND_ATTENUATION, transition_length, interpolation_factor);
    positive int subfilter_size = tap_count / interpolation_factor;

    prototype_filter = xalloc(tap_count, sizeof(float));
    if (unlikely(prototype_filter == NULL)) {
        error_set("Failed to allocate prototype sample rate conversion filter");
        goto fail;
    };

    generate_window(prototype_filter + CEIL_HALF(tap_count), tap_count, WINDOW_KAISERBESSEL_RESAMPLE);
    for (positive int itr = 0; itr < CEIL_HALF(tap_count); itr++) {
        prototype_filter[itr] = prototype_filter[(tap_count - 1) - itr];
    };

    const float passband = 1.0f - (transition_length / (float)output_sample_rate);
    for (positive int itr = 0; itr < tap_count; itr++) {
        prototype_filter[itr] *= sinc((float)(itr - tap_count / 2) * passband);
    };

    float energy = 0.0f;
    for (positive int itr = 0; itr < tap_count; itr++) {
        energy += prototype_filter[itr];
    };

    for (positive int itr = 0; itr < tap_count; itr++) {
        prototype_filter[itr] *= (float)interpolation_factor / energy;
    };

    _polyphase_filterbank = mxalloc(2, sizeof(float), interpolation_factor, subfilter_size);
    float (*polyphase_filterbank)[subfilter_size] = (void*)_polyphase_filterbank;

    if (unlikely(polyphase_filterbank == NULL)) {
        error_set("Failed to allocate sample rate conversion polyphase filterbank");
        goto fail;
    };

    const float (*prototype_matrix)[interpolation_factor] = (void*)prototype_filter;
    for (positive int subfilter_index = 0; subfilter_index < interpolation_factor; subfilter_index++) {
        for (positive int point_index = 0; point_index < subfilter_size; point_index++) {
            polyphase_filterbank[subfilter_index][point_index] = prototype_matrix[point_index][subfilter_index];
        };
    };

    free(prototype_filter);
    prototype_filter = NULL;

    positive int subfilter_index = (tap_count / 2) % interpolation_factor;
    positive long input_sample_index = 0;

    for (positive long output_sample_index = 0; output_sample_index < output_sample_count; output_sample_index++) {
        float filter_value = 0.0f;
        for (positive int itr = 0; itr < subfilter_size; itr++) {
            long index = input_sample_index + itr - (tap_count / 2);
            if (likely(index >= 0 && index < input_sample_count)) {
                filter_value += input_samples[index] * polyphase_filterbank[subfilter_index][itr];
            };
        };

        output_samples[output_sample_index] = filter_value;

        subfilter_index -= decimation_factor;
        while (subfilter_index < 0) {
            subfilter_index += interpolation_factor;
            input_sample_index++;
        };
    };

    error_pop();
    free(_polyphase_filterbank);
    return SUCCESS;

    fail:
        free(prototype_filter);
        free(_polyphase_filterbank);
        return FAIL;
}
