///|
/// One transition in an RL rollout.
pub struct Transition[S, A] {
  state : S
  action : A
  reward : Double
  next_state : S
  done : Bool
}

///|
pub fn[S, A] Transition::new(
  state : S,
  action : A,
  reward : Double,
  next_state : S,
  done : Bool,
) -> Transition[S, A] {
  { state, action, reward, next_state, done }
}

///|
pub fn[S, A] Transition::terminal(
  state : S,
  action : A,
  reward : Double,
  next_state : S,
) -> Transition[S, A] {
  Transition::new(state, action, reward, next_state, true)
}

///|
pub fn[S, A] Transition::non_terminal(
  state : S,
  action : A,
  reward : Double,
  next_state : S,
) -> Transition[S, A] {
  Transition::new(state, action, reward, next_state, false)
}

///|
/// A compact episode container for ordered transitions.
pub struct Episode[S, A] {
  transitions : Array[Transition[S, A]]
  terminated : Bool
  truncated : Bool
  total_reward : Double
}

///|
pub fn[S, A] Episode::new() -> Episode[S, A] {
  let transitions : Array[Transition[S, A]] = []
  { transitions, terminated: false, truncated: false, total_reward: 0.0 }
}

///|
pub fn[S, A] Episode::from_transitions(
  transitions : Array[Transition[S, A]],
) -> Episode[S, A] {
  let mut episode = Episode::new()
  for transition in transitions {
    episode = episode.push(transition)
  }
  episode
}

///|
pub fn[S, A] Episode::push(
  self : Episode[S, A],
  transition : Transition[S, A],
) -> Episode[S, A] {
  let transitions = self.transitions
  transitions.push(transition)
  {
    transitions,
    terminated: if transition.done {
      true
    } else {
      self.terminated
    },
    truncated: self.truncated,
    total_reward: self.total_reward + transition.reward,
  }
}

///|
pub fn[S, A] Episode::mark_terminated(self : Episode[S, A]) -> Episode[S, A] {
  { ..self, terminated: true }
}

///|
pub fn[S, A] Episode::mark_truncated(self : Episode[S, A]) -> Episode[S, A] {
  { ..self, truncated: true }
}

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

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

///|
pub fn[S, A] Episode::is_terminated(self : Episode[S, A]) -> Bool {
  self.terminated
}

///|
pub fn[S, A] Episode::is_truncated(self : Episode[S, A]) -> Bool {
  self.truncated
}

///|
pub fn[S, A] Episode::total_reward(self : Episode[S, A]) -> Double {
  self.total_reward
}

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

///|
pub fn[S, A] Episode::first_transition(
  self : Episode[S, A],
) -> Transition[S, A]? {
  if self.transitions.length() == 0 {
    None
  } else {
    Some(self.transitions[0])
  }
}

///|
pub fn[S, A] Episode::last_transition(
  self : Episode[S, A],
) -> Transition[S, A]? {
  let len = self.transitions.length()
  if len == 0 {
    None
  } else {
    Some(self.transitions[len - 1])
  }
}

///|
/// Return discounted returns in rollout order.
pub fn[S, A] Episode::discounted_returns(
  self : Episode[S, A],
  gamma : Double,
) -> Array[Double] {
  let len = self.transitions.length()
  let returns : Array[Double] = []
  let mut i = 0
  while i < len {
    returns.push(0.0)
    i = i + 1
  }

  let mut running = 0.0
  let mut idx = len
  while idx > 0 {
    idx = idx - 1
    running = self.transitions[idx].reward + gamma * running
    returns[idx] = running
  }
  returns
}

///|
/// Return N-step transitions.
/// For each transition i, the reward is accumulated for N steps (or until done),
/// and the next_state is the state after N steps (or the terminal state).
pub fn[S, A] Episode::n_step_transitions(
  self : Episode[S, A],
  n : Int,
  gamma : Double,
) -> Array[Transition[S, A]] {
  if n <= 1 {
    return self.transitions
  }
  let len = self.transitions.length()
  let result : Array[Transition[S, A]] = []
  let mut i = 0
  while i < len {
    let start_transition = self.transitions[i]
    let mut accumulated_reward = 0.0
    let mut current_gamma = 1.0
    let mut next_state = start_transition.next_state
    let mut done = start_transition.done

    let mut step = 0
    while step < n && i + step < len {
      let t = self.transitions[i + step]
      accumulated_reward = accumulated_reward + current_gamma * t.reward
      next_state = t.next_state
      done = t.done
      current_gamma = current_gamma * gamma
      if done {
        break
      }
      step = step + 1
    }

    result.push(
      Transition::new(
        start_transition.state,
        start_transition.action,
        accumulated_reward,
        next_state,
        done,
      ),
    )
    i = i + 1
  }
  result
}