// dqn.mbt — Deep Q-Network (v0.36.0).
//
// DQN (Mnih et al. 2015) on GridWorld:
//   - Q-network: Q(s, a; θ) = W · x_s  (linear, no hidden layer)
//   - Replay buffer: store (s, a, r, s', done), sample mini-batches
//   - Target network Q̂(s, a; θ⁻): periodic snapshot of θ
//   - Loss: (r + γ · max_a' Q̂(s', a'; θ⁻) · (1 - done) - Q(s, a; θ))²
//   - ε-greedy action selection with linear decay
//
// Reuses:
//   - `GridWorld` from reinforce.mbt
//   - `sample_categorical` from reinforce.mbt (for ε-greedy)

///|
/// Q-network: linear mapping from one-hot state to Q-values per
/// action. W has shape `n_actions × n_states`. Bias has shape `n_actions`.
pub struct LinearQNet {
  n_states : Int
  n_actions : Int
  w : Array[Array[Float]]
  b : Array[Float]
}

///|
pub fn LinearQNet::new(
  n_states : Int,
  n_actions : Int,
  seed : UInt64,
) -> LinearQNet {
  let std = sqrtf(0.1F / Float::from_int(n_states))
  let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let w : Array[Array[Float]] = Array::make(n_actions, [])
  for i in 0.. Array[Float] {
  let n_actions = net.n_actions
  let q : Array[Float] = Array::make(n_actions, 0.0F)
  for i in 0.. Int {
  let mut best = 0
  let mut best_v = q[0]
  for i in 1.. best_v {
      best_v = q[i]
      best = i
    }
  }
  best
}

///|
/// Experience replay buffer (uniform random sampling).
pub struct ReplayBuffer {
  capacity : Int
  states : Array[Int]
  actions : Array[Int]
  rewards : Array[Float]
  next_states : Array[Int]
  dones : Array[Bool]
  mut size : Int
  mut cursor : Int
}

///|
pub fn ReplayBuffer::new(capacity : Int) -> ReplayBuffer {
  {
    capacity,
    states: Array::make(capacity, 0),
    actions: Array::make(capacity, 0),
    rewards: Array::make(capacity, 0.0F),
    next_states: Array::make(capacity, 0),
    dones: Array::make(capacity, false),
    size: 0,
    cursor: 0,
  }
}

///|
/// Append a transition. Overwrites oldest when full.
pub fn ReplayBuffer::push(
  self : ReplayBuffer,
  s : Int,
  a : Int,
  r : Float,
  s_next : Int,
  done : Bool,
) -> Unit {
  let i = self.cursor
  self.states[i] = s
  self.actions[i] = a
  self.rewards[i] = r
  self.next_states[i] = s_next
  self.dones[i] = done
  self.cursor = (i + 1) % self.capacity
  if self.size < self.capacity {
    self.size = self.size + 1
  }
}

///|
pub fn ReplayBuffer::len(self : ReplayBuffer) -> Int {
  self.size
}

///|
/// Sample a mini-batch by random indices. Returns 5 parallel arrays.
pub fn ReplayBuffer::sample(
  self : ReplayBuffer,
  batch_size : Int,
  rng : Xoshiro,
) -> (Array[Int], Array[Int], Array[Float], Array[Int], Array[Bool]) {
  let states : Array[Int] = Array::make(batch_size, 0)
  let actions : Array[Int] = Array::make(batch_size, 0)
  let rewards : Array[Float] = Array::make(batch_size, 0.0F)
  let next_states : Array[Int] = Array::make(batch_size, 0)
  let dones : Array[Bool] = Array::make(batch_size, false)
  for k in 0..= self.size {
      self.size - 1
    } else {
      idx
    }
    states[k] = self.states[safe_idx]
    actions[k] = self.actions[safe_idx]
    rewards[k] = self.rewards[safe_idx]
    next_states[k] = self.next_states[safe_idx]
    dones[k] = self.dones[safe_idx]
  }
  (states, actions, rewards, next_states, dones)
}

///|
/// Copy weights from `src` into `dst`. Used to sync target net from
/// online net.
pub fn qnet_copy(dst : LinearQNet, src : LinearQNet) -> Unit {
  for i in 0.. Float {
  let n = states.length()
  let mut total_abs_delta = 0.0F
  for k in 0.. Float {
  let mut m = q[0]
  for i in 1.. m {
      m = q[i]
    }
  }
  m
}

///|
/// ε-greedy action. With probability ε, sample uniformly; else greedy.
pub fn eps_greedy_action(
  q_net : LinearQNet,
  state : Int,
  epsilon : Float,
  rng : Xoshiro,
) -> Int {
  if epsilon <= 0.0F {
    let x = rl_one_hot(state, q_net.n_states)
    let q = q_forward(q_net, x)
    return q_argmax(q)
  }
  let (u, _) = box_muller(rng)
  let u_pos = if u < 0.0 { -u } else { u }
  if Float::from_double(u_pos) < epsilon {
    // Uniform random.
    let (u2, _) = box_muller(rng)
    let n_a = q_net.n_actions
    let idx_raw = if u2 < 0.0 { -u2 } else { u2 }
    let idx = Float::from_double(idx_raw * Double::from_int(n_a)).to_int()
    if idx < 0 {
      0
    } else if idx >= n_a {
      n_a - 1
    } else {
      idx
    }
  } else {
    let x = rl_one_hot(state, q_net.n_states)
    let q = q_forward(q_net, x)
    q_argmax(q)
  }
}

///|
/// Train DQN for `n_episodes` episodes. Each episode:
///   1. Roll out under ε-greedy policy with linear ε-decay.
///   2. Store transitions in replay buffer.
///   3. After warmup, sample mini-batches and apply gradient step.
///   4. Periodically copy online → target.
///
/// Returns the mean undiscounted episode return over the last 10
/// episodes.
pub fn train_dqn(
  env : GridWorld,
  q_net : LinearQNet,
  target : LinearQNet,
  n_episodes : Int,
  gamma : Float,
  lr : Float,
  epsilon_start : Float,
  epsilon_end : Float,
  max_steps : Int,
  buffer_capacity : Int,
  warmup_episodes : Int,
  sync_every : Int,
  batch_size : Int,
  seed : UInt64,
) -> Float {
  let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
  let buffer = ReplayBuffer::new(buffer_capacity)
  let n_states = env.n_states()
  let mut total_return = 0.0F
  let mut return_count = 0
  for ep in 0..= warmup_episodes && buffer.len() >= batch_size {
      let (ss, aa, rr, nss, dd) = buffer.sample(batch_size, rng)
      let _ = dqn_update_step(q_net, target, ss, aa, rr, nss, dd, gamma, lr)
    }
    // Sync target net.
    if (ep + 1) % sync_every == 0 {
      qnet_copy(target, q_net)
    }
    let _ = n_states
  }
  if return_count > 0 {
    total_return / Float::from_int(return_count)
  } else {
    0.0F
  }
}