// neuron_tripod_het.mbt — Tripod with heterogeneous per-neuron
// AdEx parameters. Mirrors AdExHet (v0.10.9) and Julia's
// `Tripod(..., NMDA=EyalNMDA, param=AdExParameter{Vector{Float32}})`.
//
// Per-neuron fields: vt/vr/el/tm/r/dt_slope/tw/a/b/c/gl. Dendrite
// geometry (d1/d2) stays shared — Julia's TripodHet uses a single
// Dendrite struct across the population. Heun predictor-corrector
// integration reads per-ne parameters inside the per-neuron loop.
//
// Bit-exact note: this matches Julia's `Tripod{Vector{Float32}}`
// behavior. The soma params vary per neuron but dendrite params +
// AdExPostSpike + Heun correction order are shared across neurons
// (same as Julia).

///|
/// TripodHet — Tripod with heterogeneous per-neuron AdEx soma params.
/// Same compartments (:soma + :d1 + :d2) and same Heun integration as
/// `Tripod`, but each neuron reads its own AdExParameterHet fields.
pub struct TripodHet {
  // Soma (AdEx) parameters — per-neuron arrays.
  soma_param : AdExParameterHet
  soma_spike : AdExPostSpike
  // Soma-membrane derived scalars (per-neuron). Stored explicitly so
  // we don't recompute c = tm/(1000*r) inside the inner loop.
  c : Array[Float]
  gl : Array[Float]
  // Dendrite passive parameters (d1 + d2) — shared across neurons.
  d1 : Dendrite
  d2 : Dendrite
  // External input currents.
  i_s : Array[Float]
  i_d1 : Array[Float]
  i_d2 : Array[Float]
  // Synaptic conductances (single-exp on soma + 2 dendrites).
  ge_s : Array[Float]
  gi_s : Array[Float]
  ge_d1 : Array[Float]
  gi_d1 : Array[Float]
  ge_d2 : Array[Float]
  gi_d2 : Array[Float]
  glu_s : Array[Float]
  gaba_s : Array[Float]
  glu_d1 : Array[Float]
  gaba_d1 : Array[Float]
  glu_d2 : Array[Float]
  gaba_d2 : Array[Float]
  // Synapse parameters (scalar).
  e_e : Float
  e_i : Float
  tau_e : Float
  tau_i : Float
  gsyn_e : Float
  gsyn_i : Float
  // State.
  n : Int
  v_s : Array[Float]
  w_s : Array[Float]
  v_d1 : Array[Float]
  v_d2 : Array[Float]
  fire : Array[Bool]
  threshold : Array[Float]
  tabs : Array[Int]
  // Heun temp arrays (4 * n).
  dv : Array[Float]
  dv_temp : Array[Float]
  syn_curr_s : Array[Float]
  syn_curr_d1 : Array[Float]
  syn_curr_d2 : Array[Float]
}

///|
/// Construct a TripodHet population. `soma_param` must already be
/// filled with per-neuron arrays of length N. Dendrite parameters
/// are derived from `d1`/`d2` (Julia's `human_dend` defaults).
pub fn TripodHet::new(
  n : Int,
  soma_param : AdExParameterHet,
  rng : Xoshiro,
) -> TripodHet {
  // Per-neuron c[i] = tm[i] / r[i] (matches Julia: C = tm / R, with
  // tm in ms and R in MΩ giving C in pF when multiplied by 1e-3
  // inside the dv formula).
  let c : Array[Float] = Array::make(n, 0.0F)
  let gl : Array[Float] = Array::make(n, 0.0F)
  for k in 0.. Unit {
  let n = p.n
  for i in 0.. Unit {
  let n = p.n
  for i in 0.. Unit {
  let n = p.n
  for i in 0.. Unit {
  let n = p.n
  for i in 0.. Unit {
  let n = p.n
  let mut k = 0
  while k < n {
    let vt = p.soma_param.vt[k]
    let el = p.soma_param.el[k]
    let tm = p.soma_param.tm[k]
    let r_val = p.soma_param.r[k]
    let dt_slope = p.soma_param.dt_slope[k]
    let tw = p.soma_param.tw[k]
    let a = p.soma_param.a[k]
    let c_val = p.c[k]
    let gl_val = p.gl[k]

    let ds : Float = if store_temp { p.dv_temp[k * 4 + 0] * dt } else { 0.0F }
    let dd1 : Float = if store_temp { p.dv_temp[k * 4 + 1] * dt } else { 0.0F }
    let dd2 : Float = if store_temp { p.dv_temp[k * 4 + 2] * dt } else { 0.0F }
    let dw : Float = if store_temp { p.dv_temp[k * 4 + 3] } else { 0.0F }

    let ic1 = -(p.v_d1[k] + dd1 - p.v_s[k] - ds) * p.d1.gax[k]
    let ic2 = -(p.v_d2[k] + dd2 - p.v_s[k] - ds) * p.d2.gax[k]

    let exp_term = if dt_slope < 0.0F {
      0.0F
    } else {
      dt_slope * expf((p.v_s[k] + ds - p.threshold[k]) / dt_slope)
    }
    let dv_s_val = (gl_val * (el - p.v_s[k] - ds) + exp_term -
                    p.w_s[k] - dw - p.syn_curr_s[k] - (ic1 + ic2) +
                    p.i_s[k]) / c_val

    let dv_d1_val = ((el - p.v_d1[k] - dd1) * p.d1.gm[k] -
                     p.syn_curr_d1[k] + ic1 + p.i_d1[k]) / p.d1.c[k]
    let dv_d2_val = ((el - p.v_d2[k] - dd2) * p.d2.gm[k] -
                     p.syn_curr_d2[k] + ic2 + p.i_d2[k]) / p.d2.c[k]

    let dw_val = (a * (p.v_s[k] + ds - el) - (p.w_s[k] + dw)) / tw

    if store_temp {
      p.dv_temp[k * 4 + 0] = dv_s_val
      p.dv_temp[k * 4 + 1] = dv_d1_val
      p.dv_temp[k * 4 + 2] = dv_d2_val
      p.dv_temp[k * 4 + 3] = dw_val
    } else {
      p.dv[k * 4 + 0] = dv_s_val
      p.dv[k * 4 + 1] = dv_d1_val
      p.dv[k * 4 + 2] = dv_d2_val
      p.dv[k * 4 + 3] = dw_val
    }
    let _ = vt
    let _ = tm
    let _ = r_val
    k = k + 1
  }
}

///|
/// Update TripodHet for one timestep. Same Julia `integrate!` order
/// as `step_tripod`, but reads per-neuron AdEx params.
pub fn step_tripod_het(p : TripodHet, dt : Float) -> Unit {
  let n = p.n
  let at = p.soma_spike.at
  let tau_a = p.soma_spike.tau_a
  let tabs_const = p.soma_spike.tabs_const
  let up = p.soma_spike.up
  let ap_membrane = p.soma_spike.ap_membrane
  let tabs_steps : Int = ((up + tabs_const) / dt).to_int()

  // 1. Synapses.
  tripod_het_soma_step_synapses(p, dt)
  tripod_het_dend_step_synapses(p, dt)

  // 2. Synaptic currents.
  tripod_het_syn_curr_soma(p)
  tripod_het_syn_curr_dends(p)

  // 3. Heun integration.
  tripod_het_heun_step(p, dt, false)
  for i in 0..<(n * 4) {
    p.dv_temp[i] = p.dv[i]
  }
  tripod_het_heun_step(p, dt, true)

  // 4 + 5. Apply + spike detection (Julia's exact per-neuron order).
  for k in 0.. 0 {
      // Refractory: v_s = Vr, v_d += axial coupling.
      p.v_s[k] = vr_k
      p.v_d1[k] = p.v_d1[k] + dt * (p.v_s[k] - p.v_d1[k]) * p.d1.gax[k] / p.d1.c[k]
      p.v_d2[k] = p.v_d2[k] + dt * (p.v_s[k] - p.v_d2[k]) * p.d2.gax[k] / p.d2.c[k]
      continue
    }

    // Active phase. Detect fire BEFORE Heun apply.
    let v_s_pred = p.v_s[k] + p.dv[k * 4 + 0] * dt
    p.fire[k] = v_s_pred >= -10.0F

    if p.fire[k] {
      p.dv[k * 4 + 0] = ap_membrane - p.v_s[k]
      p.v_s[k] = ap_membrane
      p.w_s[k] = p.w_s[k] + b_k
      p.threshold[k] = p.threshold[k] + at
      p.tabs[k] = tabs_steps
      continue
    }

    // No spike: apply Heun correction.
    p.v_s[k] = p.v_s[k] + 0.5F * dt * (p.dv[k * 4 + 0] + p.dv_temp[k * 4 + 0])
    p.v_d1[k] = p.v_d1[k] + 0.5F * dt * (p.dv[k * 4 + 1] + p.dv_temp[k * 4 + 1])
    p.v_d2[k] = p.v_d2[k] + 0.5F * dt * (p.dv[k * 4 + 2] + p.dv_temp[k * 4 + 2])
    p.w_s[k] = p.w_s[k] + 0.5F * dt * (p.dv[k * 4 + 3] + p.dv_temp[k * 4 + 3])
  }
}