///|
/// Result of a measurement update.
pub(all) enum UpdateResult {
  Accepted
  RejectedByGate
  InvalidMeasurement
  SingularInnovation
  MissingMeasurement
} derive(Debug, Eq)

///|
/// A compact record of the latest update, useful for telemetry and debugging.
pub struct UpdateSummary {
  result : UpdateResult
  innovation : Array[Double]
  innovation_covariance : Matrix
  normalized_innovation_squared : Double
} derive(Debug)

///|
pub fn UpdateSummary::result(self : UpdateSummary) -> UpdateResult {
  self.result
}

///|
pub fn UpdateSummary::innovation(self : UpdateSummary) -> Array[Double] {
  self.innovation.copy()
}

///|
pub fn UpdateSummary::innovation_covariance(self : UpdateSummary) -> Matrix {
  self.innovation_covariance.copy()
}

///|
pub fn UpdateSummary::nis(self : UpdateSummary) -> Double {
  self.normalized_innovation_squared
}

///|
/// Linear Kalman filter for dense state and measurement vectors.
///
/// The implementation uses the Joseph covariance form. Compared with the
/// short textbook form, Joseph form better preserves symmetry and positive
/// semidefiniteness when a sensor has very small noise.
pub struct KalmanND {
  mut x : Array[Double]
  mut p : Matrix
  mut q : Matrix
  mut r : Matrix
  mut f : Matrix
  mut h : Matrix
  initial_state : Array[Double]
  initial_covariance : Matrix
  mut last_innovation : Array[Double]
  mut last_innovation_covariance : Matrix
  mut last_gain : Matrix
  mut last_nis : Double
  mut gate_threshold : Double
  mut predict_count : Int
  mut accepted_count : Int
  mut rejected_count : Int
  mut missing_count : Int
}

///|
pub fn KalmanND::new(
  initial_state : Array[Double],
  initial_covariance : Array[Array[Double]],
  process_noise : Array[Array[Double]],
  measurement_noise : Array[Array[Double]],
  transition_model : Array[Array[Double]],
  observation_model : Array[Array[Double]],
) -> KalmanND {
  let state = initial_state.copy()
  let covariance = Matrix::from_rows(initial_covariance)
  let measurement = Matrix::from_rows(measurement_noise)
  let transition = Matrix::from_rows(transition_model)
  let observation = Matrix::from_rows(observation_model)
  let process = Matrix::from_rows(process_noise)
  let measurement_dimension = observation.rows()
  {
    x: state.copy(),
    p: covariance.copy(),
    q: process,
    r: measurement,
    f: transition,
    h: observation,
    initial_state: state,
    initial_covariance: covariance,
    last_innovation: Array::make(measurement_dimension, 0.0),
    last_innovation_covariance: Matrix::zeros(
      measurement_dimension, measurement_dimension,
    ),
    last_gain: Matrix::zeros(initial_state.length(), measurement_dimension),
    last_nis: 0.0,
    gate_threshold: 9.210340371976184,
    predict_count: 0,
    accepted_count: 0,
    rejected_count: 0,
    missing_count: 0,
  }
}

///|
pub fn KalmanND::from_model(
  model : LinearModel,
  initial_state : Array[Double],
  initial_covariance : Matrix,
) -> KalmanND {
  KalmanND::new(
    initial_state,
    initial_covariance.to_rows(),
    model.process_noise().to_rows(),
    model.measurement_noise().to_rows(),
    model.transition().to_rows(),
    model.observation().to_rows(),
  )
}

///|
pub fn KalmanND::state(self : KalmanND) -> Array[Double] {
  self.x.copy()
}

///|
pub fn KalmanND::covariance(self : KalmanND) -> Matrix {
  self.p.copy()
}

///|
/// Restore a validated state/covariance pair, useful for checkpoint recovery.
pub fn KalmanND::restore(
  self : KalmanND,
  state : Array[Double],
  covariance : Matrix,
) -> Bool {
  if state.length() != self.x.length() ||
    covariance.rows() != self.x.length() ||
    covariance.cols() != self.x.length() {
    return false
  }
  if !vector_is_finite(state) || !covariance_is_psd(covariance, 0.0001) {
    return false
  }
  self.x = state.copy()
  self.p = covariance.symmetric_part()
  true
}

///|
pub fn KalmanND::process_noise(self : KalmanND) -> Matrix {
  self.q.copy()
}

///|
pub fn KalmanND::measurement_noise(self : KalmanND) -> Matrix {
  self.r.copy()
}

///|
pub fn KalmanND::transition_model(self : KalmanND) -> Matrix {
  self.f.copy()
}

///|
pub fn KalmanND::observation_model(self : KalmanND) -> Matrix {
  self.h.copy()
}

///|
pub fn KalmanND::state_dimension(self : KalmanND) -> Int {
  self.x.length()
}

///|
pub fn KalmanND::measurement_dimension(self : KalmanND) -> Int {
  self.h.rows()
}

///|
pub fn KalmanND::predict_count(self : KalmanND) -> Int {
  self.predict_count
}

///|
pub fn KalmanND::accepted_count(self : KalmanND) -> Int {
  self.accepted_count
}

///|
pub fn KalmanND::rejected_count(self : KalmanND) -> Int {
  self.rejected_count
}

///|
pub fn KalmanND::missing_count(self : KalmanND) -> Int {
  self.missing_count
}

///|
pub fn KalmanND::set_gate_threshold(
  self : KalmanND,
  threshold : Double,
) -> Unit {
  self.gate_threshold = if threshold < 0.0 { 0.0 } else { threshold }
}

///|
pub fn KalmanND::gate_threshold(self : KalmanND) -> Double {
  self.gate_threshold
}

///|
pub fn KalmanND::set_transition_model(
  self : KalmanND,
  transition : Matrix,
) -> Bool {
  if transition.rows() != self.x.length() ||
    transition.cols() != self.x.length() {
    return false
  }
  self.f = transition.copy()
  true
}

///|
pub fn KalmanND::set_process_noise(
  self : KalmanND,
  process_noise : Matrix,
) -> Bool {
  if process_noise.rows() != self.x.length() ||
    process_noise.cols() != self.x.length() {
    return false
  }
  self.q = process_noise.copy()
  true
}

///|
pub fn KalmanND::set_measurement_noise(
  self : KalmanND,
  measurement_noise : Matrix,
) -> Bool {
  if measurement_noise.rows() != self.h.rows() ||
    measurement_noise.cols() != self.h.rows() {
    return false
  }
  self.r = measurement_noise.copy()
  true
}

///|
pub fn KalmanND::set_observation_model(
  self : KalmanND,
  observation : Matrix,
) -> Bool {
  if observation.cols() != self.x.length() {
    return false
  }
  self.h = observation.copy()
  self.last_innovation = Array::make(observation.rows(), 0.0)
  self.last_innovation_covariance = Matrix::zeros(
    observation.rows(),
    observation.rows(),
  )
  self.last_gain = Matrix::zeros(self.x.length(), observation.rows())
  true
}

///|
pub fn KalmanND::predict(self : KalmanND) -> Unit {
  self.predict_with_control([])
}

///|
/// Predict with an optional control vector. A zero-column control matrix is a
/// valid model, so callers can use an empty vector for the no-control case.
pub fn KalmanND::predict_with_control(
  self : KalmanND,
  control : Array[Double],
) -> Unit {
  if self.f.rows() != self.x.length() || self.f.cols() != self.x.length() {
    return
  }
  let predicted = self.f.multiply_vector(self.x)
  let controlled = if control.length() == 0 {
    predicted
  } else {
    // The base type does not own a B matrix; callers that need controls can
    // use predict_with_control_matrix below. Keeping this branch explicit
    // makes malformed controls a no-op instead of corrupting state.
    let _ = control
    predicted
  }
  if controlled.length() == self.x.length() {
    self.x = controlled
  }
  let predicted_covariance = self.f
    .multiply(self.p)
    .multiply(self.f.transpose())
    .add(self.q)
  if predicted_covariance.rows() == self.x.length() {
    self.p = predicted_covariance.symmetric_part()
  }
  self.predict_count = self.predict_count + 1
}

///|
/// Predict with a caller-supplied state-space control effect matrix.
pub fn KalmanND::predict_with_control_matrix(
  self : KalmanND,
  control_effect : Matrix,
) -> Unit {
  if control_effect.rows() == self.x.length() && control_effect.cols() == 1 {
    let controlled = vector_add(
      self.f.multiply_vector(self.x),
      control_effect.column(0),
    )
    if controlled.length() == self.x.length() {
      self.x = controlled
    }
    self.p = self.f
      .multiply(self.p)
      .multiply(self.f.transpose())
      .add(self.q)
      .symmetric_part()
    self.predict_count = self.predict_count + 1
  } else {
    self.predict()
  }
}

///|
pub fn KalmanND::update(
  self : KalmanND,
  measurement : Array[Double],
) -> UpdateResult {
  self.update_gated(measurement, self.gate_threshold)
}

///|
pub fn KalmanND::update_gated(
  self : KalmanND,
  measurement : Array[Double],
  threshold : Double,
) -> UpdateResult {
  let summary = self.update_internal(measurement, threshold)
  summary.result
}

///|
pub fn KalmanND::update_and_report(
  self : KalmanND,
  measurement : Array[Double],
) -> UpdateSummary {
  self.update_internal(measurement, self.gate_threshold)
}

///|
fn KalmanND::update_internal(
  self : KalmanND,
  measurement : Array[Double],
  threshold : Double,
) -> UpdateSummary {
  let empty = UpdateSummary::{
    result: InvalidMeasurement,
    innovation: [],
    innovation_covariance: Matrix::zeros(0, 0),
    normalized_innovation_squared: 0.0,
  }
  if measurement.length() != self.h.rows() || !vector_is_finite(measurement) {
    self.rejected_count = self.rejected_count + 1
    return empty
  }
  let predicted_measurement = self.h.multiply_vector(self.x)
  let innovation = vector_sub(measurement, predicted_measurement)
  let covariance = self.h
    .multiply(self.p)
    .multiply(self.h.transpose())
    .add(self.r)
    .symmetric_part()
  let inverse = covariance.inverse()
  match inverse {
    None => {
      self.rejected_count = self.rejected_count + 1
      UpdateSummary::{
        result: SingularInnovation,
        innovation,
        innovation_covariance: covariance,
        normalized_innovation_squared: 0.0,
      }
    }
    Some(inverse_covariance) => {
      let nis = vector_dot(
        innovation,
        inverse_covariance.multiply_vector(innovation),
      )
      self.last_innovation = innovation.copy()
      self.last_innovation_covariance = covariance.copy()
      self.last_nis = nis
      let safe_threshold = if threshold < 0.0 { 0.0 } else { threshold }
      if nis > safe_threshold {
        self.rejected_count = self.rejected_count + 1
        UpdateSummary::{
          result: RejectedByGate,
          innovation,
          innovation_covariance: covariance,
          normalized_innovation_squared: nis,
        }
      } else {
        let gain = self.p
          .multiply(self.h.transpose())
          .multiply(inverse_covariance)
        let correction = gain.multiply_vector(innovation)
        if correction.length() == self.x.length() {
          self.x = vector_add(self.x, correction)
        }
        let identity = Matrix::identity(self.x.length())
        let residual_operator = identity.sub(gain.multiply(self.h))
        let joseph = residual_operator
          .multiply(self.p)
          .multiply(residual_operator.transpose())
          .add(gain.multiply(self.r).multiply(gain.transpose()))
        if joseph.rows() == self.x.length() {
          self.p = joseph.symmetric_part()
        }
        self.last_gain = gain
        self.accepted_count = self.accepted_count + 1
        UpdateSummary::{
          result: Accepted,
          innovation,
          innovation_covariance: covariance,
          normalized_innovation_squared: nis,
        }
      }
    }
  }
}

///|
/// Update a subset of measurement channels, preserving the full state model.
pub fn KalmanND::update_partial(
  self : KalmanND,
  measurement : Array[Double],
  mask : Array[Bool],
) -> UpdateResult {
  if mask.length() != self.h.rows() || measurement.length() != self.h.rows() {
    return InvalidMeasurement
  }
  let mut selected = 0
  for enabled in mask {
    if enabled {
      selected = selected + 1
    }
  }
  if selected == 0 {
    self.missing_count = self.missing_count + 1
    return MissingMeasurement
  }
  let selected_measurement = Array::make(selected, 0.0)
  let selected_h = Matrix::zeros(selected, self.h.cols())
  let selected_r = Matrix::zeros(selected, selected)
  let mut output_row = 0
  for input_row in 0.. ignore
      }
      selected_r.set(output_row, output_row, self.r.get(input_row, input_row))
      |> ignore
      output_row = output_row + 1
    }
  }
  let old_h = self.h
  let old_r = self.r
  self.h = selected_h
  self.r = selected_r
  let result = self.update(selected_measurement)
  self.h = old_h
  self.r = old_r
  result
}

///|
pub fn KalmanND::update_missing(self : KalmanND) -> UpdateResult {
  self.missing_count = self.missing_count + 1
  self.last_innovation = []
  MissingMeasurement
}

///|
pub fn KalmanND::last_update(self : KalmanND) -> UpdateSummary {
  UpdateSummary::{
    result: Accepted,
    innovation: self.last_innovation.copy(),
    innovation_covariance: self.last_innovation_covariance.copy(),
    normalized_innovation_squared: self.last_nis,
  }
}

///|
pub fn KalmanND::innovation(self : KalmanND) -> Array[Double] {
  self.last_innovation.copy()
}

///|
pub fn KalmanND::innovation_covariance(self : KalmanND) -> Matrix {
  self.last_innovation_covariance.copy()
}

///|
pub fn KalmanND::kalman_gain(self : KalmanND) -> Matrix {
  self.last_gain.copy()
}

///|
pub fn KalmanND::normalized_innovation_squared(self : KalmanND) -> Double {
  self.last_nis
}

///|
pub fn KalmanND::reset(self : KalmanND) -> Unit {
  self.x = self.initial_state.copy()
  self.p = self.initial_covariance.copy()
  self.last_innovation = Array::make(self.h.rows(), 0.0)
  self.last_innovation_covariance = Matrix::zeros(self.h.rows(), self.h.rows())
  self.last_gain = Matrix::zeros(self.x.length(), self.h.rows())
  self.last_nis = 0.0
  self.predict_count = 0
  self.accepted_count = 0
  self.rejected_count = 0
  self.missing_count = 0
}

///|
/// Run a predict/update cycle for each measurement and return state history.
pub fn KalmanND::filter(
  self : KalmanND,
  measurements : Array[Array[Double]],
) -> Array[Array[Double]] {
  let states : Array[Array[Double]] = []
  for measurement in measurements {
    self.predict()
    self.update(measurement) |> ignore
    states.push(self.state())
  }
  states
}

///|
/// Inflate covariance after a long sensor outage, with a hard cap for safety.
pub fn KalmanND::inflate_covariance(
  self : KalmanND,
  factor : Double,
  maximum : Double,
) -> Unit {
  let safe_factor = if factor < 1.0 { 1.0 } else { factor }
  let safe_maximum = if maximum < safe_factor { safe_factor } else { maximum }
  let diagonal = self.p.diagonal_max()
  if diagonal > 0.0 && diagonal * safe_factor > safe_maximum {
    self.p = self.p.scale(safe_maximum / diagonal)
  } else {
    self.p = self.p.scale(safe_factor)
  }
}