///|
/// A histogram-backed distribution snapshot.
pub struct DistributionSnapshot {
  edges : Array[Double]
  counts : Array[Int]
  probabilities : Array[Double]
  total : Int
}

///|
/// Distribution drift diagnostics between two samples.
pub struct DriftReport {
  baseline_count : Int
  current_count : Int
  mean_shift : Double
  median_shift : Double
  scale_ratio : Double
  ks_statistic : Double
  psi : Double
  js_divergence : Double
  wasserstein : Double
  drift_score : Double
  drifted : Bool
}

///|
/// Configurable thresholds for drift decisions.
pub struct DriftRule {
  psi_threshold : Double
  ks_threshold : Double
  mean_shift_threshold : Double
  scale_ratio_threshold : Double
  bins : Int
}

///|
pub fn drift_default_rule() -> DriftRule {
  {
    psi_threshold: 0.2,
    ks_threshold: 0.1,
    mean_shift_threshold: 0.2,
    scale_ratio_threshold: 0.2,
    bins: 10,
  }
}

///|
pub fn drift_rule(
  psi_threshold : Double,
  ks_threshold : Double,
  mean_shift_threshold : Double,
  scale_ratio_threshold : Double,
  bins : Int,
) -> DriftRule {
  {
    psi_threshold: if psi_threshold < 0.0 {
      0.0
    } else {
      psi_threshold
    },
    ks_threshold: if ks_threshold < 0.0 {
      0.0
    } else {
      ks_threshold
    },
    mean_shift_threshold: if mean_shift_threshold < 0.0 {
      0.0
    } else {
      mean_shift_threshold
    },
    scale_ratio_threshold: if scale_ratio_threshold < 0.0 {
      0.0
    } else {
      scale_ratio_threshold
    },
    bins: if bins < 2 {
      2
    } else {
      bins
    },
  }
}

///|
pub fn drift_range(
  baseline : Array[Double],
  current : Array[Double],
) -> Array[Double] {
  let values = []
  for value in baseline {
    values.push(value)
  }
  for value in current {
    values.push(value)
  }
  if values.length() == 0 {
    [0.0, 1.0]
  } else {
    [min_value(values), max_value(values)]
  }
}

///|
pub fn drift_edges(
  baseline : Array[Double],
  current : Array[Double],
  bins : Int,
) -> Array[Double] {
  let count = if bins < 2 { 2 } else { bins }
  let bounds = drift_range(baseline, current)
  let lower = bounds[0]
  let upper = if bounds[1] <= lower { lower + 1.0 } else { bounds[1] }
  let result = []
  for index = 0; index <= count; index = index + 1 {
    result.push(lower + (upper - lower) * index.to_double() / count.to_double())
  }
  result
}

///|
pub fn drift_bin_index(value : Double, edges : Array[Double]) -> Int {
  if edges.length() < 2 {
    return 0
  }
  if value <= edges[0] {
    return 0
  }
  for index = 1; index < edges.length(); index = index + 1 {
    if value <= edges[index] {
      return index - 1
    }
  }
  edges.length() - 2
}

///|
pub fn drift_histogram(
  data : Array[Double],
  edges : Array[Double],
) -> Array[Int] {
  let count = if edges.length() < 2 { 0 } else { edges.length() - 1 }
  let result = []
  for _ in 0.. 0 {
      let index = drift_bin_index(value, edges)
      result[index] += 1
    }
  }
  result
}

///|
pub fn drift_probabilities(counts : Array[Int]) -> Array[Double] {
  let mut total = 0
  for count in counts {
    if count > 0 {
      total += count
    }
  }
  let result = []
  for count in counts {
    result.push(
      if total == 0 {
        0.0
      } else {
        count.to_double() / total.to_double()
      },
    )
  }
  result
}

///|
fn drift_log(value : Double) -> Double {
  if value <= 0.0 {
    0.0
  } else {
    let mut scaled = value
    let mut exponent = 0
    while scaled > 2.0 {
      scaled = scaled / 2.0
      exponent += 1
    }
    while scaled < 0.5 {
      scaled = scaled * 2.0
      exponent -= 1
    }
    let ratio = (scaled - 1.0) / (scaled + 1.0)
    let square = ratio * ratio
    let mut power = ratio
    let mut total = 0.0
    for index = 0; index < 12; index = index + 1 {
      total += power / (2 * index + 1).to_double()
      power *= square
    }
    2.0 * total + exponent.to_double() * 0.6931471805599453
  }
}

///|
pub fn drift_snapshot(
  data : Array[Double],
  edges : Array[Double],
) -> DistributionSnapshot {
  let counts = drift_histogram(data, edges)
  {
    edges,
    counts,
    probabilities: drift_probabilities(counts),
    total: data.length(),
  }
}

///|
pub fn drift_snapshot_with_bins(
  data : Array[Double],
  baseline : Array[Double],
  bins : Int,
) -> DistributionSnapshot {
  drift_snapshot(data, drift_edges(baseline, data, bins))
}

///|
pub fn drift_cumulative(probabilities : Array[Double]) -> Array[Double] {
  let result = []
  let mut total = 0.0
  for probability in probabilities {
    total += probability
    result.push(total)
  }
  result
}

///|
pub fn drift_ks_statistic(
  baseline : Array[Double],
  current : Array[Double],
) -> Double {
  if baseline.length() == 0 || current.length() == 0 {
    return 0.0
  }
  let left = copy_and_sort(baseline)
  let right = copy_and_sort(current)
  let points = []
  for value in left {
    points.push(value)
  }
  for value in right {
    points.push(value)
  }
  let mut maximum = 0.0
  for value in points {
    let difference = abs_double(
      empirical_cdf(left, value) - empirical_cdf(right, value),
    )
    if difference > maximum {
      maximum = difference
    }
  }
  maximum
}

///|
pub fn drift_ks_signed(
  baseline : Array[Double],
  current : Array[Double],
) -> Double {
  if baseline.length() == 0 || current.length() == 0 {
    return 0.0
  }
  let difference = median(current) - median(baseline)
  let sign = if difference < 0.0 {
    -1.0
  } else if difference > 0.0 {
    1.0
  } else {
    0.0
  }
  sign * drift_ks_statistic(baseline, current)
}

///|
pub fn drift_psi_from_probabilities(
  expected : Array[Double],
  actual : Array[Double],
) -> Double {
  if expected.length() != actual.length() || expected.length() == 0 {
    return 0.0
  }
  let mut total = 0.0
  for index = 0; index < expected.length(); index = index + 1 {
    let left = if expected[index] <= 1.0e-12 {
      1.0e-12
    } else {
      expected[index]
    }
    let right = if actual[index] <= 1.0e-12 { 1.0e-12 } else { actual[index] }
    total += (right - left) * drift_log(right / left)
  }
  total
}

///|
pub fn drift_psi(
  baseline : Array[Double],
  current : Array[Double],
  bins : Int,
) -> Double {
  let edges = drift_edges(baseline, current, bins)
  let expected = drift_probabilities(drift_histogram(baseline, edges))
  let actual = drift_probabilities(drift_histogram(current, edges))
  drift_psi_from_probabilities(expected, actual)
}

///|
pub fn drift_js_from_probabilities(
  left : Array[Double],
  right : Array[Double],
) -> Double {
  if left.length() != right.length() || left.length() == 0 {
    return 0.0
  }
  let mut result = 0.0
  for index = 0; index < left.length(); index = index + 1 {
    let a = if left[index] <= 1.0e-12 { 1.0e-12 } else { left[index] }
    let b = if right[index] <= 1.0e-12 { 1.0e-12 } else { right[index] }
    let midpoint = (a + b) / 2.0
    result += 0.5 * a * drift_log(a / midpoint)
    result += 0.5 * b * drift_log(b / midpoint)
  }
  result
}

///|
pub fn drift_js_divergence(
  baseline : Array[Double],
  current : Array[Double],
  bins : Int,
) -> Double {
  let edges = drift_edges(baseline, current, bins)
  let left = drift_probabilities(drift_histogram(baseline, edges))
  let right = drift_probabilities(drift_histogram(current, edges))
  drift_js_from_probabilities(left, right)
}

///|
pub fn drift_total_variation_from_probabilities(
  left : Array[Double],
  right : Array[Double],
) -> Double {
  if left.length() != right.length() || left.length() == 0 {
    return 0.0
  }
  let mut total = 0.0
  for index = 0; index < left.length(); index = index + 1 {
    total += abs_double(left[index] - right[index])
  }
  total / 2.0
}

///|
pub fn drift_total_variation(
  baseline : Array[Double],
  current : Array[Double],
  bins : Int,
) -> Double {
  let edges = drift_edges(baseline, current, bins)
  drift_total_variation_from_probabilities(
    drift_probabilities(drift_histogram(baseline, edges)),
    drift_probabilities(drift_histogram(current, edges)),
  )
}

///|
pub fn drift_wasserstein_approx(
  baseline : Array[Double],
  current : Array[Double],
  grid_size : Int,
) -> Double {
  if baseline.length() == 0 || current.length() == 0 {
    return 0.0
  }
  let points = if grid_size < 2 { 2 } else { grid_size }
  let mut total = 0.0
  for index = 0; index <= points; index = index + 1 {
    let probability = index.to_double() / points.to_double()
    total += abs_double(
      quantile(baseline, probability) - quantile(current, probability),
    )
  }
  total / (points + 1).to_double()
}

///|
pub fn drift_mean_shift(
  baseline : Array[Double],
  current : Array[Double],
) -> Double {
  if baseline.length() == 0 || current.length() == 0 {
    0.0
  } else {
    mean(current) - mean(baseline)
  }
}

///|
pub fn drift_median_shift(
  baseline : Array[Double],
  current : Array[Double],
) -> Double {
  if baseline.length() == 0 || current.length() == 0 {
    0.0
  } else {
    median(current) - median(baseline)
  }
}

///|
pub fn drift_scale_ratio(
  baseline : Array[Double],
  current : Array[Double],
) -> Double {
  let left = mad(baseline) * 1.4826
  let right = mad(current) * 1.4826
  if left <= 1.0e-12 {
    if right <= 1.0e-12 {
      1.0
    } else {
      1.0e12
    }
  } else {
    right / left
  }
}

///|
pub fn drift_normalized_mean_shift(
  baseline : Array[Double],
  current : Array[Double],
) -> Double {
  let scale = mad(baseline) * 1.4826
  if scale <= 1.0e-12 {
    abs_double(drift_mean_shift(baseline, current))
  } else {
    abs_double(drift_mean_shift(baseline, current)) / scale
  }
}

///|
pub fn drift_report(
  baseline : Array[Double],
  current : Array[Double],
  rule : DriftRule,
) -> DriftReport {
  let mean_shift = drift_normalized_mean_shift(baseline, current)
  let median_shift = drift_median_shift(baseline, current)
  let scale_ratio = drift_scale_ratio(baseline, current)
  let ks_statistic = drift_ks_statistic(baseline, current)
  let psi = drift_psi(baseline, current, rule.bins)
  let js_divergence = drift_js_divergence(baseline, current, rule.bins)
  let wasserstein = drift_wasserstein_approx(baseline, current, rule.bins * 2)
  let scale_change = abs_double(scale_ratio - 1.0)
  let score = mean_shift +
    abs_double(median_shift) +
    scale_change +
    ks_statistic +
    psi +
    js_divergence
  {
    baseline_count: baseline.length(),
    current_count: current.length(),
    mean_shift,
    median_shift,
    scale_ratio,
    ks_statistic,
    psi,
    js_divergence,
    wasserstein,
    drift_score: score,
    drifted: psi >= rule.psi_threshold ||
    ks_statistic >= rule.ks_threshold ||
    mean_shift >= rule.mean_shift_threshold ||
    scale_change >= rule.scale_ratio_threshold,
  }
}

///|
pub fn drift_score(
  baseline : Array[Double],
  current : Array[Double],
  bins : Int,
) -> Double {
  drift_report(baseline, current, drift_rule(0.2, 0.1, 0.2, 0.2, bins)).drift_score
}

///|
pub fn drifted(
  baseline : Array[Double],
  current : Array[Double],
  rule : DriftRule,
) -> Bool {
  drift_report(baseline, current, rule).drifted
}

///|
pub fn drift_report_lines(report : DriftReport) -> Array[String] {
  [
    "baseline_count=" + report.baseline_count.to_string(),
    "current_count=" + report.current_count.to_string(),
    "mean_shift=" + report.mean_shift.to_string(),
    "median_shift=" + report.median_shift.to_string(),
    "scale_ratio=" + report.scale_ratio.to_string(),
    "ks_statistic=" + report.ks_statistic.to_string(),
    "psi=" + report.psi.to_string(),
    "js_divergence=" + report.js_divergence.to_string(),
    "wasserstein=" + report.wasserstein.to_string(),
    "drift_score=" + report.drift_score.to_string(),
    "drifted=" + report.drifted.to_string(),
  ]
}

///|
pub fn drift_report_string(report : DriftReport) -> String {
  drift_report_lines(report).join("\n")
}

///|
pub fn rolling_drift_score(
  data : Array[Double],
  baseline_window : Int,
  current_window : Int,
  bins : Int,
) -> Array[Double] {
  let result = []
  if baseline_window <= 0 || current_window <= 0 {
    return result
  }
  for end = baseline_window + current_window
      end <= data.length()
      end = end + 1 {
    let baseline = []
    let current = []
    for index = end - baseline_window - current_window
        index < end - current_window
        index = index + 1 {
      baseline.push(data[index])
    }
    for index = end - current_window; index < end; index = index + 1 {
      current.push(data[index])
    }
    result.push(drift_score(baseline, current, bins))
  }
  result
}

///|
pub fn rolling_drift_flags(
  data : Array[Double],
  baseline_window : Int,
  current_window : Int,
  rule : DriftRule,
) -> Array[Bool] {
  let result = []
  if baseline_window <= 0 || current_window <= 0 {
    return result
  }
  for end = baseline_window + current_window
      end <= data.length()
      end = end + 1 {
    let baseline = []
    let current = []
    for index = end - baseline_window - current_window
        index < end - current_window
        index = index + 1 {
      baseline.push(data[index])
    }
    for index = end - current_window; index < end; index = index + 1 {
      current.push(data[index])
    }
    result.push(drifted(baseline, current, rule))
  }
  result
}

///|
pub fn drift_segment_reports(
  data : Array[Double],
  segments : Int,
  rule : DriftRule,
) -> Array[DriftReport] {
  let result = []
  if segments <= 1 {
    return result
  }
  let values = segment_values(data, segments)
  for index = 1; index < values.length(); index = index + 1 {
    result.push(drift_report(values[index - 1], values[index], rule))
  }
  result
}

///|
pub fn drift_segment_scores(
  data : Array[Double],
  segments : Int,
  bins : Int,
) -> Array[Double] {
  let result = []
  for
    report in drift_segment_reports(
      data,
      segments,
      drift_rule(0.2, 0.1, 0.2, 0.2, bins),
    ) {
    result.push(report.drift_score)
  }
  result
}

///|
pub fn drift_change_indices(
  data : Array[Double],
  window : Int,
  rule : DriftRule,
) -> Array[Int] {
  let result = []
  if window <= 0 {
    return result
  }
  for index = window * 2; index <= data.length(); index = index + 1 {
    let left = []
    let right = []
    for cursor = index - window * 2
        cursor < index - window
        cursor = cursor + 1 {
      left.push(data[cursor])
    }
    for cursor = index - window; cursor < index; cursor = cursor + 1 {
      right.push(data[cursor])
    }
    if drifted(left, right, rule) {
      result.push(index - window)
    }
  }
  result
}

///|
pub fn drift_sensitivity(
  baseline : Array[Double],
  current : Array[Double],
  rule : DriftRule,
) -> Array[Double] {
  let report = drift_report(baseline, current, rule)
  [
    report.mean_shift,
    abs_double(report.median_shift),
    abs_double(report.scale_ratio - 1.0),
    report.ks_statistic,
    report.psi,
    report.js_divergence,
    report.wasserstein,
    report.drift_score,
  ]
}

///|
pub fn drift_direction(
  baseline : Array[Double],
  current : Array[Double],
) -> Double {
  let shift = drift_median_shift(baseline, current)
  if shift < 0.0 {
    -1.0
  } else if shift > 0.0 {
    1.0
  } else {
    0.0
  }
}

///|
pub fn drift_alert_level(score : Double) -> Int {
  if score < 0.1 {
    0
  } else if score < 0.25 {
    1
  } else if score < 0.5 {
    2
  } else {
    3
  }
}

///|
pub fn drift_baseline_quality(data : Array[Double]) -> Double {
  if data.length() == 0 {
    0.0
  } else {
    1.0 / (1.0 + robust_signal_quality(data) + duplicate_fraction(data))
  }
}

///|
pub fn drift_report_valid(report : DriftReport) -> Bool {
  report.baseline_count > 0 &&
  report.current_count > 0 &&
  report.psi >= 0.0 &&
  report.ks_statistic >= 0.0
}

///|
pub fn drift_stable(
  baseline : Array[Double],
  current : Array[Double],
  rule : DriftRule,
) -> Bool {
  !drifted(baseline, current, rule)
}

///|
pub fn drift_gain(
  baseline : Array[Double],
  current : Array[Double],
  reference : Array[Double],
  bins : Int,
) -> Double {
  let raw = drift_score(baseline, current, bins)
  let reference_score = drift_score(baseline, reference, bins)
  if reference_score == 0.0 {
    0.0
  } else {
    (reference_score - raw) / reference_score
  }
}

///|
pub fn drift_compare_windows(
  data : Array[Double],
  window : Int,
  bins : Int,
) -> Array[Double] {
  let result = []
  if window <= 0 {
    return result
  }
  for index = window * 2; index <= data.length(); index = index + 1 {
    let left = []
    let right = []
    for cursor = index - window * 2
        cursor < index - window
        cursor = cursor + 1 {
      left.push(data[cursor])
    }
    for cursor = index - window; cursor < index; cursor = cursor + 1 {
      right.push(data[cursor])
    }
    result.push(drift_score(left, right, bins))
  }
  result
}