// neuron_multipod.mbt — variable-dendrite-count Tripod (Multipod).
//
// Julia reference:
//   SNNModels.jl/src/populations/multicompartment/multipod.jl
//
// Provides:
//   - MultipodParameter — wraps Array[Dendrite] (the existing
//     Dendrite struct from neuron_dendrite.mbt is reused).
//   - Multipod struct (n, nd, v_s, w_s, v_d, ge_s, gi_s, he_s, hi_s,
//     he_d, hi_d, g_d, h_d, fire, after_spike, postspike, theta,
//     dv, dv_temp, cs, iv, soma_syn, dend_syn, glu_receptors,
//     gaba_receptors, alpha).
//   - Multipod::new(dendrites, n, param) — allocates all per-neuron +
//     per-dendrite arrays; initial v_s / v_d[k] in [vr, vt].
//   - Multipod::step(p, dt) — Euler-style update (he_d → h_d → g_d,
//     soma_syn, w_s, v_s, v_d[k], fire/after_spike reset).
//
// Float32 contract: every arithmetic uses Float32.
//
// **Note**: this is a simplified port. The Julia `Multipod` includes
// full NMDA multi-receptor + ReceptorArray wiring + Heun corrector.
// We provide a minimum-viable Euler step + receptor buffers (4 per
// dendrite: AMPA / NMDA / GABAa / GABAb) so the per-dendrite
// per-receptor structure is in place. Heun + NMDA voltage gating are
// TODO.

// =========================================================================
// MultipodParameter
// =========================================================================

// Dendrite is defined in neuron_dendrite.mbt — we reuse it.

///|
/// MultipodParameter — wraps an Array[Dendrite] for the multipod
/// constructor. The Julia version uses Dendrite objects directly.
pub struct MultipodParameter {
  dendrites : Array[Dendrite]
}

///|
pub fn MultipodParameter::new(dendrites : Array[Dendrite]) -> MultipodParameter {
  { dendrites: dendrites }
}

///|
/// Convenience: MultipodParameter from a single Dendrite applied
/// `nd` times. Mirrors Julia's `Multipod(d; N, Nd)` constructor.
pub fn MultipodParameter::uniform(
  d : Dendrite,
  nd : Int,
) -> MultipodParameter {
  let ds : Array[Dendrite] = Array::make(nd, d)
  { dendrites: ds }
}

// =========================================================================
// Multipod struct
// =========================================================================

///|
/// Multipod — variable-dendrite-count multi-compartment AdEx neuron.
/// `v_d` is a Array[Array[Float]] (one inner array per dendrite,
/// each of length N). `g_d` is a flat Float array indexed as
/// `g_d[i, d, n]` via `g_d[i + d*N + n*N*nd]`.
pub struct Multipod {
  n : Int
  nd : Int
  // Soma
  v_s : Array[Float]
  w_s : Array[Float]
  // Dendrite voltages (length nd; each inner array length N).
  v_d : Array[Array[Float]]
  // Soma conductances + spike-input buffers.
  ge_s : Array[Float]
  gi_s : Array[Float]
  he_s : Array[Float]
  hi_s : Array[Float]
  // Per-dendrite spike-input buffers (length nd; each inner length N).
  he_d : Array[Array[Float]]
  hi_d : Array[Array[Float]]
  // Per-dendrite per-receptor synaptic conductance + spike-input
  // state. Nested 3-level per-dendrite per-neuron: g_d[d][i][n] for
  // receptor n (0=AMPA, 1=NMDA, 2=GABAa, 3=GABAb). The 3-level
  // nesting avoids the flat-index nested loops that previously
  // triggered stack overflows in MoonBit's JIT.
  g_d : Array[Array[Array[Float]]]
  h_d : Array[Array[Array[Float]]]
  // Receptor markers.
  glu_receptors : Array[Int]
  gaba_receptors : Array[Int]
  // Alpha scaling for each of the 4 receptors.
  alpha : Array[Float]
  // Per-neuron spike bookkeeping.
  fire : Array[Bool]
  after_spike : Array[Int]
  postspike : PostSpike
  // Threshold + adaptation state.
  theta : Array[Float]
  // Heun scratch buffers (length nd + 1).
  dv : Array[Float]
  dv_temp : Array[Float]
  // Per-compartment axial + synaptic-current buffers.
  cs : Array[Float]
  iv : Array[Float]
  // Synapse parameter arrays (simplified — Julia uses Receptors).
  soma_syn : SingleExpSynapse
  dend_syn : SingleExpSynapse
  // NMDA voltage-dependency.
  nmda : NMDAVoltageDependency
  // Dendrite parameters.
  dendrites : Array[Dendrite]
}

///|
/// Allocate a Multipod. `dendrites` is the per-dendrite parameter
/// list (length = nd). `n` is the number of neurons. `param` is the
/// AdEx soma parameters (defaults to AdExParameter::new()).
pub fn Multipod::new(
  dendrites : Array[Dendrite],
  n : Int,
  param? : AdExParameter = AdExParameter::new(),
) -> Multipod {
  let nd = dendrites.length()
  let vt = param.vt
  let vr = param.vr
  let range = vt - vr
  // Initial v_s[i] in [vr, vt] uniformly (deterministic).
  let v_s : Array[Float] = Array::make(n, vr)
  let mut i = 0
  while i < n {
    v_s[i] = vr + (i % 3).to_float() * range / 3.0F
    i = i + 1
  }
  let w_s : Array[Float] = Array::make(n, 0.0F)
  // Dendrite voltages (length nd; each inner array length N).
  let v_d : Array[Array[Float]] = Array::make(nd, [vr])
  let mut d = 0
  while d < nd {
    let inner : Array[Float] = Array::make(n, vr + (d % 2).to_float() * range / 2.0F)
    ignore(v_d.set(d, inner))
    d = d + 1
  }
  // Spike-input buffers.
  let he_s : Array[Float] = Array::make(n, 0.0F)
  let hi_s : Array[Float] = Array::make(n, 0.0F)
  let ge_s : Array[Float] = Array::make(n, 0.0F)
  let gi_s : Array[Float] = Array::make(n, 0.0F)
  let he_d : Array[Array[Float]] = Array::make(nd, [0.0F])
  let hi_d : Array[Array[Float]] = Array::make(nd, [0.0F])
  let mut d2 = 0
  while d2 < nd {
    let inner_e : Array[Float] = Array::make(n, 0.0F)
    let inner_i : Array[Float] = Array::make(n, 0.0F)
    ignore(he_d.set(d2, inner_e))
    ignore(hi_d.set(d2, inner_i))
    d2 = d2 + 1
  }
  // Per-receptor per-dendrite conductance + spike-input state.
  // Nested per-dendrite per-neuron: g_d[d][i][n] for receptor n
  // (0=AMPA, 1=NMDA, 2=GABAa, 3=GABAb). The 3-level nesting avoids
  // the flat-index nested loops that previously triggered stack
  // overflows in MoonBit's JIT.
  let g_d : Array[Array[Array[Float]]] = Array::make(nd, [[0.0F]])
  let h_d : Array[Array[Array[Float]]] = Array::make(nd, [[0.0F]])
  let mut d3 = 0
  while d3 < nd {
    let g_inner : Array[Array[Float]] = Array::make(n, [0.0F])
    let h_inner : Array[Array[Float]] = Array::make(n, [0.0F])
    let mut i3 = 0
    while i3 < n {
      let g_neuron : Array[Float] = Array::make(4, 0.0F)
      let h_neuron : Array[Float] = Array::make(4, 0.0F)
      ignore(g_inner.set(i3, g_neuron))
      ignore(h_inner.set(i3, h_neuron))
      i3 = i3 + 1
    }
    ignore(g_d.set(d3, g_inner))
    ignore(h_d.set(d3, h_inner))
    d3 = d3 + 1
  }
  // Receptor markers.
  let glu_receptors : Array[Int] = [1, 2]
  let gaba_receptors : Array[Int] = [3, 4]
  let alpha : Array[Float] = [1.0F, 1.0F, 1.0F, 1.0F]
  // Spike bookkeeping.
  let fire : Array[Bool] = Array::make(n, false)
  let after_spike : Array[Int] = Array::make(n, 0)
  let postspike = PostSpike::new()
  let theta : Array[Float] = Array::make(n, param.vt)
  // Heun scratch.
  let dv : Array[Float] = Array::make(nd + 1, 0.0F)
  let dv_temp : Array[Float] = Array::make(nd + 1, 0.0F)
  let cs : Array[Float] = Array::make(nd, 0.0F)
  let iv : Array[Float] = Array::make(nd + 1, 0.0F)
  // Synapse parameters (SingleExpSynapse defaults).
  let soma_syn = SingleExpSynapse::new()
  let dend_syn = SingleExpSynapse::new()
  // NMDA voltage-dependency (Eyal defaults).
  let nmda = NMDAVoltageDependency::eyal()
  {
    n,
    nd,
    v_s,
    w_s,
    v_d,
    ge_s,
    gi_s,
    he_s,
    hi_s,
    he_d,
    hi_d,
    g_d,
    h_d,
    glu_receptors,
    gaba_receptors,
    alpha,
    fire,
    after_spike,
    postspike,
    theta,
    dv,
    dv_temp,
    cs,
    iv,
    soma_syn,
    dend_syn,
    nmda,
    dendrites,
  }
}

///|
/// Step the multipod forward one dt — Euler-style update (no Heun
/// corrector). Updates soma_syn, dend_syn receptors, w_s, v_s, v_d,
/// then spike detection.
pub fn Multipod::step(p : Multipod, dt : Float) -> Unit {
  let n = p.n
  let nd = p.nd
  let soma_decay_d : Float = expf(-dt / p.soma_syn.tau_e)
  let soma_decay_r : Float = expf(-dt / p.soma_syn.tau_e)
  let dend_decay_d : Float = expf(-dt / p.dend_syn.tau_e)
  let dend_decay_r : Float = expf(-dt / p.dend_syn.tau_e)
  // Phase 1: per-dendrite, per-receptor h update from he_d / hi_d.
  let mut d_idx = 0
  while d_idx < nd {
    let inner_e = p.he_d[d_idx]
    let inner_i = p.hi_d[d_idx]
    let mut i = 0
    while i < n {
      // Update glu_receptors [0, 1] (h_d[d][i][0] + he_d[i] * alpha[0],
      // h_d[d][i][1] + he_d[i] * alpha[1]).
      p.h_d[d_idx][i][0] = p.h_d[d_idx][i][0] + inner_e[i] * p.alpha[0]
      p.h_d[d_idx][i][1] = p.h_d[d_idx][i][1] + inner_e[i] * p.alpha[1]
      // Update gaba_receptors [2, 3].
      p.h_d[d_idx][i][2] = p.h_d[d_idx][i][2] + inner_i[i] * p.alpha[2]
      p.h_d[d_idx][i][3] = p.h_d[d_idx][i][3] + inner_i[i] * p.alpha[3]
      i = i + 1
    }
    d_idx = d_idx + 1
  }
  // Phase 2: reset he_d, hi_d.
  d_idx = 0
  while d_idx < nd {
    let mut i = 0
    while i < n {
      p.he_d[d_idx][i] = 0.0F
      p.hi_d[d_idx][i] = 0.0F
      i = i + 1
    }
    d_idx = d_idx + 1
  }
  // Phase 3: update soma_syn (ge_s, he_s, gi_s, hi_s).
  let mut i = 0
  while i < n {
    p.ge_s[i] = soma_decay_d * (p.ge_s[i] + dt * p.he_s[i])
    p.he_s[i] = soma_decay_r * p.he_s[i]
    p.gi_s[i] = soma_decay_d * (p.gi_s[i] + dt * p.hi_s[i])
    p.hi_s[i] = soma_decay_r * p.hi_s[i]
    i = i + 1
  }
  // Phase 4: update per-dendrite g_d, h_d for each of 4 receptors.
  d_idx = 0
  while d_idx < nd {
    let mut i = 0
    while i < n {
      let mut k = 0
      while k < 4 {
        p.g_d[d_idx][i][k] = dend_decay_d *
          (p.g_d[d_idx][i][k] + dt * p.h_d[d_idx][i][k])
        p.h_d[d_idx][i][k] = dend_decay_r * p.h_d[d_idx][i][k]
        k = k + 1
      }
      i = i + 1
    }
    d_idx = d_idx + 1
  }
  // Phase 5: per-neuron state update.
  multipod_step_neurons(p, dt, n, nd)
}

// =========================================================================
// Per-neuron update — extracted to flatten the call stack.
// =========================================================================

///|
fn multipod_step_neurons(p : Multipod, dt : Float, n : Int, nd : Int) -> Unit {
  let c = 281.0F
  let gl = 40.0F
  let dt_slope = 2.0F
  let el = -70.6F
  let a : Float = 4.0F
  let b : Float = 80.5F
  let tw : Float = 144.0F
  let mut i = 0
  while i < n {
    // Compute cs[d] = -((v_d[d][i] - v_s[i]) * gax[i]).
    let mut d_idx = 0
    while d_idx < nd {
      p.cs[d_idx] = -(p.v_d[d_idx][i] - p.v_s[i]) * p.dendrites[d_idx].gax[i]
      d_idx = d_idx + 1
    }
    // Compute iv[0] = soma synaptic current.
    p.iv[0] = p.soma_syn.gsyn_e * p.ge_s[i] *
      (p.v_s[i] - p.soma_syn.e_e) +
      p.soma_syn.gsyn_i * p.gi_s[i] * (p.v_s[i] - p.soma_syn.e_i)
    // Compute iv[d+1] = sum of dendrite synaptic currents.
    d_idx = 0
    while d_idx < nd {
      p.iv[d_idx + 1] = p.g_d[d_idx][i][0] +
        p.g_d[d_idx][i][1] +
        p.g_d[d_idx][i][2] +
        p.g_d[d_idx][i][3]
      d_idx = d_idx + 1
    }
    // Sum cs.
    let cs_sum = multipod_sum_cs(p, nd)
    // Δv[0] for soma.
    let theta_i = p.theta[i]
    let dv_soma : Float = (gl *
      ((-p.v_s[i] + el) + dt_slope * expf((p.v_s[i] - theta_i) / dt_slope)) -
      p.w_s[i] - p.iv[0] - cs_sum) / c
    p.v_s[i] = p.v_s[i] + dt * dv_soma
    // Δv[d+1] for each dendrite.
    d_idx = 0
    while d_idx < nd {
      let dd = p.dendrites[d_idx]
      let dv_d : Float = (-(p.v_d[d_idx][i] - el) * dd.gm[i] - p.iv[d_idx + 1] + p.cs[d_idx]) / dd.c[i]
      p.v_d[d_idx][i] = p.v_d[d_idx][i] + dt * dv_d
      d_idx = d_idx + 1
    }
    // w_s[i] += dt * (a * (v_s - El) - w_s) / τw.
    p.w_s[i] = p.w_s[i] + dt * (a * (p.v_s[i] - el) - p.w_s[i]) / tw
    // Reset fire + update threshold.
    p.fire[i] = false
    p.theta[i] = p.theta[i] - dt * (p.theta[i] - p.theta[i]) / p.postspike.tabs_const
    p.after_spike[i] = p.after_spike[i] - 1
    // Spike detection.
    if p.after_spike[i] < 0 {
      if p.v_s[i] > p.theta[i] + 10.0F {
        p.fire[i] = true
        p.theta[i] = p.theta[i] + 10.0F
        p.v_s[i] = 10.0F
        p.w_s[i] = p.w_s[i] + b
      }
    }
    i = i + 1
  }
}

///|
fn multipod_sum_cs(p : Multipod, nd : Int) -> Float {
  let mut s = 0.0F
  let mut d_idx = 0
  while d_idx < nd {
    s = s + p.cs[d_idx]
    d_idx = d_idx + 1
  }
  s
}