#include "delayline.h"
#include "catch2/catch.hpp"

DelayLine::DelayLine(int channelCount, int delayFrames)
    : writeIndex(0), delayBuffer(channelCount, delayFrames) {
  delayBuffer.fill(0.0);
}

DelayLine::~DelayLine() {}

void DelayLine::pushpull(BufferF32 buf) {
  const auto bufFrameCount = buf.getFrameCount();
  const auto bufChannelCount = buf.getChannelCount();
  const auto delayBufferFrameCount = delayBuffer.getFrameCount();

  assert(bufChannelCount == delayBuffer.getChannelCount());

  if (delayBufferFrameCount == 0 || bufFrameCount == 0) {
    return;
  }

  for (int channel = 0; channel < bufChannelCount; channel++) {
    const auto bufData = buf.getChannelData(channel);
    const auto delayBufferData = delayBuffer.getChannelData(channel);
    const auto writeIndexHere = writeIndex;

    // TODO vectorize this
    for (int i = 0; i < bufFrameCount; i++) {
      const auto origBufVal = bufData[i];
      const auto delayIndex = (writeIndexHere + i) % delayBufferFrameCount;
      bufData[i] = delayBufferData[delayIndex];
      delayBufferData[delayIndex] = origBufVal;
    }
  }

  writeIndex = (writeIndex + bufFrameCount) % delayBufferFrameCount;
}

void DelayLine::copyPlaybackState(const DelayLine &other) {
  assert(delayBuffer.getChannelCount() == other.delayBuffer.getChannelCount());
  assert(delayBuffer.getFrameCount() == other.delayBuffer.getFrameCount());
  delayBuffer.set(0, other.delayBuffer);
  writeIndex = other.writeIndex;
}

void DelayLine::zero() { delayBuffer.fill(0.0); }

void DelayLine::read(float *out, int channel, int readIndexOffset,
                     int frames) const {
  const auto delayBufferFrameCount = delayBuffer.getFrameCount();
  if (delayBufferFrameCount == 0) {
    return;
  }
  const auto channelData = delayBuffer.getChannelData(channel);
  const auto startIndex =
      (writeIndex + delayBufferFrameCount + readIndexOffset) %
      delayBufferFrameCount;
  for (int i = 0; i < frames; i++) {
    out[i] = channelData[(startIndex + i) % delayBufferFrameCount];
  }
}

float DelayLine::peakUpTo(int channel, int writeIndexOffset, int frames) {
  const auto delayBufferFrameCount = delayBuffer.getFrameCount();
  const auto channelData = delayBuffer.getChannelData(channel);
  const auto startIndex =
      (writeIndex + delayBufferFrameCount + writeIndexOffset) %
      delayBufferFrameCount;
  float peak = -INFINITY;
  for (int i = 0; i < frames; i++) {
    peak = std::max(
        peak, std::fabs(channelData[(startIndex + i) % delayBufferFrameCount]));
  }
  return peak;
}

TEST_CASE("DelayLine basic functionality", "[delayline]") {
  SECTION("Initial state - outputs zeros") {
    DelayLine delay(2, 4);
    auto input = BufferF32(2, 2);
    input.fill(1.0f);

    delay.pushpull(input);

    // First call should output zeros (initial delay buffer state)
    for (int ch = 0; ch < input.getChannelCount(); ch++) {
      const auto *data = input.getChannelData(ch);
      for (int i = 0; i < 2; i++) {
        REQUIRE(data[i] == 0.0f);
      }
    }
  }

  SECTION("Delay behavior - single channel") {
    DelayLine singleDelay(1, 3);
    auto buffer = BufferF32(1, 1);

    // Feed sequence: 1, 2, 3, 4, 5...
    // Expected output after 3 samples delay: 0, 0, 0, 1, 2, 3, 4...
    std::vector<float> expectedOutputs = {0.0f, 0.0f, 0.0f, 1.0f,
                                          2.0f, 3.0f, 4.0f};

    for (int i = 0; i < 7; i++) {
      buffer.getChannelData(0)[0] = static_cast<float>(i + 1);
      singleDelay.pushpull(buffer);

      float output = buffer.getChannelData(0)[0];
      REQUIRE(output == expectedOutputs[i]);
    }
  }

  SECTION("Multi-sample buffer processing") {
    DelayLine multiDelay(1, 4);
    auto buffer = BufferF32(1, 3);

    // First buffer: [1, 2, 3] -> should output [0, 0, 0]
    float *data = buffer.getChannelData(0);
    data[0] = 1.0f;
    data[1] = 2.0f;
    data[2] = 3.0f;

    multiDelay.pushpull(buffer);

    // expect [0, 1, 2, 3] in the line
    for (int i = 0; i < 4; i++) {
      float output[1] = {0.0f};
      multiDelay.read(output, 0, i, 1);
      REQUIRE((int)output[0] == i);
    }

    REQUIRE(data[0] == 0.0f);
    REQUIRE(data[1] == 0.0f);
    REQUIRE(data[2] == 0.0f);

    // Second buffer: [4, 5, 6] -> should output [0, 1, 2]
    data[0] = 4.0f;
    data[1] = 5.0f;
    data[2] = 6.0f;

    multiDelay.pushpull(buffer);

    REQUIRE(data[0] == 0.0f); // Still in initial delay
    REQUIRE(data[1] == 1.0f); // First sample from first buffer
    REQUIRE(data[2] == 2.0f); // Second sample from first buffer
  }

  SECTION("Multi-channel consistency") {
    DelayLine stereoDelay(2, 2);
    auto buffer = BufferF32(2, 1);

    // Feed different values to each channel
    for (int sample = 0; sample < 5; sample++) {
      buffer.getChannelData(0)[0] =
          static_cast<float>(sample * 10); // 0, 10, 20, 30, 40
      buffer.getChannelData(1)[0] =
          static_cast<float>(sample * 100); // 0, 100, 200, 300, 400

      stereoDelay.pushpull(buffer);

      if (sample < 2) {
        // During delay period
        REQUIRE(buffer.getChannelData(0)[0] == 0.0f);
        REQUIRE(buffer.getChannelData(1)[0] == 0.0f);
      } else {
        // After delay period
        REQUIRE(buffer.getChannelData(0)[0] ==
                static_cast<float>((sample - 2) * 10));
        REQUIRE(buffer.getChannelData(1)[0] ==
                static_cast<float>((sample - 2) * 100));
      }
    }
  }

  SECTION("Zero method resets delay buffer") {
    DelayLine testDelay(1, 2);
    auto buffer = BufferF32(1, 1);

    // Fill delay buffer with non-zero values
    buffer.getChannelData(0)[0] = 5.0f;
    testDelay.pushpull(buffer);
    buffer.getChannelData(0)[0] = 10.0f;
    testDelay.pushpull(buffer);

    // Zero the delay buffer
    testDelay.zero();

    // Next outputs should be zero
    buffer.getChannelData(0)[0] = 15.0f;
    testDelay.pushpull(buffer);
    REQUIRE(buffer.getChannelData(0)[0] == 0.0f);

    buffer.getChannelData(0)[0] = 20.0f;
    testDelay.pushpull(buffer);
    REQUIRE(buffer.getChannelData(0)[0] == 0.0f);
  }

  SECTION("copyPlaybackState preserves delay buffer and write position") {
    DelayLine source(2, 3);
    DelayLine dest(2, 3);

    auto buffer = BufferF32(2, 1);

    // Fill source delay with known pattern
    for (int i = 0; i < 4; i++) {
      buffer.getChannelData(0)[0] = static_cast<float>(i + 1);
      buffer.getChannelData(1)[0] = static_cast<float>((i + 1) * 10);
      source.pushpull(buffer);
    }

    // Copy state to destination
    dest.copyPlaybackState(source);

    // Both should now produce identical outputs
    auto sourceBuffer = BufferF32(2, 1);
    auto destBuffer = BufferF32(2, 1);

    for (int i = 0; i < 3; i++) {
      sourceBuffer.getChannelData(0)[0] = 99.0f;
      sourceBuffer.getChannelData(1)[0] = 999.0f;
      destBuffer.getChannelData(0)[0] = 99.0f;
      destBuffer.getChannelData(1)[0] = 999.0f;

      source.pushpull(sourceBuffer);
      dest.pushpull(destBuffer);

      REQUIRE(sourceBuffer.getChannelData(0)[0] ==
              destBuffer.getChannelData(0)[0]);
      REQUIRE(sourceBuffer.getChannelData(1)[0] ==
              destBuffer.getChannelData(1)[0]);
    }
  }
}

TEST_CASE("DelayLine edge cases", "[delayline]") {
  SECTION("Zero delay frames") {
    DelayLine zeroDelay(1, 0);
    auto buffer = BufferF32(1, 1);

    buffer.getChannelData(0)[0] = 42.0f;
    zeroDelay.pushpull(buffer);

    // With zero delay, output should equal input
    REQUIRE(buffer.getChannelData(0)[0] == 42.0f);
  }

  SECTION("Single frame delay") {
    DelayLine oneDelay(1, 1);
    auto buffer = BufferF32(1, 1);

    // First sample
    buffer.getChannelData(0)[0] = 1.0f;
    oneDelay.pushpull(buffer);
    REQUIRE(buffer.getChannelData(0)[0] == 0.0f); // Initial zero

    // Second sample
    buffer.getChannelData(0)[0] = 2.0f;
    oneDelay.pushpull(buffer);
    REQUIRE(buffer.getChannelData(0)[0] == 1.0f); // Previous input

    // Third sample
    buffer.getChannelData(0)[0] = 3.0f;
    oneDelay.pushpull(buffer);
    REQUIRE(buffer.getChannelData(0)[0] == 2.0f); // Previous input
  }

  SECTION("Large buffer processing") {
    const int bufferSize = 128;
    const int delayFrames = 64;

    DelayLine largeDelay(1, delayFrames);
    auto buffer = BufferF32(1, bufferSize);

    // Fill with ascending values
    float *data = buffer.getChannelData(0);
    for (int i = 0; i < bufferSize; i++) {
      data[i] = static_cast<float>(i);
    }

    largeDelay.pushpull(buffer);

    // First 64 samples should be zero (delay period)
    for (int i = 0; i < delayFrames; i++) {
      REQUIRE(data[i] == 0.0f);
    }

    // Remaining samples should be the first part of input
    for (int i = delayFrames; i < bufferSize; i++) {
      REQUIRE(data[i] == static_cast<float>(i - delayFrames));
    }
  }
}