///|
pub struct ReplaySample[S, A] {
  index : Int
  transition : Transition[S, A]
}

///|
pub struct ReplayBuffer[S, A] {
  capacity : Int
  mut len : Int
  mut head : Int
  data : Array[Transition[S, A]?]
}

///|
pub fn[S, A] ReplayBuffer::new(capacity : Int) -> ReplayBuffer[S, A] {
  let data : Array[Transition[S, A]?] = []
  let mut i = 0
  while i < capacity {
    data.push(None)
    i = i + 1
  }
  { capacity, len: 0, head: 0, data }
}

///|
pub fn[S, A] ReplayBuffer::capacity(self : ReplayBuffer[S, A]) -> Int {
  self.capacity
}

///|
pub fn[S, A] ReplayBuffer::len(self : ReplayBuffer[S, A]) -> Int {
  self.len
}

///|
pub fn[S, A] ReplayBuffer::is_empty(self : ReplayBuffer[S, A]) -> Bool {
  self.len == 0
}

///|
pub fn[S, A] ReplayBuffer::is_full(self : ReplayBuffer[S, A]) -> Bool {
  self.capacity > 0 && self.len == self.capacity
}

///|
pub fn[S, A] ReplayBuffer::clear(
  self : ReplayBuffer[S, A],
) -> ReplayBuffer[S, A] {
  let mut i = 0
  while i < self.data.length() {
    self.data[i] = None
    i = i + 1
  }
  { ..self, len: 0, head: 0 }
}

///|
pub fn[S, A] ReplayBuffer::get(
  self : ReplayBuffer[S, A],
  logical_index : Int,
) -> Transition[S, A]? {
  if logical_index < 0 || logical_index >= self.len || self.capacity == 0 {
    None
  } else {
    let slot = (self.head + logical_index) % self.capacity
    self.data[slot]
  }
}

///|
pub fn[S, A] ReplayBuffer::oldest(
  self : ReplayBuffer[S, A],
) -> Transition[S, A]? {
  self.get(0)
}

///|
pub fn[S, A] ReplayBuffer::newest(
  self : ReplayBuffer[S, A],
) -> Transition[S, A]? {
  if self.len == 0 {
    None
  } else {
    self.get(self.len - 1)
  }
}

///|
pub fn[S, A] ReplayBuffer::push(
  self : ReplayBuffer[S, A],
  transition : Transition[S, A],
) -> Int {
  if self.capacity == 0 {
    return 0
  }

  if self.len < self.capacity {
    self.data[self.len] = Some(transition)
    self.len = self.len + 1
  } else {
    self.data[self.head] = Some(transition)
    self.head = (self.head + 1) % self.capacity
  }
  self.len
}

///|
pub fn[S, A] ReplayBuffer::extend_episode(
  self : ReplayBuffer[S, A],
  episode : Episode[S, A],
) -> Int {
  for transition in episode.transitions() {
    let _ = self.push(transition)
  }
  self.len
}

///|
pub fn[S, A] ReplayBuffer::to_array(
  self : ReplayBuffer[S, A],
) -> Array[Transition[S, A]] {
  let result : Array[Transition[S, A]] = []
  let mut i = 0
  while i < self.len {
    match self.get(i) {
      Some(transition) => result.push(transition)
      None => ()
    }
    i = i + 1
  }
  result
}

///|
pub fn[S, A] ReplayBuffer::sample_batch(
  self : ReplayBuffer[S, A],
  batch_size : Int,
  seed : Int,
) -> Array[ReplaySample[S, A]] {
  let target = if batch_size < self.len { batch_size } else { self.len }
  let samples : Array[ReplaySample[S, A]] = []
  if target == 0 {
    return samples
  }

  let indices = shuffle_prefix(make_sequential_indices(self.len), target, seed)
  let mut i = 0
  while i < target {
    let idx = indices[i]
    match self.get(idx) {
      Some(transition) => samples.push({ index: idx, transition })
      None => ()
    }
    i = i + 1
  }
  samples
}