///|
/// A lightweight offline-RL dataset made of complete episodes.
pub struct TrajectoryDataset[S, A] {
  episodes : Array[Episode[S, A]]
  mut transition_count : Int
}

///|
pub fn[S, A] TrajectoryDataset::new() -> TrajectoryDataset[S, A] {
  { episodes: [], transition_count: 0 }
}

///|
pub fn[S, A] TrajectoryDataset::from_episodes(
  episodes : Array[Episode[S, A]],
) -> TrajectoryDataset[S, A] {
  let dataset = TrajectoryDataset::new()
  for episode in episodes {
    let _ = dataset.add(episode)
  }
  dataset
}

///|
pub fn[S, A] TrajectoryDataset::add(
  self : TrajectoryDataset[S, A],
  episode : Episode[S, A],
) -> Bool {
  if episode.validate().length() != 0 {
    return false
  }
  self.episodes.push(episode)
  self.transition_count = self.transition_count + episode.len()
  true
}

///|
pub fn[S, A] TrajectoryDataset::add_unchecked(
  self : TrajectoryDataset[S, A],
  episode : Episode[S, A],
) -> Unit {
  self.episodes.push(episode)
  self.transition_count = self.transition_count + episode.len()
}

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

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

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

///|
pub fn[S, A] TrajectoryDataset::episodes(
  self : TrajectoryDataset[S, A],
) -> Array[Episode[S, A]] {
  self.episodes
}

///|
pub fn[S, A] TrajectoryDataset::get(
  self : TrajectoryDataset[S, A],
  index : Int,
) -> Episode[S, A]? {
  if index < 0 || index >= self.episodes.length() {
    None
  } else {
    Some(self.episodes[index])
  }
}

///|
pub fn[S, A] TrajectoryDataset::flatten(
  self : TrajectoryDataset[S, A],
) -> Array[Transition[S, A]] {
  let result : Array[Transition[S, A]] = []
  for episode in self.episodes {
    for transition in episode.transitions() {
      result.push(transition)
    }
  }
  result
}

///|
pub fn[S, A] TrajectoryDataset::reward_stats(
  self : TrajectoryDataset[S, A],
) -> RewardStats {
  reward_stats_from_transitions(self.flatten())
}

///|
pub fn[S, A] TrajectoryDataset::episode_lengths(
  self : TrajectoryDataset[S, A],
) -> Array[Int] {
  let result : Array[Int] = []
  for episode in self.episodes {
    result.push(episode.len())
  }
  result
}

///|
pub fn[S, A] TrajectoryDataset::sample_episodes(
  self : TrajectoryDataset[S, A],
  count : Int,
  seed : Int,
) -> Array[Episode[S, A]] {
  let result : Array[Episode[S, A]] = []
  if count <= 0 || self.episodes.length() == 0 {
    return result
  }
  let rng = ReplayRng::new(seed)
  let mut i = 0
  while i < count {
    result.push(self.episodes[rng.next_index(self.episodes.length())])
    i = i + 1
  }
  result
}

///|
pub fn[S, A] TrajectoryDataset::split_at(
  self : TrajectoryDataset[S, A],
  first_count : Int,
) -> (TrajectoryDataset[S, A], TrajectoryDataset[S, A]) {
  let pivot = if first_count < 0 {
    0
  } else if first_count > self.episodes.length() {
    self.episodes.length()
  } else {
    first_count
  }
  let left = TrajectoryDataset::new()
  let right = TrajectoryDataset::new()
  let mut i = 0
  while i < self.episodes.length() {
    if i < pivot {
      left.add_unchecked(self.episodes[i])
    } else {
      right.add_unchecked(self.episodes[i])
    }
    i = i + 1
  }
  (left, right)
}

///|
pub fn[S, A] TrajectoryDataset::terminal_episode_count(
  self : TrajectoryDataset[S, A],
) -> Int {
  let mut count = 0
  for episode in self.episodes {
    if episode.is_terminated() {
      count = count + 1
    }
  }
  count
}

///|
pub fn[S, A] TrajectoryDataset::truncated_episode_count(
  self : TrajectoryDataset[S, A],
) -> Int {
  let mut count = 0
  for episode in self.episodes {
    if episode.is_truncated() {
      count = count + 1
    }
  }
  count
}

///|
pub fn[S, A] TrajectoryDataset::discounted_returns(
  self : TrajectoryDataset[S, A],
  gamma : Double,
) -> Array[Double] {
  let result : Array[Double] = []
  for episode in self.episodes {
    for value in episode.discounted_returns(gamma) {
      result.push(value)
    }
  }
  result
}