#pragma once
#include "buffer.h"

class RandomAccessAudioReadable {
public:
  virtual ~RandomAccessAudioReadable() = default;
  virtual int getSampleRate() const = 0;
  virtual int getChannelCount() const = 0;
  virtual int getFrameCount() const = 0;

  // Reads the audio data starting at frameOffset into the output buffer.
  // Caller promises:
  // - output.getChannelCount() == getChannelCount()
  // - output.getFrameCount() <= getFrameCount() - frameOffset
  // - frameOffset >= 0
  virtual void read(int frameOffset, BufferF32 output) = 0;
  virtual void read(int frameOffset, std::shared_ptr<BufferF32> output) = 0;
  
  virtual std::shared_ptr<RandomAccessAudioReadable> clone() const {
    assert(false && "clone() not implemented for this RandomAccessAudioReadable type");
    return nullptr;
  }

  // Reads the audio data starting at frameOffset into the output buffer.
  // output.getChannelCount() | getChannelCount() | behavior
  // 1                        | 2                 | downmix (average)
  // 2                        | 1                 | upmix (copy)
  // x                        | x                 | read
  // otherwise: assertion failure
  void readWithMuxing(int frameOffset, BufferF32 output) {
    const auto myChannels = getChannelCount();
    const auto theirChannels = output.getChannelCount();
    if (myChannels == 2 && theirChannels == 1) {
      float stereoVLA[output.getFrameCount() * myChannels];
      auto muxbuf =
          BufferF32::fromVLA(myChannels, output.getFrameCount(), stereoVLA);
      read(frameOffset, muxbuf);

      const auto a = muxbuf.getChannelData(0);
      const auto b = muxbuf.getChannelData(1);
      const auto out = output.getChannelData(0);
      const auto frameCount = output.getFrameCount();

#pragma clang loop vectorize(enable)
      for (int i = 0; i < frameCount; i++) {
        out[i] = (a[i] + b[i]) / 2.f;
      }
    } else if (myChannels == 1 && theirChannels == 2) {
      auto muxbuf = output.sliceChannel(0);
      read(frameOffset, muxbuf);
      output.cloneChannel(0, 1);
    } else if (myChannels == theirChannels) {
      read(frameOffset, output);
    } else {
      assert(false);
    }
  }

  // Reads the audio data starting at frameOffset into the output buffer.
  // Regions of the output buffer that are outside the bounds of the audio data
  // are zero-padded.
  void readZeroPadded(int frameOffset, BufferF32 output) {
    assert(output.getChannelCount() == getChannelCount());
    int frameOffsetInOutput = 0;
    if (frameOffset < 0) {
      frameOffsetInOutput = std::min(-frameOffset, (int)output.getFrameCount());
      output.fill(0, frameOffsetInOutput, 0.f);
      frameOffset = 0;
    }
    int frameCountToRead =
        std::min((int)output.getFrameCount() - frameOffsetInOutput,
                 getFrameCount() - frameOffset);
    if (frameCountToRead > 0) {
      read(frameOffset, output.slice(frameOffsetInOutput,
                                     frameOffsetInOutput + frameCountToRead));
    } else {
      frameCountToRead = 0;
    }
    if (frameOffsetInOutput + frameCountToRead < (int)output.getFrameCount()) {
      output.fill(
          frameOffsetInOutput + frameCountToRead,
          output.getFrameCount() - frameOffsetInOutput - frameCountToRead, 0.f);
    }
  }
};

class OffsetView : public RandomAccessAudioReadable {
private:
  std::shared_ptr<RandomAccessAudioReadable> underlying;
  int startFrame;
  int endFrame;

public:
  OffsetView(std::shared_ptr<RandomAccessAudioReadable> underlying,
             int startFrame, int endFrame)
      : underlying(underlying), startFrame(startFrame), endFrame(endFrame) {};

  int getSampleRate() const override { return underlying->getSampleRate(); }
  int getChannelCount() const override { return underlying->getChannelCount(); }
  int getFrameCount() const override { return endFrame - startFrame; }

  void read(int frameOffset, BufferF32 output) override {
    underlying->readZeroPadded(frameOffset + startFrame, output);
  }

  void read(int frameOffset, std::shared_ptr<BufferF32> output) override {
    read(frameOffset, output->slice(0, output->getFrameCount()));
  }
  
  std::shared_ptr<RandomAccessAudioReadable> clone() const override {
    return std::make_shared<OffsetView>(underlying->clone(), startFrame, endFrame);
  }
};