///|
/// A fixed-size rollout window. `mask[i]` is false for padding positions.
pub struct SequenceSample[S, A] {
  start : Int
  transitions : Array[Transition[S, A]?]
  mask : Array[Bool]
}

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

///|
pub fn[S, A] SequenceSample::length(self : SequenceSample[S, A]) -> Int {
  self.transitions.length()
}

///|
pub fn[S, A] SequenceSample::transitions(
  self : SequenceSample[S, A],
) -> Array[Transition[S, A]?] {
  self.transitions
}

///|
pub fn[S, A] SequenceSample::mask(self : SequenceSample[S, A]) -> Array[Bool] {
  self.mask
}

///|
pub fn[S, A] SequenceSample::valid_length(self : SequenceSample[S, A]) -> Int {
  let mut count = 0
  for valid in self.mask {
    if valid {
      count = count + 1
    }
  }
  count
}

///|
pub fn[S, A] SequenceSample::is_padded(self : SequenceSample[S, A]) -> Bool {
  self.valid_length() < self.length()
}

///|
pub fn[S, A] ReplayBuffer::sequence(
  self : ReplayBuffer[S, A],
  start : Int,
  sequence_length : Int,
  pad : Bool,
) -> SequenceSample[S, A] {
  let transitions : Array[Transition[S, A]?] = []
  let mask : Array[Bool] = []
  if sequence_length <= 0 {
    return { start, transitions, mask }
  }
  let mut offset = 0
  while offset < sequence_length {
    match self.get(start + offset) {
      Some(transition) => {
        transitions.push(Some(transition))
        mask.push(true)
      }
      None =>
        if pad {
          transitions.push(None)
          mask.push(false)
        }
    }
    offset = offset + 1
  }
  { start, transitions, mask }
}

///|
pub fn[S, A] ReplayBuffer::valid_sequence_starts(
  self : ReplayBuffer[S, A],
  sequence_length : Int,
  allow_padding : Bool,
) -> Array[Int] {
  let starts : Array[Int] = []
  if sequence_length <= 0 || self.len == 0 {
    return starts
  }
  let last_start = if allow_padding {
    self.len - 1
  } else {
    self.len - sequence_length
  }
  let mut start = 0
  while start <= last_start {
    starts.push(start)
    start = start + 1
  }
  starts
}

///|
pub fn[S, A] ReplayBuffer::sample_sequences(
  self : ReplayBuffer[S, A],
  batch_size : Int,
  sequence_length : Int,
  seed : Int,
  allow_padding : Bool,
) -> Array[SequenceSample[S, A]] {
  let result : Array[SequenceSample[S, A]] = []
  if batch_size <= 0 {
    return result
  }
  let starts = self.valid_sequence_starts(sequence_length, allow_padding)
  if starts.length() == 0 {
    return result
  }
  let rng = ReplayRng::new(seed)
  let mut i = 0
  while i < batch_size {
    let start = starts[rng.next_index(starts.length())]
    result.push(self.sequence(start, sequence_length, allow_padding))
    i = i + 1
  }
  result
}

///|
pub fn[S, A] ReplayBuffer::all_sequences(
  self : ReplayBuffer[S, A],
  sequence_length : Int,
) -> Array[SequenceSample[S, A]] {
  let result : Array[SequenceSample[S, A]] = []
  for start in self.valid_sequence_starts(sequence_length, false) {
    result.push(self.sequence(start, sequence_length, false))
  }
  result
}

///|
pub fn[S, A] SequenceSample::rewards(
  self : SequenceSample[S, A],
) -> Array[Double] {
  let result : Array[Double] = []
  for transition in self.transitions {
    match transition {
      Some(value) => result.push(value.reward)
      None => result.push(0.0)
    }
  }
  result
}

///|
pub fn[S, A] SequenceSample::terminals(
  self : SequenceSample[S, A],
) -> Array[Bool] {
  let result : Array[Bool] = []
  for transition in self.transitions {
    match transition {
      Some(value) => result.push(value.done)
      None => result.push(false)
    }
  }
  result
}