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

namespace time_transform {

// Map for arrangement offsets, (co)domain are beats, fulfills conditions:
// - Domain [-inf, timelineStart) is mapped to codomain [-inf, 0)
// - Domain [timelineStart, timelineEnd) is mapped to codomain [0, timelineEnd -
// timelineStart)
// - Domain [timelineEnd, inf) is mapped to codomain [timelineEnd -
// timelineStart, inf) where timelineStart < timelineEnd
// - Iterator iterates over segments:
//     S1: Maps [-inf, timelineStart) to [-inf, 0)
//     S2: Maps [timelineStart, timelineEnd) to [0, timelineEnd - timelineStart)
//     S3: Maps [timelineEnd, inf) to [timelineEnd - timelineStart, inf)
template <typename DomainUnit, TimeValue T = double>
class ArrangementOffsetMap {
public:
  using DomainTime = TypedTime<DomainUnit, T>;
  using CodomainTime = TypedTime<DomainUnit, T>; // Domain = Codomain
  using DomainDelta = TypedTimeDelta<DomainUnit, T>;
  using CodomainDelta = TypedTimeDelta<DomainUnit, T>;

  class iterator; // Forward declaration

private:
  DomainTime timelineStart_;
  DomainTime timelineEnd_;
  CodomainTime codomainOffsetStart_; // Cache 0
  CodomainTime codomainOffsetEnd_;   // Cache timelineEnd_ - timelineStart_

public:
  explicit ArrangementOffsetMap(DomainTime timelineStart,
                                DomainTime timelineEnd)
      : timelineStart_(timelineStart), timelineEnd_(timelineEnd),
        codomainOffsetStart_(0.0), codomainOffsetEnd_(CodomainTime(
                                       (timelineEnd_ - timelineStart_).raw())) {
    assert(timelineStart_ < timelineEnd_ &&
           "Timeline start must be before timeline end.");
    assert(timelineStart_.is_finite() && timelineEnd_.is_finite() &&
           "Timeline bounds must be finite.");
  }

  DomainTime getTimelineStart() const { return timelineStart_; }
  DomainTime getTimelineEnd() const { return timelineEnd_; }

  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 ArrangementOffsetMap *map_ptr_;
    int segment_idx_; // 0: before, 1: inside, 2: after
    mutable std::optional<value_type> current_mapped_segment_;

    void cache_current() const {
      if (current_mapped_segment_ || !map_ptr_ || segment_idx_ < 0 ||
          segment_idx_ > 2)
        return;

      const T slope = 1.0;
      const DomainTime dom_neg_inf = DomainTime::neg_inf();
      const DomainTime dom_inf = DomainTime::inf();
      const CodomainTime codom_neg_inf = CodomainTime::neg_inf();
      const CodomainTime codom_inf = CodomainTime::inf();

      if (segment_idx_ == 0) {
        // S1: Maps [-inf, timelineStart) to [-inf, 0)
        TimeRange<DomainUnit, T> source_range(dom_neg_inf,
                                              map_ptr_->timelineStart_);
        TimeRange<DomainUnit, T> target_range(codom_neg_inf,
                                              map_ptr_->codomainOffsetStart_);
        // Right anchored at timelineStart which maps to 0
        current_mapped_segment_.emplace(source_range, target_range, slope,
                                        SegmentMarks::INVALID_ARRANGEMENT_TIME);
      } else if (segment_idx_ == 1) {
        // S2: Maps [timelineStart, timelineEnd) to [0, timelineEnd -
        // timelineStart)
        TimeRange<DomainUnit, T> source_range(map_ptr_->timelineStart_,
                                              map_ptr_->timelineEnd_);
        TimeRange<DomainUnit, T> target_range(map_ptr_->codomainOffsetStart_,
                                              map_ptr_->codomainOffsetEnd_);
        // Left anchored at timelineStart which maps to 0
        current_mapped_segment_.emplace(source_range, target_range, slope,
                                        SegmentMarks::NONE);
      } else { // segment_idx_ == 2
        // S3: Maps [timelineEnd, inf) to [timelineEnd - timelineStart, inf)
        TimeRange<DomainUnit, T> source_range(map_ptr_->timelineEnd_, dom_inf);
        TimeRange<DomainUnit, T> target_range(map_ptr_->codomainOffsetEnd_,
                                              codom_inf);
        // Left anchored at timelineEnd which maps to timelineEnd -
        // timelineStart
        current_mapped_segment_.emplace(source_range, target_range, slope,
                                        SegmentMarks::INVALID_ARRANGEMENT_TIME);
      }
    }

  public:
    // Constructor for valid iterators
    iterator(const ArrangementOffsetMap *map, int idx)
        : map_ptr_(map), segment_idx_(idx) {
      assert(map_ptr_ != nullptr && segment_idx_ >= 0 && segment_idx_ <= 3 &&
             "Creating invalid iterator state");
    }

    // 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 ArrangementOffsetMap iterator");
      return *current_mapped_segment_;
    }
    pointer operator->() const {
      cache_current();
      assert(current_mapped_segment_.has_value() &&
             "Dereferencing invalid ArrangementOffsetMap iterator");
      return &(*current_mapped_segment_);
    }

    iterator &operator++() {
      if (map_ptr_ && segment_idx_ >= 0 && segment_idx_ < 3) {
        current_mapped_segment_.reset(); // Invalidate cache
        segment_idx_++;
      } else {
        // Already an end iterator or invalid, ensure canonical end state
        map_ptr_ = nullptr;
        segment_idx_ = -1;
      }
      // Make end iterators canonical
      if (segment_idx_ == 3) {
        map_ptr_ = nullptr;
        segment_idx_ = -1; // Match default constructor end state
      }
      return *this;
    }

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

    iterator &operator--() {
      if (map_ptr_ == nullptr && segment_idx_ == -1) { // Canonical end iterator
        // This case is tricky: what should --end() be? It should be the last
        // valid element. For GlobalLoopMap, the "last" element is segment 1,
        // which is infinitely looping. To make this work, we'd need the
        // map_ptr_ again. This suggests -- should not be called on the
        // default-constructed end iterator. Assuming this is called on a valid
        // iterator or map_ptr_->end().
        assert(map_ptr_ != nullptr && "Cannot decrement default-constructed "
                                      "end iterator without map context");
        current_mapped_segment_.reset();
        segment_idx_ = 1; // The segment before end() is segment 1
        return *this;
      }

      current_mapped_segment_.reset();
      if (segment_idx_ == 1) { // From S2 (loop) to S1
        segment_idx_ = 0;
      } else if (segment_idx_ == 2) { // From end() to S2 (loop)
        segment_idx_ = 1;
      }
      // If segment_idx_ is 0 (begin), it cannot be decremented further to a
      // valid segment. If segment_idx_ is -1 (and map_ptr_ is not null, e.g. an
      // invalid state), behavior is undefined.
      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 ArrangementOffsetMap<DomainUnit, T>;
  };

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

  // End iterator represents state after segment 2 (index 3 internally before
  // canonicalization)
  iterator end() const { return iterator(); } // Use default constructed end
  iterator end() { return iterator(); }       // Non-const

  iterator getSegmentIteratorAt(const DomainTime &p) const {
    if (p < timelineStart_) {
      return iterator(this, 0); // Before segment
    } else if (p < timelineEnd_) {
      return iterator(this, 1); // Inside segment
    } else {
      return iterator(this, 2); // After segment
    }
  }
};

using ArrangementBeatOffsetMap = ArrangementOffsetMap<BeatsTag, double>;

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

} // namespace time_transform