// Sparse matrix (CSR format) — bit-exact port of SNNModels.jl/src/utils/sparse_matrix.jl.
//
// The Julia source uses SparseMatrixCSC (compressed sparse column)
// but we use CSR (compressed sparse row) here for simpler pre-synaptic
// iteration: for each pre-synaptic neuron, iterate over its
// outgoing connections.
//
// Julia reference (src/utils/sparse_matrix.jl):
//   sparse_matrix(post, pre, μ, σ, p) = ...  → W, I, J, index, rowptr
//   dsparse(W, I, J)                       → SparseMatrixCSC{Float32, Int}
//   update_sparse_matrix!(c, W)            → refresh from dense matrix
//   connect!(c, j, i, μ)                   → set (i,j) weight then rebuild
//
// For our CSR layout we store (rowptr, colptr, vals) per pre-syn row.
// `connect!(c, post, pre, w)` mutates the dense weight buffer and
// then rebuilds the CSR.

///|
/// Sparse matrix in compressed-sparse-row format.
pub struct SparseMatrixCSR {
  rows : Int
  cols : Int
  // rowptr[i] = start index in colptr/vals of row i's first non-zero.
  // Length = rows + 1.
  rowptr : Array[Int]
  // Column index of each non-zero. Length = nnz.
  colptr : Array[Int]
  // Value of each non-zero. Length = nnz.
  vals : Array[Float]
}

///|
/// Number of stored non-zeros.
pub fn SparseMatrixCSR::nnz(m : SparseMatrixCSR) -> Int {
  m.vals.length()
}

///|
/// Build an empty N×M sparse matrix.
pub fn SparseMatrixCSR::empty(rows : Int, cols : Int) -> SparseMatrixCSR {
  let rowptr : Array[Int] = Array::make(rows + 1, 0)
  { rows, cols, rowptr, colptr: [], vals: [] }
}

///|
/// Build a CSR sparse matrix from a dense Float32 matrix.
/// Density = fraction of non-zero entries preserved.
pub fn SparseMatrixCSR::from_dense(
  dense : Array[Array[Float]],
  threshold : Float,
) -> SparseMatrixCSR {
  let rows = dense.length()
  let cols = if rows > 0 { dense[0].length() } else { 0 }
  let rowptr : Array[Int] = Array::make(rows + 1, 0)
  let colptr : Array[Int] = []
  let vals : Array[Float] = []
  for i in 0.. threshold {
        colptr.push(j)
        vals.push(v)
      }
    }
  }
  rowptr[rows] = vals.length()
  { rows, cols, rowptr, colptr, vals }
}

///|
/// Connection rules for random sparse matrices. Mirrors Julia SNN's
/// `rule` keyword on `sparse_matrix(post, pre, μ, σ, p; rule=...)`:
///   - `Bernoulli`: each edge independently kept with probability `p`
///   - `FixedIn`  : each post-synaptic neuron receives from exactly
///                  `p * Npre` pre-synaptic neurons (chosen uniformly
///                  without replacement). Same in-degree for all posts.
///   - `FixedOut` : each pre-synaptic neuron projects to exactly
///                  `p * Npost` post-synaptic neurons. Same out-degree
///                  for all pres.
pub(all) enum ConnectRule {
  Bernoulli
  FixedIn
  FixedOut
}

///|
/// Build a sparse matrix with random Normal(μ, σ) weights and
/// connection probability `p`. Equivalent to Julia's
/// `sparse_matrix(post, pre, μ, σ, p)`.
///
/// `rng` is the Xoshiro stream. The Float32 draw path is bit-exact
/// with Julia's `rand(rng, Float32)` for the same seed.
pub fn SparseMatrixCSR::random(
  rows : Int,
  cols : Int,
  mu : Float,
  sigma : Float,
  p : Float,
  rng : Xoshiro,
) -> SparseMatrixCSR {
  // Default rule is Bernoulli for backward compatibility.
  SparseMatrixCSR::random_with_rule(rows, cols, mu, sigma, p, Bernoulli, rng)
}

///|
/// Build a sparse matrix with the given connection rule. See
/// `ConnectRule` for the semantics of each rule.
pub fn SparseMatrixCSR::random_with_rule(
  rows : Int,
  cols : Int,
  mu : Float,
  sigma : Float,
  p : Float,
  rule : ConnectRule,
  rng : Xoshiro,
) -> SparseMatrixCSR {
  // Build the dense weight matrix first. We use Float for weights
  // (matches Julia's Float32 default).
  let dense : Array[Array[Float]] = Array::make(rows, [])
  for i in 0.. {
      // Each weight independently zeroed with prob (1 - p).
      for i in 0..= p {
            dense[i][j] = 0.0F
          }
        }
      }
    }
    FixedIn => {
      // For each post j, keep exactly `round(p*rows)` pre-synaptic
      // neurons (chosen uniformly without replacement). The rest go
      // to 0.
      let n_keep = (Float::from_int(rows) * p).to_int()
      if n_keep > 0 && n_keep <= rows {
        for j in 0..= rows { rows - 1 } else { r_idx }
            // Swap pre_idx[k] and pre_idx[r_clamped]
            let tmp = pre_idx[k]
            pre_idx[k] = pre_idx[r_clamped]
            pre_idx[r_clamped] = tmp
          }
          for k in 0.. {
      // For each pre i, keep exactly `round(p*cols)` post-synaptic
      // neurons.
      let n_keep = (Float::from_int(cols) * p).to_int()
      if n_keep > 0 && n_keep <= cols {
        for i in 0..= cols { cols - 1 } else { r_idx }
            let tmp = post_idx[k]
            post_idx[k] = post_idx[r_clamped]
            post_idx[r_clamped] = tmp
          }
          for k in 0.. Float {
  if i < 0 || i >= m.rows {
    return 0.0F
  }
  let start = m.rowptr[i]
  let end = m.rowptr[i + 1]
  // Linear scan within row i (rows are typically short)
  for k in start.. Unit {
  if i < 0 || i >= m.rows {
    return
  }
  let start = m.rowptr[i]
  let end = m.rowptr[i + 1]
  // Look for existing entry
  let mut found = false
  for k in start.. Unit {
  let rows = m.rows
  for i in 0.. Unit {
  let rows = m.rows
  for i in 0.. Int {
  if i < 0 || i >= m.rows {
    return 0
  }
  m.rowptr[i + 1] - m.rowptr[i]
}