#include "simd.h"
#include <catch2/catch.hpp>
#include <cmath>
#include <initializer_list>
#include <tuple>
#include <type_traits>
#include <wasm_simd128.h>

template <class T, class... Ts>
struct are_same : std::conjunction<std::is_same<T, Ts>...> {};

template <typename... FloatInputType>
v128_t load_up_to(const int count, const int remaining,
                  const float *const input) {
  auto v = wasm_f32x4_splat(input[count - 1]);
  switch (remaining) {
  case 3:
    v = wasm_f32x4_replace_lane(v, 2, input[count - 3]);
  case 2:
    v = wasm_f32x4_replace_lane(v, 1, input[count - 2]);
  default:
  }
  return v;
}

template <typename F, typename... FloatInputType>
void do_array_simd(F op, int count, float *const output,
                   const FloatInputType *const... inputs) {
  static_assert(are_same<float, FloatInputType...>::value,
                "Inputs must be const float * const");
  const auto full = count & ~0x3;
  const auto remaining = count & 0x3;
#pragma clang loop vectorize(disable)
  for (int i = 0; i < full; i += 4) {
    v128_t result = op(wasm_v128_load(inputs + i)...);
    wasm_v128_store(output + i, result);
  }
  if (remaining > 0) {
    auto vs = std::make_tuple(load_up_to(count, remaining, inputs)...);
    v128_t res = std::apply(op, vs);
    switch (remaining) {
    case 3:
      output[count - 3] = wasm_f32x4_extract_lane(res, 2);
    case 2:
      output[count - 2] = wasm_f32x4_extract_lane(res, 1);
    case 1:
      output[count - 1] = wasm_f32x4_extract_lane(res, 0);
    }
  }
}

template <typename F, typename... FloatInputType>
void do_array_simd_readonly(F op, int count,
                            const FloatInputType *const... inputs) {
  static_assert(are_same<float, FloatInputType...>::value,
                "Inputs must be const float * const");
  const auto full = count & ~0x3;
  const auto remaining = count & 0x3;
#pragma clang loop vectorize(disable)
  for (int i = 0; i < full; i += 4) {
    op(wasm_v128_load(inputs + i)...);
  }
  if (remaining > 0) {
    auto vs = std::make_tuple(load_up_to(count, remaining, inputs)...);
    std::apply(op, vs);
  }
}

static inline v128_t vfastpow2(const v128_t p) {
  v128_t offset = wasm_v128_and(wasm_f32x4_lt(p, wasm_f32x4_const_splat(0.f)),
                                wasm_f32x4_const_splat(1.f));
  v128_t clipp = wasm_f32x4_max(p, wasm_f32x4_const_splat(-126.f));
  v128_t z = wasm_f32x4_sub(clipp, wasm_f32x4_trunc(clipp));
  z = wasm_f32x4_add(z, offset);

  v128_t res_f = wasm_f32x4_mul(
      wasm_f32x4_const_splat(1 << 23),
      wasm_f32x4_sub(
          wasm_f32x4_add(
              wasm_f32x4_add(clipp, wasm_f32x4_const_splat(121.2740575f)),
              wasm_f32x4_div(
                  wasm_f32x4_const_splat(27.7280233f),
                  wasm_f32x4_sub(wasm_f32x4_const_splat(4.84252568f), z))),
          wasm_f32x4_mul(wasm_f32x4_const_splat(1.49012907f), z)));

  return wasm_u32x4_trunc_sat_f32x4(res_f);
}

static inline v128_t vfastlog2(const v128_t vx) {
  const v128_t mx =
      wasm_v128_or(wasm_v128_and(vx, wasm_i32x4_const_splat(0x007FFFFF)),
                   wasm_i32x4_const_splat(0x3f000000));
  const v128_t y =
      wasm_f32x4_mul(wasm_f32x4_convert_i32x4(vx),
                     wasm_f32x4_const_splat(1.1920928955078125e-7f));

  const v128_t c_124_22551499 = wasm_f32x4_const_splat(124.22551499f);
  const v128_t c_1_498030302 = wasm_f32x4_const_splat(1.498030302f);
  const v128_t c_1_725877999 = wasm_f32x4_const_splat(1.72587999f);
  const v128_t c_0_3520087068 = wasm_f32x4_const_splat(0.3520887068f);

  return wasm_f32x4_sub(
      wasm_f32x4_sub(wasm_f32x4_sub(y, c_124_22551499),
                     wasm_f32x4_mul(c_1_498030302, mx)),
      wasm_f32x4_div(c_1_725877999, wasm_f32x4_add(c_0_3520087068, mx)));
}

static inline v128_t vfastsin_0topi(const v128_t vx) {
  const v128_t fouroverpi = wasm_f32x4_const_splat(1.2732395447351627f);
  const v128_t fouroverpisq = wasm_f32x4_const_splat(0.40528473456935109f);
  const v128_t q = wasm_f32x4_const_splat(0.78444488374548933f);
  const v128_t p = wasm_f32x4_const_splat(0.20363937680730309f);
  const v128_t r = wasm_f32x4_const_splat(0.015124940802184233f);
  const v128_t s = wasm_f32x4_const_splat(-0.0032225901625579573f);

  const v128_t qpprox =
      wasm_f32x4_sub(wasm_f32x4_mul(fouroverpi, vx),
                     wasm_f32x4_mul(fouroverpisq, wasm_f32x4_mul(vx, vx)));
  const v128_t qpproxsq = wasm_f32x4_mul(qpprox, qpprox);

  return wasm_f32x4_add(
      wasm_f32x4_mul(q, qpprox),
      wasm_f32x4_mul(
          qpproxsq,
          wasm_f32x4_add(
              p,
              wasm_f32x4_mul(qpproxsq,
                             wasm_f32x4_add(r, wasm_f32x4_mul(qpproxsq, s))))));
}

static inline v128_t vfastcos_0topi(const v128_t vx) {
  // TODO this could be optimized
  const v128_t halfpi = wasm_f32x4_const_splat(1.5707963267948966f);
  const v128_t halfpiorless = wasm_f32x4_le(vx, halfpi);
  const v128_t lag = wasm_v128_and(wasm_f32x4_add(vx, halfpi), halfpiorless);
  const v128_t lead =
      wasm_v128_andnot(wasm_f32x4_sub(vx, halfpi), halfpiorless);
  const v128_t sined = vfastsin_0topi(wasm_v128_or(lag, lead));
  return wasm_f32x4_mul(
      sined, wasm_v128_or(
                 wasm_v128_and(wasm_f32x4_const_splat(1.f), halfpiorless),
                 wasm_v128_andnot(wasm_f32x4_const_splat(-1.f), halfpiorless)));
}

// pow(2, input)
void __attribute__((noinline)) vfastpow2_array(const float *input,
                                               float *output, int count) {
  do_array_simd(vfastpow2, count, output, input);
}

// log2(input)
void __attribute__((noinline)) vfastlog2_array(const float *input,
                                               float *output, int count) {
  do_array_simd(vfastlog2, count, output, input);
}

// pow(a, b)
void __attribute__((noinline)) vfastpow_array(const float *a, const float *b,
                                              float *output, int count) {
  do_array_simd(
      [](auto va, auto vb) {
        return vfastpow2(wasm_f32x4_mul(vfastlog2(va), vb));
      },
      count, output, a, b);
}

// pow(10, input)
void __attribute__((noinline)) vfastpow10_array(const float *input,
                                                float *output, int count) {
  const v128_t log10 = wasm_f32x4_const_splat(3.32192809489f); // log2(10)
  do_array_simd(
      [log10](auto va) { return vfastpow2(wasm_f32x4_mul(log10, va)); }, count,
      output, input);
}

// sqrt(input)
void __attribute__((noinline)) vfastsqrt_array(const float *input,
                                               float *output, int count) {
  do_array_simd([](auto va) { return wasm_f32x4_sqrt(va); }, count, output,
                input);
}

// sin(input) for input in [0, pi]
void __attribute__((noinline)) vfastsin_0topi_array(const float *input,
                                                    float *output, int count) {
  do_array_simd(vfastsin_0topi, count, output, input);
}

// cos(input) for input in [0, pi]
void __attribute__((noinline)) vfastcos_0topi_array(const float *input,
                                                    float *output, int count) {
  do_array_simd(vfastcos_0topi, count, output, input);
}

// fmod(input, 1) NOTE: output has same sign as input
void __attribute__((noinline)) vfastfmod1_array(const float *input,
                                                float *output, int count) {

  do_array_simd([](auto v) { return wasm_f32x4_sub(v, wasm_f32x4_trunc(v)); },
                count, output, input);
}

// mod(input, 1) positive output (0 - \epsilon -> 1 - \epsilon)
void __attribute__((noinline)) vfastwrapmod1_array(const float *input,
                                                   float *output, int count) {

  do_array_simd(
      [](auto v) {
        const auto firstmod =
            wasm_f32x4_sub(v, wasm_f32x4_trunc(v)); // -1 < firstmod < 1
        const auto translated = wasm_f32x4_add(
            firstmod, wasm_f32x4_const_splat(1.f)); // 0 < translated < 2
        return wasm_f32x4_sub(translated, wasm_f32x4_trunc(translated));
      },
      count, output, input);
}

// clamp(input, min, max)
void __attribute__((noinline)) vfastclamp_array(const float *input, float min,
                                                float max, float *output,
                                                int count) {
  v128_t vmin = wasm_f32x4_splat(min);
  v128_t vmax = wasm_f32x4_splat(max);

  do_array_simd(
      [=](auto v) { return wasm_f32x4_pmin(vmax, wasm_f32x4_pmax(vmin, v)); },
      count, output, input);
}

// clamp(input, min, max) and replace NaNs with 0s
void __attribute__((noinline)) vfastclamp_nanzero_array(const float *input,
                                                        float min, float max,
                                                        float *output,
                                                        int count) {
  v128_t vmin = wasm_f32x4_splat(min);
  v128_t vmax = wasm_f32x4_splat(max);

  do_array_simd(
      [=](auto v) {
        v128_t result = wasm_f32x4_pmin(vmax, wasm_f32x4_pmax(vmin, v));
        v128_t nanMask = wasm_f32x4_eq(v, v);
        return wasm_v128_and(result, nanMask);
      },
      count, output, input);
}

// replace NaNs with 0s
void __attribute__((noinline)) vfastnanzero_array(const float *input,
                                                  float *output, int count) {
  do_array_simd(
      [=](auto v) {
        v128_t nanMask = wasm_f32x4_eq(v, v);
        return wasm_v128_and(v, nanMask);
      },
      count, output, input);
}

// abs(input)
void __attribute__((noinline)) vfastabs_array(const float *input, float *output,
                                              int count) {
  do_array_simd(wasm_f32x4_abs, count, output, input);
}

// min(a, b)
void __attribute__((noinline)) vfastmin_array(const float *a, const float *b,
                                              float *output, int count) {
  do_array_simd(wasm_f32x4_pmin, count, output, a, b);
}

// max(a, b)
void __attribute__((noinline)) vfastmax_array(const float *a, const float *b,
                                              float *output, int count) {
  do_array_simd(wasm_f32x4_pmax, count, output, a, b);
}

// max(*input)
float __attribute__((noinline)) vfastmax_elem_array(const float *input,
                                                    int count) {
  v128_t accum = wasm_f32x4_splat(-INFINITY);
  do_array_simd_readonly([&](auto v) { accum = wasm_f32x4_pmax(accum, v); },
                         count, input);
  return std::max(wasm_f32x4_extract_lane(accum, 0),
                  std::max(wasm_f32x4_extract_lane(accum, 1),
                           std::max(wasm_f32x4_extract_lane(accum, 2),
                                    wasm_f32x4_extract_lane(accum, 3))));
}

// min(*input)
float __attribute__((noinline)) vfastmin_elem_array(const float *input,
                                                    int count) {
  v128_t accum = wasm_f32x4_splat(INFINITY);
  do_array_simd_readonly([&](auto v) { accum = wasm_f32x4_pmin(accum, v); },
                         count, input);
  return std::min(wasm_f32x4_extract_lane(accum, 0),
                  std::min(wasm_f32x4_extract_lane(accum, 1),
                           std::min(wasm_f32x4_extract_lane(accum, 2),
                                    wasm_f32x4_extract_lane(accum, 3))));
}

// max(*(abs(x) for x in input))
float __attribute__((noinline)) vfastmax_abs_elem_array(const float *input,
                                                        int count) {
  v128_t accum = wasm_f32x4_splat(-INFINITY);
  do_array_simd_readonly(
      [&](auto v) { accum = wasm_f32x4_pmax(accum, wasm_f32x4_abs(v)); }, count,
      input);
  return std::max(wasm_f32x4_extract_lane(accum, 0),
                  std::max(wasm_f32x4_extract_lane(accum, 1),
                           std::max(wasm_f32x4_extract_lane(accum, 2),
                                    wasm_f32x4_extract_lane(accum, 3))));
}

// min(*(abs(x) for x in input))
float __attribute__((noinline)) vfastmin_abs_elem_array(const float *input,
                                                        int count) {
  v128_t accum = wasm_f32x4_splat(INFINITY);
  do_array_simd_readonly(
      [&](auto v) { accum = wasm_f32x4_pmin(accum, wasm_f32x4_abs(v)); }, count,
      input);
  return std::min(wasm_f32x4_extract_lane(accum, 0),
                  std::min(wasm_f32x4_extract_lane(accum, 1),
                           std::min(wasm_f32x4_extract_lane(accum, 2),
                                    wasm_f32x4_extract_lane(accum, 3))));
}

void __attribute__((noinline))
vfast_curve_fade_array(const float *input, float gain, float startTime,
                       float endTime, float fadeInDuration,
                       float fadeOutDuration, float fadeInExponent,
                       float fadeOutExponent, float *output, int count) {
  v128_t vstart = wasm_f32x4_splat(startTime);
  v128_t vend = wasm_f32x4_splat(endTime);
  v128_t vgain = wasm_f32x4_splat(gain);
  v128_t vfadeInDuration = wasm_f32x4_splat(fadeInDuration);
  v128_t vfadeOutDuration = wasm_f32x4_splat(fadeOutDuration);
  v128_t vfadeInExponent = wasm_f32x4_splat(fadeInExponent);
  v128_t vfadeOutExponent = wasm_f32x4_splat(fadeOutExponent);

  v128_t vzero = wasm_f32x4_const_splat(0.f);
  v128_t vone = wasm_f32x4_const_splat(1.f);

  do_array_simd(
      [=](auto v) {
        v128_t vtimeIntoClip = wasm_f32x4_sub(v, vstart);
        v128_t vtimeLeftInClip = wasm_f32x4_sub(vend, v);
        v128_t vfadeInBase = wasm_f32x4_pmin(
            vone, wasm_f32x4_pmax(
                      vzero, wasm_f32x4_div(vtimeIntoClip, vfadeInDuration)));
        v128_t vfadeOutBase = wasm_f32x4_pmin(
            vone, wasm_f32x4_pmax(vzero, wasm_f32x4_div(vtimeLeftInClip,
                                                        vfadeOutDuration)));
        v128_t vfadeInGain =
            vfastpow2(wasm_f32x4_mul(vfastlog2(vfadeInBase), vfadeInExponent));
        v128_t vfadeOutGain = vfastpow2(
            wasm_f32x4_mul(vfastlog2(vfadeOutBase), vfadeOutExponent));
        v128_t vcombinedGain =
            wasm_f32x4_mul(vgain, wasm_f32x4_mul(vfadeInGain, vfadeOutGain));
        v128_t nanMask = wasm_f32x4_eq(vcombinedGain, vcombinedGain);
        return wasm_v128_and(vcombinedGain, nanMask);
      },
      count, output, input);
}

float __attribute__((noinline)) vfast_ema_squared_array(const float *input,
                                                        float alpha,
                                                        float initialValue,
                                                        int count) {
  const auto oneMinusAlpha = 1.f - alpha;
  const auto full = count & ~0x3;

  float currentValue = initialValue;

  if (full > 0) {
    const v128_t alphaVec = wasm_f32x4_splat(alpha);
    const v128_t oneMinusAlphaVec = wasm_f32x4_splat(oneMinusAlpha);

    const auto alpha2 = alpha * alpha;
    const auto alpha3 = alpha2 * alpha;
    const auto alpha4 = alpha3 * alpha;

    const v128_t alphaPowers = wasm_f32x4_make(alpha3, alpha2, alpha, 1.0f);

#pragma clang loop vectorize(disable)
    for (int i = 0; i < full; i += 4) {
      v128_t samples_v = wasm_v128_load(input + i);
      v128_t squared_v = wasm_f32x4_mul(samples_v, samples_v);

      v128_t contributions = wasm_f32x4_mul(
          wasm_f32x4_mul(oneMinusAlphaVec, squared_v), alphaPowers);

      float contrib_sum = wasm_f32x4_extract_lane(contributions, 0) +
                          wasm_f32x4_extract_lane(contributions, 1) +
                          wasm_f32x4_extract_lane(contributions, 2) +
                          wasm_f32x4_extract_lane(contributions, 3);

      currentValue = alpha4 * currentValue + contrib_sum;
    }
  }

  for (int i = full; i < count; i++) {
    const auto sample = input[i];
    const auto squared = sample * sample;
    currentValue = alpha * currentValue + oneMinusAlpha * squared;
  }

  return currentValue;
}

TEST_CASE("do_array_simd weird lengths", "[simd]") {
  constexpr const int N = 16;
  float x0[N];
  float x[N];
  float y[N];
  for (int i = 0; i < N; i++) {
    memset(y, 0, sizeof(float) * N);
    for (int j = 0; j < i; j++) {
      x0[j] = j + 1;
    }
    for (int j = i; j < N; j++) {
      x0[j] = 1;
    }
    memcpy(x, x0, sizeof(float) * N);
    vfastpow2_array(x, y, i);
    REQUIRE(std::equal(x0, x0 + N, x));
    REQUIRE(std::all_of(y + i, y + N, [](const float &v) { return v == 0.f; }));
    for (int j = 0; j < i; j++) {
      REQUIRE_THAT(y[j],
                   Catch::Matchers::WithinRel(std::pow(2.f, x[j]), 1e-5f));
    }
  }
}

TEST_CASE("wasm simd ops", "[simd]") {
  auto test_op_range = [](auto simd_op, std::function<float(float)> &&ref_op,
                          float lo, float hi, float step, float eps = 1e-5f) {
    constexpr const int N = 1024;
    float x[N];
    float y[N];
    int i = 0;

    auto testblock = [&]() {
      simd_op(x, y, i);
      for (int j = 0; j < i; j++) {
        REQUIRE_THAT(y[j], Catch::Matchers::WithinRel(ref_op(x[j]), eps));
      }
    };

    while (lo <= hi) {
      x[i++] = lo;
      lo += step;

      if (i == N) {
        testblock();
        i = 0;
      }
    }
    if (i > 0) {
      testblock();
    }
  };

  auto test_op_list = [](auto simd_op, std::function<float(float)> &&ref_op,
                         std::initializer_list<float> xs, float eps = 1e-5) {
    float x[xs.size()];
    std::copy(xs.begin(), xs.end(), x);
    float y[xs.size()];
    simd_op(x, y, xs.size());
    for (int i = 0; i < xs.size(); ++i) {
      REQUIRE_THAT(y[i], Catch::Matchers::WithinRel(ref_op(x[i]), eps));
    }
  };

  auto test_op_list_2 = [](auto simd_op,
                           std::function<float(float, float)> &&ref_op,
                           std::initializer_list<std::pair<float, float>> xs,
                           float eps = 1e-5) {
    float x1[xs.size()];
    float x2[xs.size()];
    {
      int i = 0;
      for (const auto [x1e, x2e] : xs) {
        x1[i] = x1e;
        x2[i++] = x2e;
      }
    }
    float y[xs.size()];
    simd_op(x1, x2, y, xs.size());
    for (int i = 0; i < xs.size(); ++i) {
      REQUIRE_THAT(y[i], Catch::Matchers::WithinRel(ref_op(x1[i], x2[i]), eps));
    }
  };

  auto test_elem_op = [](auto simd_op, auto ref_op,
                         std::initializer_list<float> xs) {
    float x[xs.size()];
    std::copy(xs.begin(), xs.end(), x);
    float result = simd_op(x, xs.size());
    float expected = ref_op(xs.begin(), xs.end());
    REQUIRE_THAT(result, Catch::Matchers::WithinRel(expected, 1e-6f));
  };

  auto abs_comparator = [](float a, float b) {
    return std::abs(a) < std::abs(b);
  };

  auto ref_min_abs = [&](auto begin, auto end) {
    return std::abs(*std::min_element(begin, end, abs_comparator));
  };

  auto ref_max_abs = [&](auto begin, auto end) {
    return std::abs(*std::max_element(begin, end, abs_comparator));
  };

  auto ref_curve_fade = [](float time, float gain, float startTime,
                           float endTime, float fadeInDuration,
                           float fadeOutDuration, float fadeInExponent,
                           float fadeOutExponent) -> float {
    float timeIntoClip = time - startTime;
    float timeLeftInClip = endTime - time;

    float fadeInBase = std::clamp(timeIntoClip / fadeInDuration, 0.0f, 1.0f);
    float fadeOutBase =
        std::clamp(timeLeftInClip / fadeOutDuration, 0.0f, 1.0f);

    float fadeInGain = std::pow(fadeInBase, fadeInExponent);
    float fadeOutGain = std::pow(fadeOutBase, fadeOutExponent);

    float combinedGain = gain * fadeInGain * fadeOutGain;

    return std::isnan(combinedGain) ? 0.0f : combinedGain;
  };

  auto test_curve_fade =
      [&](std::initializer_list<float> times, float gain, float startTime,
          float endTime, float fadeInDuration, float fadeOutDuration,
          float fadeInExponent, float fadeOutExponent, float eps = 1e-4f) {
        float input[times.size()];
        float output[times.size()];
        std::copy(times.begin(), times.end(), input);

        vfast_curve_fade_array(input, gain, startTime, endTime, fadeInDuration,
                               fadeOutDuration, fadeInExponent, fadeOutExponent,
                               output, times.size());

        int i = 0;
        for (float time : times) {
          float expected =
              ref_curve_fade(time, gain, startTime, endTime, fadeInDuration,
                             fadeOutDuration, fadeInExponent, fadeOutExponent);
          REQUIRE_THAT(output[i++], Catch::Matchers::WithinAbs(expected, eps));
        }
      };

  auto ref_ema_squared = [](const float *input, float alpha, float initialValue,
                            int count) -> float {
    float currentValue = initialValue;
    float oneMinusAlpha = 1.0f - alpha;

    for (int i = 0; i < count; i++) {
      float sample = input[i];
      float squared = sample * sample;
      currentValue = alpha * currentValue + oneMinusAlpha * squared;
    }

    return currentValue;
  };

  auto test_ema_squared = [&](std::initializer_list<float> samples, float alpha,
                              float initialValue, float eps = 1e-5f) {
    float input[samples.size()];
    std::copy(samples.begin(), samples.end(), input);

    float result =
        vfast_ema_squared_array(input, alpha, initialValue, samples.size());
    float expected =
        ref_ema_squared(input, alpha, initialValue, samples.size());

    REQUIRE_THAT(result, Catch::Matchers::WithinRel(expected, eps));
  };

  SECTION("pow") {
    test_op_range(
        vfastpow2_array, [](float x) { return std::pow(2.f, x); }, -15.f, +15.f,
        0.01f, 1e-4);
    test_op_list(
        vfastpow2_array, [](float x) { return std::pow(2.f, x); },
        {0.f, -100.f, 100.f}, 1e-4);
    test_op_list(
        vfastpow10_array, [](float x) { return std::pow(10.f, x); },
        {-3.f, -2.f, -0.33f, -1.f, 0.f, 0.33f, 1.f, 2.f, 3.f}, 1e-4);
    test_op_list(
        vfastsqrt_array, [](float x) { return std::sqrtf(x); },
        {0.f, 0.33f, 1.f, 2.f, 3.f, 1234.f}, 1e-4);
    test_op_list_2(
        vfastpow_array, [](float a, float b) { return std::pow(a, b); },
        {{1, -99}, {1, 99}, {1, 0}, {2, 0}, {2, 2}, {3, 3}, {4, 4}, {5, 5}},
        5e-4);
  }

  SECTION("log") {
    test_op_list(
        vfastlog2_array, [](float x) { return std::log2f(x); },
        {1e-20f, 3e-20f, 1e-10f, 1e-3f, 1e-1f, 3, 10, 30, 1e5, 3e5, 1e10, 3e10},
        1e-4);
  }

  SECTION("trig") {
    test_op_range(
        vfastsin_0topi_array, [](float x) { return std::sinf(x); }, 0.f, M_PI,
        0.01f, 5e-3);
    test_op_range(
        vfastcos_0topi_array, [](float x) { return std::cosf(x); }, 0.f, M_PI,
        0.01f, 5e-3);
  }

  SECTION("fmod") {
    test_op_list(vfastfmod1_array, [](float x) { return std::fmod(x, 1.f); },
                 {-1e30f, -2.f, -1.5f, -1.33f, -1.f, -0.33f, -1e-30f, 0.f,
                  1e-30f, 0.33f, 0.999f, 1.f, 1.0000001f, 1.5f, 12345.f,
                  1e30f});

    test_op_list(
        vfastwrapmod1_array,
        [](float x) { return std::fmod(std::fmod(x, 1.f) + 1.f, 1.f); },
        {-999.f, -2.5f, -2.1f, -1.f, -0.9f, -0.1f, 0.f, 0.1f, 1.0f, 1.1f, 1.5f,
         1.9f, 2.0f, 2.3f, 999.f});
  }

  SECTION("abs") {
    test_op_list(vfastabs_array, [](float x) { return std::fabs(x); },
                 {-100.f, -1.f, 0.f, 1.f, 100.f});
  }

  SECTION("clamp") {
    test_op_list([](const float *x, float *y,
                    int count) { vfastclamp_array(x, -1.f, 1.f, y, count); },
                 [](float x) { return std::clamp(x, -1.f, 1.f); },
                 {-INFINITY, -2, -1, -0.5, 0, 0.5, 1, 2, INFINITY});
    test_op_list(
        [](const float *x, float *y, int count) {
          vfastclamp_nanzero_array(x, -1.f, 1.f, y, count);
        },
        [](float x) { return std::isnan(x) ? 0.f : std::clamp(x, -1.f, 1.f); },
        {-INFINITY, -2, -1, -0.5, 0, 0.5, 1, 2, INFINITY, 0.f / 0.f});
  }

  SECTION("extrema") {
    test_op_list_2(vfastmin_array,
                   [](float a, float b) { return std::min(a, b); },
                   {{-INFINITY, INFINITY},
                    {INFINITY, -INFINITY},
                    {-1, -2},
                    {-2, -1},
                    {0, 0},
                    {1, 2},
                    {2, 1}});
    test_op_list_2(vfastmax_array,
                   [](float a, float b) { return std::max(a, b); },
                   {{-INFINITY, INFINITY},
                    {INFINITY, -INFINITY},
                    {-1, -2},
                    {-2, -1},
                    {0, 0},
                    {1, 2},
                    {2, 1}});
  }

  SECTION("element extrema") {
    test_elem_op(
        vfastmin_elem_array,
        [](auto begin, auto end) { return *std::min_element(begin, end); },
        {1.0f, 2.0f, 3.0f, 4.0f});
    test_elem_op(
        vfastmax_elem_array,
        [](auto begin, auto end) { return *std::max_element(begin, end); },
        {1.0f, 2.0f, 3.0f, 4.0f});

    test_elem_op(
        vfastmin_elem_array,
        [](auto begin, auto end) { return *std::min_element(begin, end); },
        {-3.0f, -1.0f, -5.0f, -2.0f});
    test_elem_op(
        vfastmax_elem_array,
        [](auto begin, auto end) { return *std::max_element(begin, end); },
        {-3.0f, -1.0f, -5.0f, -2.0f});

    test_elem_op(
        vfastmin_elem_array,
        [](auto begin, auto end) { return *std::min_element(begin, end); },
        {-2.0f, 1.0f, -5.0f, 3.0f});
    test_elem_op(
        vfastmax_elem_array,
        [](auto begin, auto end) { return *std::max_element(begin, end); },
        {-2.0f, 1.0f, -5.0f, 3.0f});

    test_elem_op(
        vfastmin_elem_array,
        [](auto begin, auto end) { return *std::min_element(begin, end); },
        {-INFINITY, 1.0f, 2.0f, 3.0f});
    test_elem_op(
        vfastmax_elem_array,
        [](auto begin, auto end) { return *std::max_element(begin, end); },
        {1.0f, 2.0f, 3.0f, INFINITY});

    test_elem_op(
        vfastmin_elem_array,
        [](auto begin, auto end) { return *std::min_element(begin, end); },
        {42.0f});
    test_elem_op(
        vfastmax_elem_array,
        [](auto begin, auto end) { return *std::max_element(begin, end); },
        {42.0f});

    for (int len = 1; len <= 17; len++) {
      float x[17];
      for (int i = 0; i < len; i++) {
        x[i] = sin(i * 0.5f) * 10.0f;
      }

      float expected_min = *std::min_element(x, x + len);
      float expected_max = *std::max_element(x, x + len);

      REQUIRE_THAT(vfastmin_elem_array(x, len),
                   Catch::Matchers::WithinRel(expected_min, 1e-6f));
      REQUIRE_THAT(vfastmax_elem_array(x, len),
                   Catch::Matchers::WithinRel(expected_max, 1e-6f));
    }
  }

  SECTION("absolute element extrema") {
    test_elem_op(vfastmin_abs_elem_array, ref_min_abs,
                 {1.0f, -2.0f, 3.0f, -0.5f});
    test_elem_op(vfastmax_abs_elem_array, ref_max_abs,
                 {1.0f, -2.0f, 3.0f, -0.5f});

    test_elem_op(vfastmin_abs_elem_array, ref_min_abs,
                 {-3.0f, -1.0f, -5.0f, -2.0f});
    test_elem_op(vfastmax_abs_elem_array, ref_max_abs,
                 {-3.0f, -1.0f, -5.0f, -2.0f});

    test_elem_op(vfastmin_abs_elem_array, ref_min_abs,
                 {0.0f, -2.0f, 3.0f, -1.0f});
    test_elem_op(vfastmax_abs_elem_array, ref_max_abs,
                 {0.0f, -2.0f, 3.0f, -1.0f});

    test_elem_op(vfastmin_abs_elem_array, ref_min_abs, {-42.0f});
    test_elem_op(vfastmax_abs_elem_array, ref_max_abs, {-42.0f});

    for (int len = 1; len <= 17; len++) {
      float x[17];
      for (int i = 0; i < len; i++) {
        x[i] = sin(i * 0.3f + 1.0f) * 10.0f;
      }

      float expected_min_abs = ref_min_abs(x, x + len);
      float expected_max_abs = ref_max_abs(x, x + len);

      REQUIRE_THAT(vfastmin_abs_elem_array(x, len),
                   Catch::Matchers::WithinRel(expected_min_abs, 1e-6f));
      REQUIRE_THAT(vfastmax_abs_elem_array(x, len),
                   Catch::Matchers::WithinRel(expected_max_abs, 1e-6f));
    }
  }

  SECTION("curve fade") {
    test_curve_fade({0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 2.5f, 3.0f}, 1.0f, 0.5f,
                    2.5f, 0.5f, 0.5f, 1.0f, 1.0f);

    test_curve_fade({0.0f, 0.25f, 0.5f, 0.75f, 1.0f}, 2.0f, 0.0f, 1.0f, 0.5f,
                    0.5f, 2.0f, 0.5f);

    test_curve_fade({0.0f, 0.5f, 1.0f}, 1.5f, 0.0f, 1.0f, 0.0f, 0.0f, 1.0f,
                    1.0f);

    test_curve_fade({-1.0f, -0.5f, 2.0f, 3.0f}, 1.0f, 0.0f, 1.0f, 0.3f, 0.3f,
                    1.0f, 1.0f);

    test_curve_fade({0.0f, 0.001f, 0.999f, 1.0f}, 1.0f, 0.0f, 1.0f, 0.001f,
                    0.001f, 1.0f, 1.0f, 1e-3f);

    for (int len = 1; len <= 17; len++) {
      float times[17];
      float output[17];
      for (int i = 0; i < len; i++) {
        times[i] = i * 0.1f;
      }

      vfast_curve_fade_array(times, 1.0f, 0.2f, 1.4f, 0.3f, 0.3f, 1.5f, 1.5f,
                             output, len);

      for (int i = 0; i < len; i++) {
        float expected =
            ref_curve_fade(times[i], 1.0f, 0.2f, 1.4f, 0.3f, 0.3f, 1.5f, 1.5f);
        REQUIRE_THAT(output[i], Catch::Matchers::WithinAbs(expected, 1e-4f));
      }
    }
  }

  SECTION("ema squared") {
    test_ema_squared({1.0f, 2.0f, 3.0f, 4.0f}, 0.5f, 0.0f);

    test_ema_squared({1.0f, -1.0f, 2.0f, -2.0f}, 0.1f, 1.0f);
    test_ema_squared({1.0f, -1.0f, 2.0f, -2.0f}, 0.9f, 1.0f);
    test_ema_squared({1.0f, -1.0f, 2.0f, -2.0f}, 0.99f, 1.0f);

    test_ema_squared({2.0f, 3.0f}, 0.0f, 5.0f);
    test_ema_squared({2.0f, 3.0f}, 1.0f, 5.0f);

    test_ema_squared({-1.5f, 2.3f, -0.8f, 1.2f, -3.1f}, 0.3f, 2.0f);
    test_ema_squared({0.01f, -0.02f, 0.03f, -0.04f}, 0.7f, 0.001f);
    test_ema_squared({100.0f, -200.0f, 150.0f}, 0.4f, 10000.0f);
    test_ema_squared({5.0f}, 0.6f, 2.0f);

    {
      float dummy = 0.0f;
      float result = vfast_ema_squared_array(&dummy, 0.5f, 3.0f, 0);
      REQUIRE_THAT(result, Catch::Matchers::WithinRel(3.0f, 1e-6f));
    }

    for (int len = 1; len <= 17; len++) {
      float input[17];
      for (int i = 0; i < len; i++) {
        input[i] = sin(i * 0.7f) * 2.0f + cos(i * 0.3f);
      }

      float alpha = 0.25f + (len % 5) * 0.15f;
      float initialValue = len * 0.1f;

      float result = vfast_ema_squared_array(input, alpha, initialValue, len);
      float expected = ref_ema_squared(input, alpha, initialValue, len);

      REQUIRE_THAT(result, Catch::Matchers::WithinRel(expected, 1e-4f));
    }

    test_ema_squared({1.0f, 2.0f, 3.0f}, 1e-6f, 1000.0f, 1e-3f);
    test_ema_squared({1.0f, 2.0f, 3.0f}, 0.999999f, 5.0f, 1e-3f);

    {
      std::vector<float> long_input;
      long_input.reserve(100);
      for (int i = 0; i < 100; i++) {
        long_input.push_back(sin(i * 0.1f) * 3.0f);
      }

      float result = vfast_ema_squared_array(long_input.data(), 0.2f, 1.5f,
                                             long_input.size());
      float expected =
          ref_ema_squared(long_input.data(), 0.2f, 1.5f, long_input.size());

      REQUIRE_THAT(result, Catch::Matchers::WithinRel(expected, 1e-4f));
    }
  }
}