// AdEx (Adaptive Exponential Integrate-and-Fire) neuron — bit-exact
// port of SNNModels.jl/src/populations/generized_if/adex.jl.
//
// Julia reference:
//   src/populations/generized_if/adex.jl
//   src/populations/spike/postspike.jl
//
// Float32 contract: every arithmetic operation uses `Float` (Float32).
// The exponential term uses `expf` (libm) for bit-exact match with
// Julia's `exp(Float32, x)`.

///|
/// AdExParameter — biophysical constants of an AdEx neuron.
pub struct AdExParameter {
  c : Float
  gl : Float
  vt : Float
  vr : Float
  el : Float
  tm : Float
  r : Float
  dt_slope : Float
  tw : Float
  a : Float
  b : Float
}

///|
/// Default AdExParameter, matching Julia's `AdExParameter()`.
pub fn AdExParameter::new() -> AdExParameter {
  // Julia defaults:
  //   C = 281pF
  //   gl = 40nS
  //   Vt = -50mV
  //   Vr = -70.6mV
  //   El = -70.6mV
  //   τm = C / gl
  //   R = nS / gl
  //   ΔT = 2mV
  //   τw = 144ms
  //   a = 4nS
  //   b = 80.5pA
  //
  // All in Float32 (units normalised to 1.0).
  let c : Float = 281.0F
  let gl : Float = 40.0F
  let tm : Float = 281.0F / 40.0F
  let r : Float = 1.0F / 40.0F
  { c, gl, vt: -50.0F, vr: -70.6F, el: -70.6F, tm, r, dt_slope: 2.0F,
    tw: 144.0F, a: 4.0F, b: 80.5F }
}

///|
/// AdExParameter with a custom reset potential `vr`. Matches Julia's
/// `AdExParameter(; Vr = -50mV)`. Other fields use defaults.
pub fn AdExParameter::with_vr(vr : Float) -> AdExParameter {
  let c : Float = 281.0F
  let gl : Float = 40.0F
  let tm : Float = 281.0F / 40.0F
  let r : Float = 1.0F / 40.0F
  { c, gl, vt: -50.0F, vr, el: -70.6F, tm, r, dt_slope: 2.0F,
    tw: 144.0F, a: 4.0F, b: 80.5F }
}

///|
/// AdExParameter with full custom tm / vt / vr / el / r fields.
/// Matches Julia's `AdExParameter(; tm, vt, vr, el, r)`. Other
/// fields use defaults. Required for non-default Litwin-Kumar-Doiron
/// 2014 parameters (300pF/15nS=20ms tm, -70mV El, -52mV Vt, -60mV Vr,
/// R=1/15nS).
pub fn AdExParameter::custom(
  tm~ : Float,
  vt~ : Float,
  vr~ : Float,
  el~ : Float,
  r~ : Float,
) -> AdExParameter {
  { c: tm * r * 0.0F + 281.0F, gl: 1.0F / r, vt, vr, el, tm, r,
    dt_slope: 2.0F, tw: 144.0F, a: 4.0F, b: 80.5F }
}

///|
/// AdEx PostSpike — adds `At` (threshold jump) and `τA` (threshold decay)
/// on top of the IF PostSpike fields.
pub struct AdExPostSpike {
  at : Float
  tau_a : Float
  ap_membrane : Float
  tabs_const : Float
  up : Float
}

///|
pub fn AdExPostSpike::new() -> AdExPostSpike {
  // Julia: PostSpike{Float32}(; At = 0mV, τA = 10ms, AP_membrane = 10.0f0mV,
  //                       τabs = 1ms, up = 1ms)
  { at: 0.0F, tau_a: 10.0F, ap_membrane: 10.0F, tabs_const: 1.0F, up: 1.0F }
}

///|
/// AdExPostSpike with custom `at` (after-spike threshold jump) and
/// `tau_a` (threshold time constant). Matches Julia's
/// `PostSpike(; At = 10mV, τA = 10ms)`.
pub fn AdExPostSpike::with_at(at : Float, tau_a : Float) -> AdExPostSpike {
  { at, tau_a, ap_membrane: 10.0F, tabs_const: 1.0F, up: 1.0F }
}

///|
/// AdEx neuron state — a population of N adaptive exponential LIF neurons.
pub struct AdEx {
  param : AdExParameter
  spike : AdExPostSpike
  n : Int
  v : Array[Float]
  w : Array[Float]
  fire : Array[Bool]
  threshold : Array[Float]
  tabs : Array[Int]
  i : Array[Float]
  syn_curr : Array[Float]
  // Synapse state
  ge : Array[Float]
  gi : Array[Float]
  he : Array[Float]
  hi : Array[Float]
  glu : Array[Float]
  gaba : Array[Float]
  gsyn_e : Array[Float]
  gsyn_i : Array[Float]
  e_e : Float
  e_i : Float
  tre : Float
  tde : Float
  tri : Float
  tdi : Float
}

///|
/// Construct a new AdEx population with `n` neurons, using the
/// default PostSpike (At=0mV, τA=10ms).
pub fn AdEx::new(n : Int, param : AdExParameter, rng : Xoshiro) -> AdEx {
  AdEx::new_with_spike(n, param, AdExPostSpike::new(), rng)
}

///|
/// Construct a new AdEx population with a custom PostSpike.
pub fn AdEx::new_with_spike(
  n : Int,
  param : AdExParameter,
  spike : AdExPostSpike,
  rng : Xoshiro,
) -> AdEx {
  let v = Array::make(n, 0.0F)
  let spread = param.vt - param.vr
  for k in 0.. Unit {
  let n = p.n
  let p_ = p.param
  let tm = p_.tm
  let vt = p_.vt
  let vr = p_.vr
  let el = p_.el
  let r = p_.r
  let dt_slope = p_.dt_slope
  let tw = p_.tw
  let a = p_.a
  let b = p_.b
  let at = p.spike.at
  let tau_a = p.spike.tau_a
  let tabs_const = p.spike.tabs_const
  let tabs_steps : Int = (tabs_const / dt).to_int()

  for i in 0.. 0 {
      continue
    }

    // Adaptation current
    p.w[i] = p.w[i] + dt * (a * (p.v[i] - el) - p.w[i]) / tw

    // Membrane potential: leakage + exponential + synapses + adaptation + ext I
    let exp_term = if dt_slope < 0.0F {
      0.0F
    } else {
      dt_slope * expf((p.v[i] - p.threshold[i]) / dt_slope)
    }
    p.v[i] = p.v[i] +
      dt *
      (
        -(p.v[i] - el) + exp_term - r * p.syn_curr[i] - r * p.w[i] +
        r * p.i[i]
      ) / tm

    // Threshold dynamics
    p.threshold[i] = p.threshold[i] + dt * (vt - p.threshold[i]) / tau_a

    // Spike detection
    p.fire[i] = p.v[i] >= 0.0F
    p.v[i] = if p.fire[i] { 20.0F } else { p.v[i] }
    p.w[i] = if p.fire[i] { p.w[i] + b } else { p.w[i] }
    p.threshold[i] = if p.fire[i] { p.threshold[i] + at } else { p.threshold[i] }
    p.tabs[i] = if p.fire[i] { tabs_steps } else { p.tabs[i] }
  }
  ()
}

///|
/// Re-export the synapse helpers for AdEx.
pub fn adex_step_synapses(p : AdEx, dt : Float) -> Unit {
  let n = p.n
  for i in 0.. Unit {
  let n = p.n
  for i in 0.. AdExParameterHet {
  let vt_arr : Array[Float] = Array::make(n, p.vt)
  let vr_arr : Array[Float] = Array::make(n, p.vr)
  let el_arr : Array[Float] = Array::make(n, p.el)
  let tm_arr : Array[Float] = Array::make(n, p.tm)
  let r_arr : Array[Float] = Array::make(n, p.r)
  let dt_slope_arr : Array[Float] = Array::make(n, p.dt_slope)
  let tw_arr : Array[Float] = Array::make(n, p.tw)
  let a_arr : Array[Float] = Array::make(n, p.a)
  let b_arr : Array[Float] = Array::make(n, p.b)
  { vt: vt_arr, vr: vr_arr, el: el_arr, tm: tm_arr, r: r_arr,
    dt_slope: dt_slope_arr, tw: tw_arr, a: a_arr, b: b_arr }
}

///|
/// AdExHet — population with per-neuron heterogeneous AdEx parameters.
/// Same synaptic state layout as AdEx (ge/gi/he/hi/glu/gaba/gsyn_e/gsyn_i).
pub struct AdExHet {
  param : AdExParameterHet
  spike : AdExPostSpike
  n : Int
  v : Array[Float]
  w : Array[Float]
  fire : Array[Bool]
  threshold : Array[Float]
  tabs : Array[Int]
  i : Array[Float]
  syn_curr : Array[Float]
  ge : Array[Float]
  gi : Array[Float]
  he : Array[Float]
  hi : Array[Float]
  glu : Array[Float]
  gaba : Array[Float]
  gsyn_e : Array[Float]
  gsyn_i : Array[Float]
  e_e : Float
  e_i : Float
  tre : Float
  tde : Float
  tri : Float
  tdi : Float
}

///|
/// Construct a heterogeneous AdEx population with per-neuron params.
pub fn AdExHet::new(
  n : Int,
  param : AdExParameterHet,
  rng : Xoshiro,
) -> AdExHet {
  let v : Array[Float] = Array::make(n, 0.0F)
  // Init: v[i] = param.vr[i] + rand * (param.vt[i] - param.vr[i]).
  let mut k = 0
  while k < n {
    let spread = param.vt[k] - param.vr[k]
    v[k] = param.vr[k] + next_f32(rng) * spread
    k = k + 1
  }
  let w : Array[Float] = Array::make(n, 0.0F)
  let fire : Array[Bool] = Array::make(n, false)
  let threshold : Array[Float] = Array::make(n, 0.0F)
  let mut k2 = 0
  while k2 < n {
    threshold[k2] = param.vt[k2]
    k2 = k2 + 1
  }
  let tabs : Array[Int] = Array::make(n, 1)
  let i : Array[Float] = Array::make(n, 0.0F)
  let syn_curr : Array[Float] = Array::make(n, 0.0F)
  let ge : Array[Float] = Array::make(n, 0.0F)
  let gi : Array[Float] = Array::make(n, 0.0F)
  let he : Array[Float] = Array::make(n, 0.0F)
  let hi : Array[Float] = Array::make(n, 0.0F)
  let glu : Array[Float] = Array::make(n, 0.0F)
  let gaba : Array[Float] = Array::make(n, 0.0F)
  let gsyn_e : Array[Float] = Array::make(n, 1.0F)
  let gsyn_i : Array[Float] = Array::make(n, 1.0F)
  { param, spike: AdExPostSpike::new(), n, v, w, fire, threshold,
    tabs, i, syn_curr, ge, gi, he, hi, glu, gaba, gsyn_e, gsyn_i,
    e_e: 0.0F, e_i: -75.0F,
    tre: 1.0F, tde: 6.0F, tri: 0.5F, tdi: 2.0F }
}

///|
/// Update the heterogeneous AdEx neuron state for one step.
/// Bit-exact port of Julia's `update_neuron!` for
/// `AdEx{Float32} = Vector{Float32}`:
///   v[i] = ifelse(fire[i], vr[i], v[i])
///   tabs[i] -= 1; if tabs[i] > 0 continue
///   w[i] += dt * (a[i] * (v[i] - el[i]) - w[i]) / τw[i]
///   v[i] += dt * (-(v[i] - el[i]) + ΔT[i]*exp((v[i]-θ[i])/ΔT[i])
///              - R[i] * (syn_curr[i] + w[i]) + R[i] * I[i]) / τm[i]
///   θ[i] += dt * (Vt[i] - θ[i]) / τA      (τA scalar from spike)
///   fire[i] = v[i] >= 0
///   ...
pub fn step_adex_het(p : AdExHet, dt : Float) -> Unit {
  let n = p.n
  let spike = p.spike
  let at = spike.at
  let tau_a = spike.tau_a
  let tabs_const = spike.tabs_const
  let tabs_steps : Int = (tabs_const / dt).to_int()

  let mut i = 0
  while i < n {
    let vr_i = p.param.vr[i]
    p.v[i] = if p.fire[i] { vr_i } else { p.v[i] }

    p.fire[i] = false
    p.tabs[i] = p.tabs[i] - 1
    if p.tabs[i] > 0 {
      i = i + 1
      continue
    }

    let a_i = p.param.a[i]
    let el_i = p.param.el[i]
    let tw_i = p.param.tw[i]
    p.w[i] = p.w[i] + dt * (a_i * (p.v[i] - el_i) - p.w[i]) / tw_i

    let exp_term = if p.param.dt_slope[i] < 0.0F {
      0.0F
    } else {
      p.param.dt_slope[i] * expf(
        (p.v[i] - p.threshold[i]) / p.param.dt_slope[i],
      )
    }
    let tm_i = p.param.tm[i]
    let r_i = p.param.r[i]
    p.v[i] = p.v[i] +
      dt *
      (
        -(p.v[i] - el_i) + exp_term - r_i * p.syn_curr[i] - r_i * p.w[i] +
        r_i * p.i[i]
      ) / tm_i

    p.threshold[i] = p.threshold[i] + dt * (p.param.vt[i] - p.threshold[i]) / tau_a

    p.fire[i] = p.v[i] >= 0.0F
    p.v[i] = if p.fire[i] { 20.0F } else { p.v[i] }
    let b_i = p.param.b[i]
    p.w[i] = if p.fire[i] { p.w[i] + b_i } else { p.w[i] }
    p.threshold[i] = if p.fire[i] { p.threshold[i] + at } else { p.threshold[i] }
    p.tabs[i] = if p.fire[i] { tabs_steps } else { p.tabs[i] }

    i = i + 1
  }
  ()
}