// dqn_n_step.mbt — N-step DQN (Sutton 1988; Mnih et al. 2016 §4.2).
//
// 1-step TD target:    y_t = r_t + γ · max_a' Q̂(s_{t+1}, a')
// n-step TD target:    y_t = Σ_{k=0}^{n-1} γ^k · r_{t+k} + γ^n · max_a' Q̂(s_{t+n}, a')
//
// With truncation: if any transition r_{t+k} is terminal, the sum stops at k
// and γ^n · max ... is replaced by 0 (no bootstrap from a terminal future).
//
// We combine this with the standard (uniform) `ReplayBuffer` from `dqn.mbt`.
// Pairing with `PrioritizedReplayBuffer` is also possible — the priority can
// be the magnitude of the n-step TD error.
//
// Per-step pipeline (during a rollout):
//   1. nstep.push(s, a, r, s', done)
//   2. If nstep.len() == n OR done:
//        extract n-step transition G, h, s_root, s_lookahead
//        if (h == n) or (s_lookahead corresponds to truncated terminal):
//          replay.push(s_root, a_root, G, s_lookahead, done_flag)
//        if done: nstep.reset()
//   3. If replay.len() >= batch_size and ep >= warmup:
//        sample mini-batch
//        dqn_update_step_n_step(...)
//   4. Periodically copy online → target.

///|
/// N-step variant of `dqn_update_step`. Each row in the batch is an n-step
/// transition:
///
///   (s_root, a_root, G_n, s_lookahead, h_eff, use_bootstrap)
///
/// where:
///
/// * `G_n`            — discounted reward sum Σ γ^k r_{root+k}
/// * `s_lookahead`    — the s_{root + h_eff} used for bootstrap
/// * `h_eff`          — effective horizon (number of rewards summed, 1..=n)
/// * `use_bootstrap`  — true if `h_eff == n` and the n-step window is non-terminal
///
/// Loss for a row (with bootstrap):
///
///     δ = G_n + γ^h_eff · max_a' Q̂(s_lookahead, a') - Q(s_root, a_root)
///
/// Without bootstrap (terminal reached inside window):
///
///     δ = G_n - Q(s_root, a_root)
///
/// Returns mean |δ| over the batch.
pub fn dqn_update_step_n_step(
  online : LinearQNet,
  target : LinearQNet,
  states : Array[Int],
  actions : Array[Int],
  rewards : Array[Float], // G_n
  next_states : Array[Int], // s_lookahead
  h_effs : Array[Int], // effective horizons
  use_bootstraps : Array[Bool],
  gamma : Float,
  lr : Float,
) -> Float {
  let n = states.length()
  let mut total_abs_delta = 0.0F
  for k in 0.. Float {
  let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
  let buffer = ReplayBuffer::new(buffer_capacity)
  let nstep = NStepBuffer::new(n_step)
  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_n_step_via_replay(q_net, target, ss, aa, rr, nss, dd, gamma, lr, n_step)
    }
    // Sync target.
    if (ep + 1) % sync_every == 0 {
      qnet_copy(target, q_net)
    }
  }
  if return_count > 0 {
    total_return / Float::from_int(return_count)
  } else {
    0.0F
  }
}

///|
/// Adapter: replay buffer rows don't carry h_eff or use_bootstrap, but we can
/// reconstruct them when n == h_eff and the row's `done` flag is false. For
/// rows with done == true, h_eff is 1 and bootstrap is disabled.
///
/// This adapter assumes the training pipeline produced rows with h_eff = n_step
/// (no early truncation). The flag `done == true` is interpreted as
/// "do not bootstrap".
fn dqn_update_step_n_step_via_replay(
  online : LinearQNet,
  target : LinearQNet,
  states : Array[Int],
  actions : Array[Int],
  rewards : Array[Float],
  next_states : Array[Int],
  dones : Array[Bool],
  gamma : Float,
  lr : Float,
  n_step : Int,
) -> Float {
  let n = states.length()
  let h_arr : Array[Int] = Array::make(n, n_step)
  let use_arr : Array[Bool] = Array::make(n, true)
  for k in 0..