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

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

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

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

///|
pub fn ExpectedSARSAAgent::learn(
  self : ExpectedSARSAAgent,
  transition : Transition,
  _next_action : Int?,
) -> Unit {
  let mut expected_q = 0.0
  let next_state = transition.next_state

  if !transition.done {
    let next_q_values = self.table.row(next_state)
    let best_action = self.table.best_action_index(next_state)
    let action_count = 4.0
    let epsilon = self.policy.epsilon

    for a in 0..<4 {
      let prob = if a == best_action {
        1.0 - epsilon + epsilon / action_count
      } else {
        epsilon / action_count
      }
      expected_q += prob * next_q_values[a]
    }
  }

  let target = transition.reward + self.gamma * expected_q
  self.table.update(transition.state, transition.action, target, self.alpha)
}

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

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