// trajectory_buffer.mbt — TrajectoryBuffer: offline RL replay storing
// full (state, action, reward, done) trajectories with returns-to-go
// computation and window batching (v0.78.0).
//
// Reference: Chen et al. 2021 "Decision Transformer" offline-RL
// pipeline. The DT trainer (v0.79.0) consumes this buffer's
// `sample_trajectory_window` output and feeds it into `dt_forward`.
//
// Layout:
//   Each trajectory is a row-major flat Float buffer plus an Int length.
//   The buffer holds up to `capacity` trajectories, each of arbitrary
//   length (different trajectories can have different lengths).
//
//   - state_seq[t * state_dim + k]    = state value
//   - action_seq[t * action_dim + k]  = action value
//   - rewards[t]                      = scalar reward
//   - dones[t]                        = 0.0F or 1.0F
//   - rtg_seq[t]                      = returns-to-go at t = sum_{t'=t}^{T-1} γ^{t'-t} r_{t'}
//   - timesteps[t]                    = absolute step index in trajectory (for
//                                       DT timestep embedding)
//
// API:
//   TrajectoryBuffer::new(capacity, state_dim, action_dim, gamma)
//   load_trajectory(state_seq, action_seq, reward_seq, done_seq, traj_len)
//   num_trajectories() -> Int
//   sample_trajectory_window(batch_size, seq_len, rng)
//     -> (rtg_batch, state_batch, action_batch, timesteps_batch, seq_len)
//        shapes:
//          rtg_batch       [batch_size × seq_len]
//          state_batch     [batch_size × seq_len × state_dim]
//          action_batch    [batch_size × seq_len × action_dim]
//          timesteps_batch [batch_size × seq_len]   (Int)
//   compute_returns_to_go(reward_seq, done_seq, traj_len) -> rtg_seq
//   reset() -> Unit

///|
/// One trajectory: a flat row-major Float buffer plus its length.
pub struct Trajectory {
  state_dim : Int
  action_dim : Int
  state_seq : Array[Float]
  action_seq : Array[Float]
  reward_seq : Array[Float]
  done_seq : Array[Float]
  rtg_seq : Array[Float]
  length : Int
}

///|
/// Offline trajectory buffer. Holds up to `capacity` trajectories,
/// each with its own length. Returns-to-go are computed once per
/// trajectory at load time.
pub(all) struct TrajectoryBuffer {
  capacity : Int
  state_dim : Int
  action_dim : Int
  gamma : Float
  mut trajectories : Array[Trajectory]
  mut num_loaded : Int
}

///|
/// Construct a fresh TrajectoryBuffer. `gamma` is the discount used in
/// returns-to-go (typically 1.0 for DT, since DT does not bootstrap
/// over its own value function — the returns themselves encode the
/// horizon).
pub fn TrajectoryBuffer::new(
  capacity : Int,
  state_dim : Int,
  action_dim : Int,
  gamma : Float,
) -> TrajectoryBuffer {
  { capacity, state_dim, action_dim, gamma, trajectories: [], num_loaded: 0 }
}

///|
/// Compute returns-to-go for a single trajectory:
///   R_t = r_t + γ · R_{t+1}     (if t < T - 1)
///   R_{T-1} = r_{T-1}            (terminal)
/// For trajectory boundaries with done=1, the next-step return is zero
/// (the episode ended — no further rewards). Returns an Array[Float]
/// of length `traj_len`.
pub fn compute_returns_to_go(
  reward_seq : Array[Float],
  done_seq : Array[Float],
  traj_len : Int,
  gamma : Float,
) -> Array[Float] {
  let rtg : Array[Float] = Array::make(traj_len, 0.0F)
  if traj_len == 0 {
    return rtg
  }
  // Walk from end to start.
  let mut running = 0.0F
  for t_rev in 0.. 0.5F {
      // Terminal: bootstrap from zero (episode ended).
      running = r
    } else {
      running = r + gamma * running
    }
    rtg[t] = running
  }
  rtg
}

///|
/// Load a single trajectory into the buffer. Computes returns-to-go
/// at load time (caller passes the raw rewards). If the buffer is full,
/// the trajectory is silently dropped (returns false).
pub fn load_trajectory(
  buf : TrajectoryBuffer,
  state_seq : Array[Float],
  action_seq : Array[Float],
  reward_seq : Array[Float],
  done_seq : Array[Float],
  traj_len : Int,
) -> Bool {
  if buf.num_loaded >= buf.capacity {
    return false
  }
  let rtg = compute_returns_to_go(reward_seq, done_seq, traj_len, buf.gamma)
  let traj : Trajectory = {
    state_dim: buf.state_dim,
    action_dim: buf.action_dim,
    state_seq,
    action_seq,
    reward_seq,
    done_seq,
    rtg_seq: rtg,
    length: traj_len,
  }
  buf.trajectories.push(traj)
  buf.num_loaded = buf.num_loaded + 1
  true
}

///|
/// Number of trajectories currently loaded.
pub fn trajectory_buffer_num_loaded(buf : TrajectoryBuffer) -> Int {
  buf.num_loaded
}

///|
/// Sample a batch of trajectory windows. Each sample is a contiguous
/// `seq_len` slice starting at a random position within a random trajectory.
/// If the chosen trajectory is shorter than seq_len, the window is
/// zero-padded on the right (the timesteps reflect the actual position).
/// Returns:
///   - rtg_batch:       [batch_size × seq_len]
///   - state_batch:     [batch_size × seq_len × state_dim]
///   - action_batch:    [batch_size × seq_len × action_dim]
///   - timesteps_batch: [batch_size × seq_len]   (Int)
/// All four are newly-allocated Float / Int arrays.
pub fn sample_trajectory_window(
  buf : TrajectoryBuffer,
  batch_size : Int,
  seq_len : Int,
  rng : Xoshiro,
) -> (Array[Float], Array[Float], Array[Float], Array[Int]) {
  let rtg_batch : Array[Float] = Array::make(batch_size * seq_len, 0.0F)
  let state_batch : Array[Float] = Array::make(
    batch_size * seq_len * buf.state_dim, 0.0F,
  )
  let action_batch : Array[Float] = Array::make(
    batch_size * seq_len * buf.action_dim, 0.0F,
  )
  let timesteps_batch : Array[Int] = Array::make(batch_size * seq_len, 0)
  for b in 0.. 0 {
      traj.length - 1
    } else {
      0
    }
    let start = if max_start == 0 {
      0
    } else {
      (next_u64(rng) % (max_start + 1).to_uint64()).to_int()
    }
    // Copy seq_len steps; timesteps are absolute indices (used by DT
    // timestep embedding).
    for t in 0.. Unit {
  buf.trajectories = []
  buf.num_loaded = 0
}