// julia_parity.mbt — Cross-implementation parity verification framework.
//
// Goal: every example in `SpikingNeuralNetworks.jl/examples/` should produce
// the same numerical trajectories (last-bit Float32) when run under MoonBit
// and Julia. This module is the verification backbone.
//
// Layout:
//   * A `ParityCase` describes one named experiment: name, expected output
//     (either a Float32 time-series, an Int spike-train, or a final-state
//     vector), and per-value tolerances.
//   * A `ParityResult` summarises the comparison: per-step max error in
//     ULPs (for floats) or in step counts (for spike trains), and a pass/
//     fail boolean.
//
// Reference data lives in `mbt/testdata/parity/.csv` as one value
// per line (Floats in their full Float32 repr, or Ints). See
// `mbt/testdata/parity/chain.csv` and `mbt/testdata/parity/single_if.csv`
// for examples generated by hand from analytical closed-form solutions.

// Float32 → ULP distance.
// Two Float32 numbers `a` and `b` are `n` ULPs apart if the difference of
// their bit-patterns (interpreted as Int32) is `n` (taking sign into
// account). ULP distance is the canonical "bit-exact" measure for Float
// comparisons because it ignores scale and only cares about precision.
///
///|
pub fn float32_ulp_distance(a : Float, b : Float) -> Int {
  // NaN: not a ULP distance — return -1 sentinel.
  if a != a || b != b {
    return -1
  }
  let ai = Float::reinterpret_as_int(a)
  let bi = Float::reinterpret_as_int(b)
  let diff = if ai < bi { bi - ai } else { ai - bi }
  diff
}

///|
/// Compare two Float32 values with a ULP tolerance. Returns true if the
/// distance in ULPs is ≤ `tol_ulps`.
pub fn float32_close(a : Float, b : Float, tol_ulps : Int) -> Bool {
  let d = float32_ulp_distance(a, b)
  if d < 0 {
    // NaN vs anything: only equal if both NaN.
    return a != a && b != b
  }
  d <= tol_ulps
}

///|
/// One reference trace loaded from a parity CSV file. `kind` is
/// `"float"` or `"int"`. `values` are in chronological order.
pub struct ParityTrace {
  kind : String
  values : Array[Float]
}

///|
/// Result of comparing two traces element-by-element.
pub struct ParityResult {
  name : String
  n_compared : Int
  max_error : Float // max abs error in raw units (or max ULP distance for floats)
  max_ulp : Int // max ULP distance across all compared elements (for floats)
  passed : Bool
}

///|
/// Compare two float traces element-by-element. Both arrays must have the
/// same length. Tolerance is in ULPs (`tol_ulps`) and absolute raw units
/// (`atol`); passes iff every element is within both.
pub fn compare_float_traces(
  name : String,
  actual : Array[Float],
  expected : Array[Float],
  atol : Float,
  tol_ulps : Int,
) -> ParityResult {
  let n_actual = actual.length()
  let n_expected = expected.length()
  let n_compared = if n_actual < n_expected { n_actual } else { n_expected }
  let mut max_abs = 0.0F
  let mut max_ulp = 0
  let mut passed = true
  for k in 0.. e { a - e } else { e - a }
    if abs_diff > max_abs {
      max_abs = abs_diff
    }
    let ulp = float32_ulp_distance(a, e)
    if ulp > max_ulp {
      max_ulp = ulp
    }
    let close = float32_close(a, e, tol_ulps)
    if !close && abs_diff > atol {
      passed = false
    }
  }
  if n_actual != n_expected {
    passed = false
  }
  { name, n_compared, max_error: max_abs, max_ulp, passed }
}

///|
/// Compare two spike trains (arrays of spike-time step indices). Tolerance
/// is `tol_steps` — each actual spike must match an expected spike within
/// `tol_steps` of it. Returns the count of matched spikes and unmatched
/// spikes (greedy nearest-neighbour matching).
pub fn compare_spike_trains(
  name : String,
  actual : Array[Int],
  expected : Array[Int],
  tol_steps : Int,
) -> ParityResult {
  let n_actual = actual.length()
  let n_expected = expected.length()
  let mut matched = 0
  let mut j = 0
  let mut max_err = 0
  for i in 0.. e { a - e } else { e - a }
      if d <= tol_steps {
        if best_j < 0 || d < best_d {
          best_j = k
          best_d = d
        }
      }
      k = k + 1
    }
    if best_j >= 0 {
      matched = matched + 1
      if best_d > max_err {
        max_err = best_d
      }
      j = best_j + 1
    }
  }
  let passed = matched == n_actual && matched == n_expected
  { name, n_compared: matched, max_error: Float::from_int(max_err), max_ulp: matched, passed }
}

///|
/// Load a parity reference trace from a CSV file. Format: one value per
/// line; lines starting with `#` are comments; blank lines skipped.
/// Returns ParityTrace with kind="float" (use parse_float on each line) or
/// kind="int" (parse_int). The kind is detected from the first non-comment
/// line — if it parses as Int without losing info, it's int; else float.
///
/// The file path is relative to the package root (where moon.mod lives).
/// This function does NOT touch the filesystem — the caller passes the
/// already-loaded file content as a String.
pub fn parse_parity_csv(content : String) -> ParityTrace {
  let mut kind = "float"
  let values : Array[Float] = []
  // We split on \n and process each line. (No `lines` helper in MoonBit core.)
  let mut i = 0
  let n = content.length()
  while i < n {
    // Find next newline.
    let mut j = i
    while j < n && content[j] != '\n' && content[j] != '\r' {
      j = j + 1
    }
    let line_view = substring(content, i, j)
    // Skip leading whitespace.
    let s = line_view
    // Detect kind on first non-empty, non-comment line.
    if !is_blank_or_comment(s) {
      // Try parsing as Int.
      let parsed_int = try_parse_int(s)
      let mut is_int = false
      match parsed_int {
        Some(_) => is_int = true
        None => ()
      }
      if is_int && values.length() == 0 {
        kind = "int"
      } else if !is_int && values.length() == 0 {
        kind = "float"
      }
      if kind == "int" {
        match parsed_int {
          Some(v) => values.push(Float::from_int(v))
          None => ()
        }
      } else {
        let parsed_f = try_parse_float(s)
        match parsed_f {
          Some(v) => values.push(v)
          None => ()
        }
      }
    }
    i = j + 1
    // Skip \r\n as one.
    if i < n && content[i - 1] == '\r' && i < n && content[i] == '\n' {
      i = i + 1
    }
  }
  { kind, values }
}

///|
/// Helper: substring by char indices. Bounds-checked.
fn substring(s : String, start : Int, end : Int) -> String {
  let mut out = ""
  let mut i = start
  if end > s.length() {
    return s
  }
  while i < end {
    out = out + s[i].unsafe_to_char().to_string()
    i = i + 1
  }
  out
}

///|
/// Helper: blank line or `#` comment.
fn is_blank_or_comment(s : String) -> Bool {
  let n = s.length()
  let mut i = 0
  while i < n {
    let c = s[i]
    if c == ' ' || c == '\t' {
      i = i + 1
      continue
    }
    if c == '#' {
      return true
    }
    return false
  }
  true
}

///|
/// Parse an Int from a trimmed string. Returns `Some(value)` on success.
fn try_parse_int(s : String) -> Int? {
  let trimmed = trim(s)
  if trimmed.length() == 0 {
    return None
  }
  let mut neg = false
  let mut i = 0
  if trimmed[0] == '-' {
    neg = true
    i = 1
  } else if trimmed[0] == '+' {
    i = 1
  }
  if i >= trimmed.length() {
    return None
  }
  let mut v = 0
  while i < trimmed.length() {
    let c = trimmed[i]
    if c < '0' || c > '9' {
      return None
    }
    v = v * 10 + (c.to_int() - '0'.to_int())
    i = i + 1
  }
  if neg {
    Some(-v)
  } else {
    Some(v)
  }
}

///|
/// Parse a Float from a trimmed string. Supports integer and decimal forms,
/// optional sign, optional exponent (`e±NN`).
fn try_parse_float(s : String) -> Float? {
  let trimmed = trim(s)
  if trimmed.length() == 0 {
    return None
  }
  let mut i = 0
  let mut neg = false
  if trimmed[0] == '-' {
    neg = true
    i = 1
  } else if trimmed[0] == '+' {
    i = 1
  }
  if i >= trimmed.length() {
    return None
  }
  // Integer part.
  let mut int_part = 0
  let mut saw_digit = false
  while i < trimmed.length() && trimmed[i] >= '0' && trimmed[i] <= '9' {
    int_part = int_part * 10 + (trimmed[i].to_int() - '0'.to_int())
    saw_digit = true
    i = i + 1
  }
  // Fractional part.
  let mut frac_part = 0.0F
  if i < trimmed.length() && trimmed[i] == '.' {
    i = i + 1
    let mut frac = 0.0F
    let mut div = 1.0F
    while i < trimmed.length() && trimmed[i] >= '0' && trimmed[i] <= '9' {
      let digit = Float::from_int(trimmed[i].to_int() - '0'.to_int())
      div = div * 10.0F
      frac = frac + digit / div
      saw_digit = true
      i = i + 1
    }
    frac_part = frac
  }
  if !saw_digit {
    return None
  }
  // Optional exponent.
  let mut exp_part = 0
  if i < trimmed.length() && (trimmed[i] == 'e' || trimmed[i] == 'E') {
    i = i + 1
    let mut exp_neg = false
    if i < trimmed.length() && trimmed[i] == '-' {
      exp_neg = true
      i = i + 1
    } else if i < trimmed.length() && trimmed[i] == '+' {
      i = i + 1
    }
    if i >= trimmed.length() {
      return None
    }
    while i < trimmed.length() && trimmed[i] >= '0' && trimmed[i] <= '9' {
      exp_part = exp_part * 10 + (trimmed[i].to_int() - '0'.to_int())
      i = i + 1
    }
    if exp_neg {
      exp_part = -exp_part
    }
  }
  // Trailing garbage: reject.
  if i < trimmed.length() {
    return None
  }
  let combined = Float::from_int(int_part) + frac_part
  let final_v = if neg { -combined } else { combined }
  if exp_part == 0 {
    return Some(final_v)
  }
  // Apply exponent via pow.
  let exp_d = Float::from_int(exp_part).to_double()
  let scaled = final_v.to_double() * @math.pow(10.0, exp_d)
  Some(Float::from_double(scaled))
}

///|
/// Trim leading/trailing whitespace from a String.
fn trim(s : String) -> String {
  let n = s.length()
  let mut start = 0
  while start < n && (s[start] == ' ' || s[start] == '\t') {
    start = start + 1
  }
  let mut end = n
  while end > start && (s[end - 1] == ' ' || s[end - 1] == '\t') {
    end = end - 1
  }
  substring(s, start, end)
}