// HetRec — heterogeneous-timescale non-recurrent neuron population.
// Bit-exact port of SNNModels.jl/src/populations/hetrec.jl.
//
// A HetRec population is a set of N neurons, each with Nd dendritic
// compartments that filter input on heterogeneous timescales. Each
// soma reads from N*Nd dendritic compartments via a sparse mapping
// matrix M (soma i → dendrite (j-1)*N + k, with weight 1 if k==i
// or weight 1 with probability `overlap` otherwise). The soma then
// fires stochastically based on its baseline rate r[i] and the
// deviation of v_s[i] from a slow adaptation trace.
//
// Float32 contract: every arithmetic uses `Float` (Float32). The
// firing criterion uses libm `expf` via the formula
// `r * sigmoid(steepness * (v_s - trace)) * dt`.
//
// Storage layout (matches Julia):
//   v_d : Array[Float] size N*Nd   — dendritic voltages (filtered input)
//   v_s : Array[Float] size N      — somatic voltages
//   is  : Array[Float] size N*Nd   — synaptic input currents (set externally)
//   r   : Array[Float] size N      — per-neuron baseline firing rate
//   τd  : Array[Float] size N*Nd   — per-compartment dendritic time constants
//   fire   : Array[Bool]  size N   — spike indicators
//   tabs   : Array[Int]   size N   — refractory counters
//   trace  : Array[Float] size N   — slow adaptation trace
//   randcache : Array[Float] size N — random cache for stochastic firing
//   colptr : Array[Int]   size N+1 — column pointer for CSC of M'
//   I      : Array[Int]   size nnz — dendrite indices per synapse
//   W      : Array[Float] size nnz — synapse weights (= 1 in default M)
//
// We use a CSC-like representation directly (colptr + I + W) rather
// than CSR to match Julia's `dsparse(w)` output that the integrate!
// loop expects: `for s in colptr[i]:(colptr[i+1]-1)` iterates over
// the dendrites that soma i reads from.

///|
/// HetRecParameter — parameters for the HetRec layer.
pub(all) struct HetRecParameter {
  nd : Int              // dendritic compartments per neuron
  overlap : Float       // dendritic overlap across neurons [0, 1]
  tau_d_low : Float     // lower bound of τd distribution (ms)
  tau_d_high : Float    // upper bound of τd distribution (ms)
  rate_low : Float      // lower bound of r distribution (Hz)
  rate_high : Float     // upper bound of r distribution (Hz)
  tau_abs : Float       // absolute refractory period (ms)
  steepness : Float     // soma firing nonlinearity steepness
  tau_m : Float         // soma integration time constant (ms)
  tau_rate : Float      // adaptation trace time constant (ms)
}

///|
/// HetRecParameter with Julia's defaults.
pub fn HetRecParameter::new() -> HetRecParameter {
  // Julia: HetRecParameter() defaults:
  //   Nd=2, overlap=0.5, τd=Uniform(10, 100), rate=Uniform(0, 1),
  //   τabs=5ms, steepness=1, τm=20ms, τrate=100ms
  {
    nd: 2,
    overlap: 0.5F,
    tau_d_low: 10.0F,
    tau_d_high: 100.0F,
    rate_low: 0.0F,
    rate_high: 1.0F,
    tau_abs: 5.0F,
    steepness: 1.0F,
    tau_m: 20.0F,
    tau_rate: 100.0F,
  }
}

///|
/// Custom HetRecParameter with non-default Nd, τm, and overlap.
/// Mirrors Julia's `HetRecParameter(Nd = 4, τm = 30ms, overlap = 0.2)`.
/// Other fields keep their defaults.
pub fn HetRecParameter::custom(
  nd : Int,
  tau_m : Float,
  overlap : Float,
) -> HetRecParameter {
  {
    nd,
    overlap,
    tau_d_low: 10.0F,
    tau_d_high: 100.0F,
    rate_low: 0.0F,
    rate_high: 1.0F,
    tau_abs: 5.0F,
    steepness: 1.0F,
    tau_m,
    tau_rate: 100.0F,
  }
}

///|
/// HetRec — heterogeneous-timescale non-recurrent population.
pub struct HetRec {
  param : HetRecParameter
  n : Int                          // number of neurons
  v_d : Array[Float]               // dendritic voltages (size N*Nd)
  v_s : Array[Float]               // somatic voltages (size N)
  is_ : Array[Float]               // synaptic input currents (size N*Nd)
  r : Array[Float]                 // baseline firing rates (size N)
  tau_d : Array[Float]             // dendritic time constants (size N*Nd)
  fire : Array[Bool]               // spike indicators (size N)
  tabs : Array[Int]                // refractory counters (size N)
  trace : Array[Float]             // adaptation trace (size N)
  randcache : Array[Float]         // random cache (size N)
  // Sparse mapping M' (CSC layout: somas = columns, dendrites = rows).
  colptr : Array[Int]              // column pointer (size N+1)
  i_syn : Array[Int]               // dendrite index per synapse (size nnz)
  w_syn : Array[Float]             // weight per synapse (size nnz; binary 0/1)
}

///|
/// Construct a HetRec population. Initialises:
///   - v_d, v_s = 0
///   - is_ = 0
///   - r = Uniform(rate_low, rate_high) per neuron (in Hz; matches
///     Julia's `Uniform(0, 1)` so r values are in (0, 1)).
///   - τd = Uniform(tau_d_low, tau_d_high) per compartment (in ms).
///   - fire = false, tabs = 0, trace = 0
///   - randcache = uniform (will be re-filled each step)
///   - colptr / i_syn / w_syn: sparse M where
///       M[soma i, dendrite (j-1)*N + k] = 1 if k==i (own dendrite)
///                                  or 1 with probability `overlap`
///       M is then transposed to M' (dendrites × somas) and stored
///       as CSC with colptr = soma indices, i_syn = dendrite indices.
pub fn HetRec::new(n : Int, param : HetRecParameter, rng : Xoshiro) -> HetRec {
  let nd = param.nd
  let total_d = n * nd
  let v_d : Array[Float] = Array::make(total_d, 0.0F)
  let v_s : Array[Float] = Array::make(n, 0.0F)
  let is_ : Array[Float] = Array::make(total_d, 0.0F)
  // Sample baseline firing rates from Uniform(rate_low, rate_high).
  let r : Array[Float] = Array::make(n, 0.0F)
  for i in 0.. Unit {
  let n = p.n
  for i in 0.. 0, skip firing logic
///   4. update trace: trace += dt * (-trace / τrate)
///      if not refractory: trace += (v_s[i] - trace) / τrate
///   5. stochastic fire: if randcache[i] < r[i] * sigmoid(steepness * (v_s - trace)) * dt
///      then fire[i] = true; tabs[i] = τabs/dt; trace[i] += 1
///
/// The caller is expected to:
///   - inject synaptic currents into `p.is_` before calling step
///     (e.g., via SpikingSynapse forward to the `:is` target)
///   - call hetrec_refresh_random(p, rng) once per step to refresh
///     `p.randcache` (matches Julia's `rand!(randcache)`).
pub fn step_hetrec(p : HetRec, dt : Float) -> Unit {
  let n = p.n
  let nd = p.param.nd
  let total_d = n * nd
  let steepness = p.param.steepness
  let tau_m = p.param.tau_m
  let tau_rate = p.param.tau_rate
  let tau_abs = p.param.tau_abs
  let tabs_steps : Int = (tau_abs / dt).to_int()

  // 1. Dendritic Euler step.
  for i in 0.. 0 {
      continue
    }
    // trace catch-up to v_s
    p.trace[i] = p.trace[i] + (p.v_s[i] - p.trace[i]) / tau_rate
    // Stochastic fire: sigmoid(steepness * (v_s - trace)) = 1 / (1 + exp(-steepness * (v_s - trace)))
    let sigmoid_arg = -steepness * (p.v_s[i] - p.trace[i])
    let rate = if sigmoid_arg > 88.0F {
      // Saturating: exp(-large) ≈ 0 → sigmoid ≈ 1
      p.r[i] * dt
    } else if sigmoid_arg < -88.0F {
      // Saturating: exp(large) → Inf → sigmoid ≈ 0
      0.0F
    } else {
      p.r[i] * dt / (1.0F + expf(sigmoid_arg))
    }
    if p.randcache[i] < rate {
      p.fire[i] = true
      p.tabs[i] = tabs_steps
      p.trace[i] = p.trace[i] + 1.0F
    }
  }
}