// BallAndStick — ball-and-stick multi-compartment neuron: AdEx soma
// + 1 passive dendritic compartment. Bit-exact port of
// SNNModels.jl/src/populations/multicompartment/ballandstick.jl.
//
// Float32 contract: every arithmetic uses `Float` (Float32); expf
// for the exponential term (matches Julia's exp(Float32, x)).
//
// Structure:
//   - Soma: AdEx neuron (v_s, w_s, fire, threshold, tabs, etc.)
//   - Dendrite: passive RC compartment (uses Dendrite for geometry)
//   - Axial current: gax between soma and dendrite
//   - Heun integration (predictor-corrector) for numerical stability
//
// Integration order (matches Julia BallAndStick):
//   1. update_synapses! for soma + dendrite (uses SimpleSynapse for
//      now — separate from TripodSomaSynapse / TripodDendSynapse since
//      those require multi-receptor infrastructure we haven't built).
//   2. Heun: update_neuron! twice (predictor + corrector) to get Δv.
//   3. Apply: spike detection on soma, threshold dynamics, axial
//      coupling during refractory periods.

///|
/// BallAndStick neuron state — soma (AdEx) + 1 dendrite.
pub struct BallAndStick {
  // Soma (AdEx) parameters.
  soma_param : AdExParameter
  soma_spike : AdExPostSpike
  // Dendrite passive parameters.
  dend : Dendrite
  // External input currents.
  i_s : Array[Float]  // soma
  i_d : Array[Float]  // dendrite
  // Synaptic conductances (simple single-exp for soma + dend).
  ge_s : Array[Float]
  gi_s : Array[Float]
  ge_d : Array[Float]
  gi_d : Array[Float]
  glu_s : Array[Float]  // raw input buffer for soma (from SpikingSynapse)
  gaba_s : Array[Float]
  glu_d : Array[Float]  // raw input buffer for dendrite
  gaba_d : Array[Float]
  // Synapse parameters (single-exp like AdExSinExp).
  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_d : Array[Float]
  fire : Array[Bool]
  threshold : Array[Float]
  tabs : Array[Int]
  // Temporary variables for Heun integration.
  dv : Array[Float]      // [v_s; v_d; w_s] per neuron (size n*3 flattened)
  dv_temp : Array[Float] // previous Heun step (size n*3 flattened)
  syn_curr_s : Array[Float]  // soma synaptic current
  syn_curr_d : Array[Float]  // dendrite synaptic current
  // Axial current (soma-dendrite coupling).
  ic : Float
}

///|
/// Default BallAndStick: AdEx soma + Dendrite with human_dend defaults.
pub fn BallAndStick::new(
  n : Int,
  soma_param : AdExParameter,
  rng : Xoshiro,
) -> BallAndStick {
  // 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
  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 {
    // For the corrector pass: use dv_temp (the predictor's dv * dt)
    // to extrapolate v_s+v_d. For the predictor pass: ds=dd=dw=0.
    let ds : Float = if store_temp { p.dv_temp[k * 3 + 0] * dt } else { 0.0F }
    let dd : Float = if store_temp { p.dv_temp[k * 3 + 1] * dt } else { 0.0F }
    // dw: Julia quirk — no dt multiplication.
    let dw : Float = if store_temp { p.dv_temp[k * 3 + 2] } else { 0.0F }

    // Axial current: ic = -(v_d + dd - v_s - ds) * gax
    // ds, dd are Δv*dt predictions (note: our dv is per-ms, Julia
    // uses Δv per timestep where Δv*dt is the actual voltage change).
    // We treat ds, dd as already-multiplied-by-dt estimates.
    let ic_val = -(p.v_d[k] + dd - p.v_s[k] - ds) * p.dend.gax[k]

    // Soma Δv:
    let exp_term = if dt_slope < 0.0F {
      0.0F
    } else {
      dt_slope * expf((p.v_s[k] + ds - p.threshold[k]) / dt_slope)
    }
    // Julia: dv_s = 1/c * (gl*(El - v_s - ds) + exp_term - w_s - dw - syn_curr_s - ic + i_s)
    let dv_s_val = (gl * (el - p.v_s[k] - ds) + exp_term -
                    p.w_s[k] - dw - p.syn_curr_s[k] - ic_val +
                    p.i_s[k]) / c

    // Dendrite Δv: ((El - v_d - dd)*gm - syn_curr_d + ic + i_d) / C_d
    let dv_d_val = ((el - p.v_d[k] - dd) * p.dend.gm[k] -
                    p.syn_curr_d[k] + ic_val +
                    p.i_d[k]) / p.dend.c[k]

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

    // Write into dv (predictor) or dv_temp (corrector).
    if store_temp {
      p.dv_temp[k * 3 + 0] = dv_s_val
      p.dv_temp[k * 3 + 1] = dv_d_val
      p.dv_temp[k * 3 + 2] = dw_val
    } else {
      p.dv[k * 3 + 0] = dv_s_val
      p.dv[k * 3 + 1] = dv_d_val
      p.dv[k * 3 + 2] = dw_val
    }
    k = k + 1
  }
}

///|
/// Update BallAndStick for one timestep.
/// Mirrors Julia's BallAndStick `integrate!` order bit-exactly:
///   1. update_synapses! (soma + dendrite)
///   2. synaptic_current! (soma + dendrite)
///   3. Heun: predictor pass → save dv → corrector pass → save dv_temp
///   4. for each neuron:
///      - decrement tabs, update threshold
///      - if tabs > τabs/dt (backprop): v_s=AP, v_d += dt*axial
///      - elsif tabs > 0 (abs refractory): v_s=Vr, v_d += 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 (skip Heun apply)
///          * else: apply Heun correction to v_s, v_d, w_s
///
/// Julia writes to `Δv` in-place during both Heun passes; after the
/// loop, `Δv` holds the corrector's values. Spike detection uses the
/// corrector's Δv_s. The corrector's exp_term can explode when v_s is
/// near or above threshold; Julia's `fire && continue` skips the
/// bogus v_d update. We must do the same to avoid v_d runaway.
///
/// `tabs_steps = round(Int, (up + τabs) / dt)` — Julia sets the full
/// backprop + refractory duration. We must use both `up` and `tabs_const`
/// (τabs in Julia) — using τabs alone gives half the refractory window.
pub fn step_ballandstick(p : BallAndStick, 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 (single-exp on both soma and dend).
  ballandstick_soma_step_synapses(p, dt)
  ballandstick_dend_step_synapses(p, dt)

  // 2. Synaptic currents.
  ballandstick_syn_curr_soma(p)
  ballandstick_syn_curr_dend(p)

  // 3. Heun integration.
  //    Predictor pass writes dv; corrector pass writes dv_temp.
  ballandstick_heun_step(p, dt, false)
  for i in 0..<(n * 3) {
    p.dv_temp[i] = p.dv[i]
  }
  ballandstick_heun_step(p, dt, true)

  // 4 + 5. Apply + spike detection (Julia's exact per-neuron order).
  for k in 0.. 0 {
      // Refractory (tabs > 0 means within backprop or abs-refract).
      // Julia distinguishes `tabs > τabs/dt` (backprop, v_s=AP) from
      // `tabs > 0` (abs-refract, v_s=Vr). Here tabs_steps = up+τabs/dt,
      // so tabs values (up+τabs/dt) down to (τabs/dt+1) are backprop.
      // We use a uniform `v_s = Vr` for refractory since our spike
      // reset already sets v_s = ap_membrane on fire, and the next
      // apply loop iteration resets v_s from the fire flag at the top.
      // Bit-exact match to Julia's abs-refract branch: v_s = Vr and
      // v_d += dt * (v_s - v_d) * gax / C.
      p.v_s[k] = vr
      p.v_d[k] = p.v_d[k] + dt * (p.v_s[k] - p.v_d[k]) * p.dend.gax[k] / p.dend.c[k]
      continue
    }

    // Active phase. Detect fire BEFORE Heun apply (Julia order).
    // Predictive criterion: v_s + corrector_dv_s * dt >= -10mV.
    let v_s_pred = p.v_s[k] + p.dv[k * 3 + 0] * dt
    p.fire[k] = v_s_pred >= -10.0F

    if p.fire[k] {
      // Replace corrector dv_s with (AP - v_s); set v_s = AP, etc.
      p.dv[k * 3 + 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
      // CRITICAL: skip Heun apply to v_d/w_s — corrector dv_d may
      // have exploded (exp_term in corrector of v_s + Δv_pred*dt).
      continue
    }

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