#include "wavtool_fft.h"
#include "pffft.h"
#include <catch2/catch.hpp>

WavtoolFFT::WavtoolFFT() {
  for (auto sz : {64,   96,   128,   160,   192,   256,    384,   480,
                  512,  640,  768,   800,   1024,  2048,   2400,  4096,
                  8192, 9216, 16384, 32768, 65536, 131072, 262144}) {
    pffft_setups[sz] = std::make_pair(pffft_new_setup(sz, PFFFT_REAL),
                                      pffft_new_setup(sz, PFFFT_COMPLEX));
  }
}

WavtoolFFT::~WavtoolFFT() {
  for (const auto &x : pffft_setups) {
    pffft_destroy_setup(x.second.first);
    pffft_destroy_setup(x.second.second);
  }
}

std::vector<int> WavtoolFFT::get_supported_sizes() {
  std::vector<int> res;
  res.reserve(pffft_setups.size());
  for (const auto &x : pffft_setups) {
    res.push_back(x.first);
  }
  return res;
}

int WavtoolFFT::fft_ordered(int size, TransformType type,
                            TransformDirection direction, const float *input,
                            float *output) {
  auto it = pffft_setups.find(size);
  if (it == pffft_setups.end()) {
    return -1;
  }
  pffft_transform_ordered(
      type == TransformType::REAL ? it->second.first : it->second.second, input,
      output, nullptr,
      direction == TransformDirection::FORWARD ? PFFFT_FORWARD
                                               : PFFFT_BACKWARD);
  return 0;
}

int WavtoolFFT::fft_unordered(int size, TransformType type,
                              TransformDirection direction, const float *input,
                              float *output) {
  auto it = pffft_setups.find(size);
  if (it == pffft_setups.end()) {
    return -1;
  }
  pffft_transform(type == TransformType::REAL ? it->second.first
                                              : it->second.second,
                  input, output, nullptr,
                  direction == TransformDirection::FORWARD ? PFFFT_FORWARD
                                                           : PFFFT_BACKWARD);
  return 0;
}

int WavtoolFFT::fft_convolve_accumulate(int size, TransformType type,
                                        float scale, const float *in1,
                                        const float *in2, float *output) {
  auto it = pffft_setups.find(size);
  if (it == pffft_setups.end()) {
    return -1;
  }
  pffft_zconvolve_accumulate(type == TransformType::REAL ? it->second.first
                                                         : it->second.second,
                             in1, in2, output, scale);
  return 0;
}

int WavtoolFFT::fft_reorder(int size, TransformType type, const float *input,
                            float *output) {
  auto it = pffft_setups.find(size);
  if (it == pffft_setups.end()) {
    return -1;
  }
  pffft_zreorder(type == TransformType::REAL ? it->second.first
                                             : it->second.second,
                 input, output, PFFFT_FORWARD);
  return 0;
}

PFFFT_Setup *WavtoolFFT::get_pffft_setup(int size, TransformType type) {
  auto it = pffft_setups.find(size);
  if (it == pffft_setups.end()) {
    return nullptr;
  }
  return type == TransformType::REAL ? it->second.first : it->second.second;
}

WavtoolFFT wavtool_fft;

extern "C" {
int fft_get_supported_size_count() {
  auto x = wavtool_fft.get_supported_sizes();
  return x.size();
}

void fft_get_supported_sizes(int *output) {
  auto x = wavtool_fft.get_supported_sizes();
  for (int i = 0; i < x.size(); i++) {
    output[i] = x[i];
  }
}

int fft_ordered(int size, int type, int direction, float *input,
                float *output) {
  return wavtool_fft.fft_ordered(
      size, static_cast<WavtoolFFT::TransformType>(type),
      static_cast<WavtoolFFT::TransformDirection>(direction), input, output);
}

int fft_convolve_real(int size, float *in_fft, float *in_real, float *output) {
  auto in_res = wavtool_fft.fft_ordered(size, WavtoolFFT::TransformType::REAL,
                                        WavtoolFFT::TransformDirection::FORWARD,
                                        in_real, output);
  if (in_res != 0) {
    return in_res;
  }
  float scale = 1.0f / size;
  for (int i = 0; i < size; i += 2) {
    auto in_fft_re = in_fft[i];
    auto in_fft_im = in_fft[i + 1];
    auto in_real_re = output[i];
    auto in_real_im = output[i + 1];
    output[i] = (in_fft_re * in_real_re - in_fft_im * in_real_im) * scale;
    output[i + 1] = (in_fft_re * in_real_im + in_fft_im * in_real_re) * scale;
  }
  return wavtool_fft.fft_ordered(size, WavtoolFFT::TransformType::REAL,
                                 WavtoolFFT::TransformDirection::BACKWARD,
                                 output, output);
}

float *fft_alloc_aligned(int size) {
  return (float *)pffft_aligned_malloc(size * sizeof(float));
}

void fft_free_aligned(float *ptr) { pffft_aligned_free(ptr); }
}

TEST_CASE("fft tests", "[fft]") {
  WavtoolFFT f;
  for (auto size : f.get_supported_sizes()) {
    // complex
    {
      std::vector<float> time(2 * size);
      std::vector<float> freq(2 * size);
      for (int i = 0; i < size; i++) {
        time[i * 2] = std::sin(i);
        time[i * 2 + 1] = 0.f;
      }
      auto res = f.fft_ordered(size, WavtoolFFT::TransformType::COMPLEX,
                               WavtoolFFT::TransformDirection::FORWARD,
                               time.data(), freq.data());
      REQUIRE(res == 0);
      for (int i = 0; i < 2 * size; i++) {
        time[i] = 0.f;
      }
      res = f.fft_ordered(size, WavtoolFFT::TransformType::COMPLEX,
                          WavtoolFFT::TransformDirection::BACKWARD, freq.data(),
                          time.data());
      REQUIRE(res == 0);
      for (int i = 0; i < size; ++i) {
        const auto delta = std::fabs(time[i * 2] / (float)size - std::sin(i));
        if (delta > 1e-5f) {
          FAIL("FFT error");
        }
      }
    }

    // real
    {
      std::vector<float> time(size);
      std::vector<float> freq(size);
      for (int i = 0; i < size; i++) {
        time[i] = std::sin(i);
      }
      auto res = f.fft_ordered(size, WavtoolFFT::TransformType::REAL,
                               WavtoolFFT::TransformDirection::FORWARD,
                               time.data(), freq.data());
      REQUIRE(res == 0);
      for (int i = 0; i < size; i++) {
        time[i] = 0.f;
      }
      res = f.fft_ordered(size, WavtoolFFT::TransformType::REAL,
                          WavtoolFFT::TransformDirection::BACKWARD, freq.data(),
                          time.data());
      REQUIRE(res == 0);
      for (int i = 0; i < size; ++i) {
        const auto delta = std::fabs(time[i] / (float)size - std::sin(i));
        if (delta > 1e-5f) {
          FAIL("FFT error");
        }
      }
    }
  }
}