// Tripod — Tripod multi-compartment neuron: AdEx soma + 2 passive
// dendritic compartments (d1, d2). Bit-exact port of
// SNNModels.jl/src/populations/multicompartment/tripod.jl.
//
// Structure:
//   - Soma: AdEx neuron (v_s, w_s, fire, threshold, tabs, etc.)
//   - 2 dendrites: passive RC compartments (each with own Dendrite)
//   - Axial currents: gax between soma and each dendrite
//   - Heun integration (predictor-corrector) for numerical stability
//
// Bit-exact ordering: same as BallAndStick but with one extra
// dendrite Δv component. Per neuron we have 4 Δv values:
//   dv[k*4 + 0] = dv_s    (soma)
//   dv[k*4 + 1] = dv_d1   (dendrite 1)
//   dv[k*4 + 2] = dv_d2   (dendrite 2)
//   dv[k*4 + 3] = dw_s    (soma adaptation current)
// dv length = n * 4.

///|
/// Tripod neuron state — soma (AdEx) + 2 dendrites (d1, d2).
pub struct Tripod {
  // Soma (AdEx) parameters.
  soma_param : AdExParameter
  soma_spike : AdExPostSpike
  // Dendrite passive parameters (d1 + d2).
  d1 : Dendrite
  d2 : Dendrite
  // External input currents.
  i_s : Array[Float]  // soma
  i_d1 : Array[Float] // dendrite 1
  i_d2 : Array[Float] // dendrite 2
  // Synaptic conductances (simple 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.
  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]
  // Temporary variables for Heun integration (n*4 per neuron).
  dv : Array[Float]
  dv_temp : Array[Float]
  syn_curr_s : Array[Float]
  syn_curr_d1 : Array[Float]
  syn_curr_d2 : Array[Float]
}

///|
/// Default Tripod: AdEx soma + 2 Dendrites with human_dend defaults.
pub fn Tripod::new(
  n : Int,
  soma_param : AdExParameter,
  rng : Xoshiro,
) -> Tripod {
  // Soma initial v_s in [vr, vt].
  let v_s = Array::make(n, 0.0F)
  let spread = soma_param.vt - soma_param.vr
  for k in 0.. Unit {
  let n = p.n
  for i in 0.. Unit {
  let n = p.n
  // dendrite 1.
  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 p_ = p.soma_param
  let c = p_.c
  let gl = p_.gl
  let el = p_.el
  let dt_slope = p_.dt_slope
  let tw = p_.tw
  let a = p_.a
  let mut k = 0
  while k < n {
    // Read predictor's Δv (or 0 for predictor pass).
    // ds/dd1/dd2 multiplied by dt (soma + dendrite voltages).
    // dw NOT multiplied by dt (Julia quirk).
    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 }

    // Axial currents: ic_d1 = -(v_d1 + dd1 - v_s - ds) * gax_d1
    //                 ic_d2 = -(v_d2 + dd2 - v_s - ds) * gax_d2
    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]

    // Soma Δv (sum of both axial currents; Julia uses `sum(ic)`).
    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 * (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

    // Dendrite 1 Δv: ((El - v_d1 - dd1)*gm_d1 - syn_curr_d1 + ic1 + i_d1) / C_d1
    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]
    // Dendrite 2 Δv: same with ic2.
    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]

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

    // Write.
    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
    }
    k = k + 1
  }
}

///|
/// Update Tripod for one timestep.
/// Mirrors Julia's Tripod `integrate!` order bit-exactly:
///   1. update_synapses! (soma + both dendrites)
///   2. synaptic_current! (soma + both dendrites)
///   3. Heun: predictor → save → corrector → save
///   4. for each neuron:
///      - decrement tabs, update threshold
///      - if tabs > 0 (refractory): v_s=Vr, v_d1,v_d2 += dt*axial
///      - else (active):
///          * detect fire (predictive): v_s + corrector_dv*dt >= -10mV
///          * if fire: v_s=AP, w+=b, θ+=At, tabs=tabs_steps, continue
///          * else: apply Heun correction to v_s, v_d1, v_d2, w_s
///
/// See `step_ballandstick` doc for why we must skip v_d apply on
/// fire (corrector exp_term can explode when v_s is near/above
/// threshold; Julia's `fire && continue` prevents v_d runaway).
///
/// `tabs_steps = round(Int, (up + τabs) / dt)` — Julia's full
/// backprop + refractory duration. Using only `τabs` gives half.
pub fn step_tripod(p : Tripod, dt : Float) -> Unit {
  let n = p.n
  let p_ = p.soma_param
  let vt = p_.vt
  let vr = p_.vr
  let b = p_.b
  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_soma_step_synapses(p, dt)
  tripod_dend_step_synapses(p, dt)

  // 2. Synaptic currents.
  tripod_syn_curr_soma(p)
  tripod_syn_curr_dends(p)

  // 3. Heun integration.
  tripod_heun_step(p, dt, false)
  for i in 0..<(n * 4) {
    p.dv_temp[i] = p.dv[i]
  }
  tripod_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
      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
      p.threshold[k] = p.threshold[k] + at
      p.tabs[k] = tabs_steps
      // Skip Heun apply (avoid corrector dv_d runaway).
      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])
  }
}