#pragma once
#include "timing/time_transform.h"

namespace time_transform {

// Map for arrangement loops, (co)domain are beats, fulfills conditions:
// - For each integer i, domain
//   [prefixLength + i * loopLength, prefixLength + (i + 1) * loopLength)
//   is mapped to codomain
//   [loopStart, loopEnd)
// where
//   prefixLength = loopEnd - readStart,
//   loopLength = loopEnd - loopStart,
//   loopStart <= readStart < loopEnd
// - Iterator iterates over segments:
//     Si: Maps
//       [prefixLength + i * loopLength, prefixLength + (i + 1) * loopLength)
//     to [loopStart, loopEnd)
// The iterators are bidirectionally infinite, so can never be equal to begin()
// or end()
template <typename DomainUnit, TimeValue T = double> class ArrangementLoopMap {
public:
  using DomainTime = TypedTime<DomainUnit, T>;
  using CodomainTime = TypedTime<DomainUnit, T>;
  using DomainDelta = TypedTimeDelta<DomainUnit, T>;
  using CodomainDelta = TypedTimeDelta<DomainUnit, T>;

  class iterator; // Forward declaration

private:
  DomainTime readStart_;
  DomainTime loopStart_;
  DomainTime loopEnd_;
  DomainDelta loopLength_;
  DomainDelta prefixLength_;
  DomainTime prefixEndDomainTime_;

public:
  explicit ArrangementLoopMap(DomainTime readStart, DomainTime loopStart,
                              DomainTime loopEnd)
      : readStart_(readStart), loopStart_(loopStart), loopEnd_(loopEnd),
        loopLength_(loopEnd - loopStart), prefixLength_(loopEnd - readStart),
        prefixEndDomainTime_(DomainTime(0.0) + prefixLength_) {
    assert(loopStart_ < loopEnd_ && "Loop start must be before loop end.");
    // Allow readStart == loopStart
    assert(readStart_ >= loopStart_ &&
           "Read start must be at or after loop start.");
    assert(readStart_ < loopEnd_ && "Read start must be before loop end.");
    assert(loopStart_.is_finite() && loopEnd_.is_finite() &&
           readStart_.is_finite() && "Loop bounds must be finite.");
    assert(loopLength_.raw() > 0 && "Loop length must be positive.");
    assert(prefixLength_.raw() >= 0 && "Prefix length must be non-negative.");
  }

public:
  class iterator {
  public:
    using iterator_category = std::bidirectional_iterator_tag;
    using value_type = MappedSegment<DomainUnit, DomainUnit, T>;
    using difference_type = std::ptrdiff_t;
    using pointer = const value_type *;
    using reference = const value_type &;

  private:
    const ArrangementLoopMap *map_ptr_;
    int segment_idx_; // 0: prefix segment, >=1: loop segments
    mutable std::optional<value_type> current_mapped_segment_;

    void cache_current() const {
      if (current_mapped_segment_ || !map_ptr_)
        return;

      const T slope = 1.0;
      int i = segment_idx_;

      DomainTime domain_start =
          map_ptr_->prefixEndDomainTime_ + (i * map_ptr_->loopLength_);
      DomainTime domain_end = domain_start + map_ptr_->loopLength_;

      TimeRange<DomainUnit, T> source_range(domain_start, domain_end);
      TimeRange<DomainUnit, T> target_range(map_ptr_->loopStart_,
                                            map_ptr_->loopEnd_);

      current_mapped_segment_.emplace(source_range, target_range, slope,
                                      SegmentMarks::NONE);
    }

  public:
    // Constructor for valid iterators
    iterator(const ArrangementLoopMap *map, int idx)
        : map_ptr_(map), segment_idx_(idx) {
      assert(map_ptr_ != nullptr && "Creating iterator with null map_ptr_");
    }

    // Default constructor for placeholder/end iterator
    iterator() : map_ptr_(nullptr), segment_idx_(-1) {}

    reference operator*() const {
      cache_current();
      assert(current_mapped_segment_.has_value() &&
             "Dereferencing invalid ArrangementLoopMap iterator");
      return *current_mapped_segment_;
    }
    pointer operator->() const {
      cache_current();
      assert(current_mapped_segment_.has_value() &&
             "Dereferencing invalid ArrangementLoopMap iterator");
      return &(*current_mapped_segment_);
    }

    iterator &operator++() {
      if (map_ptr_) {
        current_mapped_segment_.reset();
        segment_idx_++;
      }
      return *this;
    }

    iterator operator++(int) {
      iterator tmp = *this;
      ++(*this);
      return tmp;
    }

    iterator &operator--() {
      if (map_ptr_) {
        current_mapped_segment_.reset();
        segment_idx_--;
      }
      return *this;
    }

    iterator operator--(int) {
      iterator tmp = *this;
      --(*this);
      return tmp;
    }

    bool operator==(const iterator &other) const {
      // Both are end iterators
      if (!map_ptr_ && !other.map_ptr_ && segment_idx_ == -1 &&
          other.segment_idx_ == -1) {
        return true;
      }
      // Otherwise compare map pointer and index
      return map_ptr_ == other.map_ptr_ && segment_idx_ == other.segment_idx_;
    }

    bool operator!=(const iterator &other) const { return !(*this == other); }
    friend class ArrangementLoopMap<DomainUnit, T>;
  };

  iterator begin() const { return iterator(this, 0); }
  iterator begin() { return iterator(this, 0); } // Non-const

  // End iterator is a sentinel value, as the sequence is infinite
  iterator end() const { return iterator(); }
  iterator end() { return iterator(); } // Non-const

  iterator getSegmentIteratorAt(const DomainTime &p) const {
    if (!p.is_finite()) {
      int i;
      if (p.raw() < 0) { // Negative infinity
        i = std::numeric_limits<int>::min();
      } else { // Positive infinity
        i = std::numeric_limits<int>::max();
      }
      return iterator(this, i);
    }

    DomainDelta time_relative_to_prefix_end = p - prefixEndDomainTime_;
    double i_double = time_relative_to_prefix_end.raw() / loopLength_.raw();
    int i = static_cast<int>(std::floor(i_double));

    return iterator(this, i);
  }
};

using ArrangementBeatLoopMap = ArrangementLoopMap<BeatsTag, double>;

static_assert(
    IsTimeTransformer<ArrangementBeatLoopMap>,
    "ArrangementBeatLoopMap does not satisfy the IsTimeTransformer concept.");

} // namespace time_transform