///|
enum EpisodeStage {
  Running
  Finished
} derive(Debug, Eq)

///|
pub struct Transition {
  state : Int
  action : Int
  reward : Double
  next_state : Int
  done : Bool
  mut step : Int
} derive(Debug, Eq)

///|
pub struct EpisodeRecord {
  episode : Int
  steps : Int
  reward : Double
  goal_reached : Bool
  final_state : Int
} derive(Debug, Eq)

///|
pub struct TrainingReport {
  label : String
  episodes : Int
  mut rewards : Array[Double]
  mut steps : Array[Int]
  mut goal_hits : Int
  mut final_epsilon : Double
} derive(Debug)

///|
fn TrainingReport::summary(self : TrainingReport) -> String {
  if self.episodes == 0 {
    "label=\{self.label}\nepisodes=0\ngoal_hits=0\naverage_reward=0.0\naverage_steps=0.0\nrecent_12_reward=0.0\nbest_episode=0\nbest_reward=0.0\nfinal_epsilon=\{self.final_epsilon}"
  } else {
    let mut total_reward = 0.0
    let mut total_steps = 0
    let mut best_reward = -999999.0
    let mut best_episode = 0
    let recent_start = if self.episodes > 12 { self.episodes - 12 } else { 0 }
    let mut recent_reward = 0.0
    let mut recent_count = 0
    let mut i = 0
    while i < self.episodes {
      let reward = self.rewards[i]
      let step_count = self.steps[i]
      total_reward = total_reward + reward
      total_steps = total_steps + step_count
      if reward > best_reward {
        best_reward = reward
        best_episode = i + 1
      }
      if i >= recent_start {
        recent_reward = recent_reward + reward
        recent_count = recent_count + 1
      }
      i = i + 1
    }
    let average_reward = total_reward / self.episodes.to_double()
    let average_steps = total_steps.to_double() / self.episodes.to_double()
    let recent_average = if recent_count == 0 {
      0.0
    } else {
      recent_reward / recent_count.to_double()
    }
    let summary0 = "label=\{self.label}\n"
    let summary1 = summary0 + "episodes=\{self.episodes}\n"
    let summary2 = summary1 + "goal_hits=\{self.goal_hits}\n"
    let summary3 = summary2 + "average_reward=\{average_reward}\n"
    let summary4 = summary3 + "average_steps=\{average_steps}\n"
    let summary5 = summary4 + "recent_12_reward=\{recent_average}\n"
    let summary6 = summary5 + "best_episode=\{best_episode}\n"
    let summary7 = summary6 + "best_reward=\{best_reward}\n"
    let summary8 = summary7 + "final_epsilon=\{self.final_epsilon}"
    summary8
  }
}

///|
pub fn TrainingReport::compact_line(self : TrainingReport) -> String {
  if self.episodes == 0 {
    "\{self.label}: avg_reward=0.0, tail10=0.0, goal_hits=0"
  } else {
    let mut total_reward = 0.0
    let mut tail_reward = 0.0
    let mut tail_count = 0
    let tail_start = if self.episodes > 10 { self.episodes - 10 } else { 0 }
    let mut i = 0
    while i < self.episodes {
      total_reward = total_reward + self.rewards[i]
      if i >= tail_start {
        tail_reward = tail_reward + self.rewards[i]
        tail_count = tail_count + 1
      }
      i = i + 1
    }
    let tail_average = if tail_count == 0 {
      0.0
    } else {
      tail_reward / tail_count.to_double()
    }
    "\{self.label}: avg_reward=\{total_reward / self.episodes.to_double()}, tail10=\{tail_average}, goal_hits=\{self.goal_hits}"
  }
}

///|
pub fn TrainingReport::episode_count(self : TrainingReport) -> Int {
  self.episodes
}

///|
pub fn TrainingReport::goal_count(self : TrainingReport) -> Int {
  self.goal_hits
}

///|
pub fn TrainingReport::reward_at(
  self : TrainingReport,
  episode : Int,
) -> Double {
  if episode < 0 || episode >= self.rewards.length() {
    0.0
  } else {
    self.rewards[episode]
  }
}

///|
pub fn TrainingReport::steps_at(self : TrainingReport, episode : Int) -> Int {
  if episode < 0 || episode >= self.steps.length() {
    0
  } else {
    self.steps[episode]
  }
}

///|
pub fn TrainingReport::final_epsilon_value(self : TrainingReport) -> Double {
  self.final_epsilon
}

///|
pub fn Transition::new(
  state : Int,
  action : Int,
  reward : Double,
  next_state : Int,
  done : Bool,
  step : Int,
) -> Transition {
  { state, action, reward, next_state, done, step }
}

///|
struct QTable {
  state_count : Int
  action_count : Int
  values : Array[Array[Double]]
} derive(Debug)

///|
fn QTable::new(state_count : Int, action_count : Int) -> QTable {
  let safe_state_count = if state_count < 1 { 1 } else { state_count }
  let safe_action_count = if action_count < 1 { 1 } else { action_count }
  let values = Array::make(
    safe_state_count,
    Array::make(safe_action_count, 0.0),
  )
  let mut i = 0
  while i < safe_state_count {
    values[i] = Array::make(safe_action_count, 0.0)
    i = i + 1
  }
  { state_count: safe_state_count, action_count: safe_action_count, values }
}

///|
fn QTable::row(self : QTable, state : Int) -> Array[Double] {
  if state < 0 || state >= self.state_count {
    Array::make(self.action_count, 0.0)
  } else {
    self.values[state]
  }
}

///|
fn QTable::value(self : QTable, state : Int, action : Int) -> Double {
  if state < 0 ||
    state >= self.state_count ||
    action < 0 ||
    action >= self.action_count {
    0.0
  } else {
    self.values[state][action]
  }
}

///|
fn QTable::set_value(
  self : QTable,
  state : Int,
  action : Int,
  value : Double,
) -> Unit {
  if state >= 0 &&
    state < self.state_count &&
    action >= 0 &&
    action < self.action_count {
    self.values[state][action] = value
  }
}

///|
fn QTable::best_action_index(self : QTable, state : Int) -> Int {
  let row = self.row(state)
  let mut best = 0
  let mut best_value = row[0]
  let mut i = 1
  while i < row.length() {
    if row[i] > best_value {
      best = i
      best_value = row[i]
    }
    i = i + 1
  }
  best
}

///|
fn QTable::max_value(self : QTable, state : Int) -> Double {
  self.row(state)[self.best_action_index(state)]
}

///|
fn QTable::update(
  self : QTable,
  state : Int,
  action : Int,
  target : Double,
  alpha : Double,
) -> Unit {
  let old = self.value(state, action)
  self.set_value(state, action, old + alpha * (target - old))
}

///|
struct LcgRng {
  mut seed : Int
} derive(Debug)

///|
fn LcgRng::new(seed : Int) -> LcgRng {
  { seed: if seed <= 0 { 20260711 } else { seed } }
}

///|
fn LcgRng::next_seed(self : LcgRng) -> Int {
  self.seed = (self.seed * 1103515245 + 12345) % 2147483647
  if self.seed < 0 {
    self.seed = -self.seed
  }
  self.seed
}

///|
fn LcgRng::next_int(self : LcgRng, bound : Int) -> Int {
  if bound <= 1 {
    0
  } else {
    self.next_seed() % bound
  }
}

///|
fn LcgRng::next_double(self : LcgRng) -> Double {
  self.next_int(10000).to_double() / 10000.0
}

///|
pub struct EpsilonGreedyPolicy {
  epsilon : Double
  rng : LcgRng
} derive(Debug)

///|
pub fn EpsilonGreedyPolicy::new(
  epsilon : Double,
  seed : Int,
) -> EpsilonGreedyPolicy {
  { epsilon, rng: LcgRng::new(seed) }
}

///|
pub fn EpsilonGreedyPolicy::choose_action(
  self : EpsilonGreedyPolicy,
  _state : Int,
  actions : Array[Int],
  q_values : Array[Double],
) -> Int {
  if actions.length() == 0 || q_values.length() == 0 {
    0
  } else {
    let explore = self.rng.next_double() < self.epsilon
    if explore {
      actions[self.rng.next_int(actions.length())]
    } else {
      let mut best = 0
      let mut best_value = q_values[0]
      let mut i = 1
      while i < q_values.length() {
        if q_values[i] > best_value {
          best = i
          best_value = q_values[i]
        }
        i = i + 1
      }
      actions[best]
    }
  }
}

///|
pub(open) trait Environment {
  fn reset(Self) -> Int
  fn actions(Self) -> Array[Int]
  fn state_space(Self) -> Array[Int]
  fn step(Self, Int) -> Transition
  fn render(Self) -> String
}

///|
pub(open) trait Policy {
  fn choose_action(Self, Int, Array[Int], Array[Double]) -> Int
}

///|
pub(open) trait Agent {
  fn reset_episode(Self) -> Unit
  fn choose_action(Self, Int) -> Int
  fn learn(Self, Transition, Int?) -> Unit
  fn epsilon(Self) -> Double
  fn q_report(Self, Int) -> String
}

///|
pub(open) trait Logger {
  fn start_episode(Self, Int) -> Unit
  fn step(Self, Int, Int, Int, Double, Int, Bool, Int) -> Unit
  fn finish_episode(Self, EpisodeRecord) -> Unit
  fn finish(Self, TrainingReport) -> Unit
}

///|
pub struct GridWorldEnv {
  width : Int
  height : Int
  start_x : Int
  start_y : Int
  goal_x : Int
  goal_y : Int
  mut x : Int
  mut y : Int
  mut _last_stage : EpisodeStage
} derive(Debug)

///|
pub fn GridWorldEnv::new() -> GridWorldEnv {
  {
    width: 4,
    height: 4,
    start_x: 0,
    start_y: 0,
    goal_x: 3,
    goal_y: 3,
    x: 0,
    y: 0,
    _last_stage: Running,
  }
}

///|
fn GridWorldEnv::encode(self : GridWorldEnv, x : Int, y : Int) -> Int {
  y * self.width + x
}

///|
fn GridWorldEnv::decode_x(self : GridWorldEnv, state : Int) -> Int {
  state % self.width
}

///|
fn GridWorldEnv::decode_y(self : GridWorldEnv, state : Int) -> Int {
  state / self.width
}

///|
fn GridWorldEnv::state_label(self : GridWorldEnv, state : Int) -> String {
  let x = self.decode_x(state)
  let y = self.decode_y(state)
  "\{x},\{y}"
}

///|
fn GridWorldEnv::action_name(action : Int) -> String {
  match action {
    0 => "up"
    1 => "down"
    2 => "left"
    3 => "right"
    _ => "stay"
  }
}

///|
fn GridWorldEnv::reset(self : GridWorldEnv) -> Int {
  self.x = self.start_x
  self.y = self.start_y
  self._last_stage = Running
  self.encode(self.x, self.y)
}

///|
pub fn GridWorldEnv::actions(_self : GridWorldEnv) -> Array[Int] {
  [0, 1, 2, 3]
}

///|
pub fn GridWorldEnv::state_space(self : GridWorldEnv) -> Array[Int] {
  let size = self.width * self.height
  let states = Array::make(size, 0)
  let mut i = 0
  while i < size {
    states[i] = i
    i = i + 1
  }
  states
}

///|
fn GridWorldEnv::is_goal(self : GridWorldEnv) -> Bool {
  self.x == self.goal_x && self.y == self.goal_y
}

///|
pub fn GridWorldEnv::step(self : GridWorldEnv, action : Int) -> Transition {
  let from_state = self.encode(self.x, self.y)
  let mut next_x = self.x
  let mut next_y = self.y
  match action {
    0 => next_y = if next_y > 0 { next_y - 1 } else { next_y }
    1 => next_y = if next_y + 1 < self.height { next_y + 1 } else { next_y }
    2 => next_x = if next_x > 0 { next_x - 1 } else { next_x }
    3 => next_x = if next_x + 1 < self.width { next_x + 1 } else { next_x }
    _ => ()
  }
  self.x = next_x
  self.y = next_y
  let done = self.is_goal()
  self._last_stage = if done { Finished } else { Running }
  let reward = if done { 1.0 } else { -0.04 }
  Transition::{
    state: from_state,
    action,
    reward,
    next_state: self.encode(self.x, self.y),
    done,
    step: 0,
  }
}

///|
pub fn GridWorldEnv::render(self : GridWorldEnv) -> String {
  let mut buffer = "GridWorld \{self.width}x\{self.height}\n"
  let mut y = 0
  while y < self.height {
    let mut x = 0
    while x < self.width {
      let state = self.encode(x, y)
      let cell = if x == self.x && y == self.y {
        "A"
      } else if x == self.goal_x && y == self.goal_y {
        "G"
      } else {
        "."
      }
      buffer = buffer + cell + "(" + self.state_label(state) + ")"
      if x + 1 < self.width {
        buffer = buffer + " "
      }
      x = x + 1
    }
    buffer = buffer + "\n"
    y = y + 1
  }
  buffer
}

///|
pub struct QLearningAgent {
  table : QTable
  policy : EpsilonGreedyPolicy
  alpha : Double
  gamma : Double
} derive(Debug)

///|
pub fn QLearningAgent::new(
  states : Array[Int],
  actions : Array[Int],
  policy : EpsilonGreedyPolicy,
  alpha : Double,
  gamma : Double,
) -> QLearningAgent {
  {
    table: QTable::new(states.length(), actions.length()),
    policy,
    alpha,
    gamma,
  }
}

///|
pub fn QLearningAgent::reset_episode(_self : QLearningAgent) -> Unit {
  ()
}

///|
pub fn QLearningAgent::choose_action(self : QLearningAgent, state : Int) -> Int {
  self.policy.choose_action(state, [0, 1, 2, 3], self.table.row(state))
}

///|
pub fn QLearningAgent::learn(
  self : QLearningAgent,
  transition : Transition,
  _next_action : Int?,
) -> Unit {
  let target = if transition.done {
    transition.reward
  } else {
    transition.reward + self.gamma * self.table.max_value(transition.next_state)
  }
  self.table.update(transition.state, transition.action, target, self.alpha)
}

///|
pub fn QLearningAgent::epsilon(self : QLearningAgent) -> Double {
  self.policy.epsilon
}

///|
pub fn QLearningAgent::q_report(self : QLearningAgent, state : Int) -> String {
  let best = self.table.best_action_index(state)
  "\{state} -> \{best}"
}

///|
pub fn QLearningAgent::q_value(
  self : QLearningAgent,
  state : Int,
  action : Int,
) -> Double {
  self.table.value(state, action)
}

///|
pub fn QLearningAgent::set_q_value(
  self : QLearningAgent,
  state : Int,
  action : Int,
  value : Double,
) -> Unit {
  self.table.set_value(state, action, value)
}

///|
pub struct SARSAAgent {
  table : QTable
  policy : EpsilonGreedyPolicy
  alpha : Double
  gamma : Double
} derive(Debug)

///|
pub fn SARSAAgent::new(
  states : Array[Int],
  actions : Array[Int],
  policy : EpsilonGreedyPolicy,
  alpha : Double,
  gamma : Double,
) -> SARSAAgent {
  {
    table: QTable::new(states.length(), actions.length()),
    policy,
    alpha,
    gamma,
  }
}

///|
pub fn SARSAAgent::reset_episode(_self : SARSAAgent) -> Unit {
  ()
}

///|
pub fn SARSAAgent::choose_action(self : SARSAAgent, state : Int) -> Int {
  self.policy.choose_action(state, [0, 1, 2, 3], self.table.row(state))
}

///|
pub fn SARSAAgent::learn(
  self : SARSAAgent,
  transition : Transition,
  next_action : Int?,
) -> Unit {
  let next_value = match next_action {
    Some(action) => self.table.value(transition.next_state, action)
    None => 0.0
  }
  let target = if transition.done {
    transition.reward
  } else {
    transition.reward + self.gamma * next_value
  }
  self.table.update(transition.state, transition.action, target, self.alpha)
}

///|
pub fn SARSAAgent::epsilon(self : SARSAAgent) -> Double {
  self.policy.epsilon
}

///|
pub fn SARSAAgent::q_report(self : SARSAAgent, state : Int) -> String {
  let best = self.table.best_action_index(state)
  "\{state} -> \{best}"
}

///|
pub struct ConsoleLogger {
  prefix : String
} derive(Debug)

///|
pub fn ConsoleLogger::new(prefix : String) -> ConsoleLogger {
  { prefix, }
}

///|
fn ConsoleLogger::start_episode(self : ConsoleLogger, episode : Int) -> Unit {
  println("\{self.prefix} episode \{episode} start")
}

///|
fn ConsoleLogger::step(
  self : ConsoleLogger,
  episode : Int,
  state : Int,
  action : Int,
  reward : Double,
  next_state : Int,
  done : Bool,
  step : Int,
) -> Unit {
  if step <= 3 {
    println(
      "\{self.prefix} ep=\{episode} \{state} --\{GridWorldEnv::action_name(action)}/\{reward}--> \{next_state} done=\{done}",
    )
  }
}

///|
fn ConsoleLogger::finish_episode(
  self : ConsoleLogger,
  record : EpisodeRecord,
) -> Unit {
  println(
    "\{self.prefix} episode \{record.episode}: reward=\{record.reward}, steps=\{record.steps}, goal=\{record.goal_reached}",
  )
}

///|
fn ConsoleLogger::finish(self : ConsoleLogger, report : TrainingReport) -> Unit {
  println("\{self.prefix} training finished")
  println(report.summary())
}

///|
pub struct Trainer {
  episodes : Int
  max_steps : Int
} derive(Debug)

///|
pub fn Trainer::new(episodes : Int, max_steps : Int) -> Trainer {
  {
    episodes: if episodes < 0 {
      0
    } else {
      episodes
    },
    max_steps: if max_steps < 1 {
      1
    } else {
      max_steps
    },
  }
}

///|
pub fn Trainer::train_q_learning(
  self : Trainer,
  env : GridWorldEnv,
  agent : QLearningAgent,
  logger : ConsoleLogger,
) -> TrainingReport {
  let report = TrainingReport::{
    label: "Q-learning",
    episodes: self.episodes,
    rewards: Array::make(self.episodes, 0.0),
    steps: Array::make(self.episodes, 0),
    goal_hits: 0,
    final_epsilon: 0.0,
  }
  let mut episode = 0
  while episode < self.episodes {
    agent.reset_episode()
    let mut state = env.reset()
    logger.start_episode(episode + 1)
    let mut total_reward = 0.0
    let mut step_count = 0
    let mut done = false
    while step_count < self.max_steps && !done {
      let action = agent.choose_action(state)
      let transition = env.step(action)
      transition.step = step_count + 1
      let next_state = transition.next_state
      let done_flag = transition.done
      total_reward = total_reward + transition.reward
      logger.step(
        episode + 1,
        transition.state,
        transition.action,
        transition.reward,
        next_state,
        done_flag,
        transition.step,
      )
      agent.learn(transition, None)
      state = next_state
      done = done_flag
      step_count = step_count + 1
    }
    if done {
      report.goal_hits = report.goal_hits + 1
    }
    report.rewards[episode] = total_reward
    report.steps[episode] = step_count
    logger.finish_episode(EpisodeRecord::{
      episode: episode + 1,
      steps: step_count,
      reward: total_reward,
      goal_reached: done,
      final_state: state,
    })
    episode = episode + 1
  }
  report.final_epsilon = agent.epsilon()
  logger.finish(report)
  report
}

///|
pub fn Trainer::train_sarsa(
  self : Trainer,
  env : GridWorldEnv,
  agent : SARSAAgent,
  logger : ConsoleLogger,
) -> TrainingReport {
  let report = TrainingReport::{
    label: "SARSA",
    episodes: self.episodes,
    rewards: Array::make(self.episodes, 0.0),
    steps: Array::make(self.episodes, 0),
    goal_hits: 0,
    final_epsilon: 0.0,
  }
  let mut episode = 0
  while episode < self.episodes {
    agent.reset_episode()
    let mut state = env.reset()
    logger.start_episode(episode + 1)
    let mut action = agent.choose_action(state)
    let mut total_reward = 0.0
    let mut step_count = 0
    let mut done = false
    while step_count < self.max_steps && !done {
      let transition = env.step(action)
      transition.step = step_count + 1
      let next_state = transition.next_state
      let done_flag = transition.done
      total_reward = total_reward + transition.reward
      let next_action = if done_flag {
        None
      } else {
        Some(agent.choose_action(next_state))
      }
      logger.step(
        episode + 1,
        transition.state,
        transition.action,
        transition.reward,
        next_state,
        done_flag,
        transition.step,
      )
      agent.learn(transition, next_action)
      state = next_state
      done = done_flag
      step_count = step_count + 1
      action = match next_action {
        Some(next) => next
        None => 0
      }
    }
    if done {
      report.goal_hits = report.goal_hits + 1
    }
    report.rewards[episode] = total_reward
    report.steps[episode] = step_count
    logger.finish_episode(EpisodeRecord::{
      episode: episode + 1,
      steps: step_count,
      reward: total_reward,
      goal_reached: done,
      final_state: state,
    })
    episode = episode + 1
  }
  report.final_epsilon = agent.epsilon()
  logger.finish(report)
  report
}

///|
pub fn tutorial_blurb() -> String {
  "MoonRLLab combines finite environments, tabular control, and a small trainer so that new MoonBit contributors can inspect the whole learning loop in one place."
}