///|
/// Extended Kalman Filter for differentiable non-linear models.
///
/// Callers provide the state transition and observation functions together
/// with their Jacobians.  All covariance updates use the same numerically
/// stable Joseph form as `KalmanND`.
pub struct EKF {
  mut x : Array[Double]
  mut p : Matrix
  mut q : Matrix
  mut r : 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 EKF::new(
  initial_state : Array[Double],
  initial_covariance : Array[Array[Double]],
  process_noise : Array[Array[Double]],
  measurement_noise : Array[Array[Double]],
) -> EKF {
  let state = initial_state.copy()
  let covariance = Matrix::from_rows(initial_covariance)
  let measurement = Matrix::from_rows(measurement_noise)
  let process = Matrix::from_rows(process_noise)
  {
    x: state.copy(),
    p: covariance.copy(),
    q: process,
    r: measurement,
    initial_state: state,
    initial_covariance: covariance,
    last_innovation: [],
    last_innovation_covariance: Matrix::zeros(0, 0),
    last_gain: Matrix::zeros(state.length(), 0),
    last_nis: 0.0,
    gate_threshold: 9.210340371976184,
    predict_count: 0,
    accepted_count: 0,
    rejected_count: 0,
    missing_count: 0,
  }
}

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

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

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

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

///|
pub fn EKF::set_process_noise(self : EKF, 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 EKF::set_measurement_noise(
  self : EKF,
  measurement_noise : Matrix,
) -> Bool {
  if !measurement_noise.is_square() {
    return false
  }
  self.r = measurement_noise.copy()
  true
}

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

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

///|
pub fn EKF::predict(
  self : EKF,
  f : (Array[Double]) -> Array[Double],
  jacobian_f : (Array[Double]) -> Array[Array[Double]],
) -> Unit {
  let next_state = f(self.x.copy())
  let jacobian = Matrix::from_rows(jacobian_f(self.x.copy()))
  if next_state.length() != self.x.length() ||
    jacobian.rows() != self.x.length() ||
    jacobian.cols() != self.x.length() {
    return
  }
  self.x = next_state
  self.p = jacobian
    .multiply(self.p)
    .multiply(jacobian.transpose())
    .add(self.q)
    .symmetric_part()
  self.predict_count = self.predict_count + 1
}

///|
pub fn EKF::predict_with_control(
  self : EKF,
  f : (Array[Double], Array[Double]) -> Array[Double],
  jacobian_f : (Array[Double], Array[Double]) -> Array[Array[Double]],
  control : Array[Double],
) -> Unit {
  let old_state = self.x.copy()
  let next_state = f(old_state.copy(), control.copy())
  let jacobian = Matrix::from_rows(jacobian_f(old_state, control.copy()))
  if next_state.length() != self.x.length() ||
    jacobian.rows() != self.x.length() ||
    jacobian.cols() != self.x.length() {
    return
  }
  self.x = next_state
  self.p = jacobian
    .multiply(self.p)
    .multiply(jacobian.transpose())
    .add(self.q)
    .symmetric_part()
  self.predict_count = self.predict_count + 1
}

///|
pub fn EKF::update(
  self : EKF,
  z : Array[Double],
  h : (Array[Double]) -> Array[Double],
  jacobian_h : (Array[Double]) -> Array[Array[Double]],
) -> UpdateResult {
  self.update_gated(z, h, jacobian_h, self.gate_threshold)
}

///|
pub fn EKF::update_gated(
  self : EKF,
  z : Array[Double],
  h : (Array[Double]) -> Array[Double],
  jacobian_h : (Array[Double]) -> Array[Array[Double]],
  threshold : Double,
) -> UpdateResult {
  let predicted = h(self.x.copy())
  let jacobian = Matrix::from_rows(jacobian_h(self.x.copy()))
  if predicted.length() != z.length() ||
    !vector_is_finite(z) ||
    jacobian.cols() != self.x.length() ||
    jacobian.rows() != z.length() {
    self.rejected_count = self.rejected_count + 1
    return InvalidMeasurement
  }
  let innovation = vector_sub(z, predicted)
  let covariance = jacobian
    .multiply(self.p)
    .multiply(jacobian.transpose())
    .add(self.r)
    .symmetric_part()
  let safe_threshold = if threshold < 0.0 { 0.0 } else { threshold }
  match covariance.inverse() {
    None => {
      self.rejected_count = self.rejected_count + 1
      SingularInnovation
    }
    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
      if nis > safe_threshold {
        self.rejected_count = self.rejected_count + 1
        RejectedByGate
      } else {
        let gain = self.p
          .multiply(jacobian.transpose())
          .multiply(inverse_covariance)
        self.x = vector_add(self.x, gain.multiply_vector(innovation))
        let identity = Matrix::identity(self.x.length())
        let residual_operator = identity.sub(gain.multiply(jacobian))
        self.p = residual_operator
          .multiply(self.p)
          .multiply(residual_operator.transpose())
          .add(gain.multiply(self.r).multiply(gain.transpose()))
          .symmetric_part()
        self.last_gain = gain
        self.accepted_count = self.accepted_count + 1
        Accepted
      }
    }
  }
}

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

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

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

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

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

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

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

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

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

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

///|
/// Run a sequence of EKF updates against a caller-owned model.
pub fn EKF::filter(
  self : EKF,
  measurements : Array[Array[Double]],
  f : (Array[Double]) -> Array[Double],
  jacobian_f : (Array[Double]) -> Array[Array[Double]],
  h : (Array[Double]) -> Array[Double],
  jacobian_h : (Array[Double]) -> Array[Array[Double]],
) -> Array[Array[Double]] {
  let states : Array[Array[Double]] = []
  for measurement in measurements {
    self.predict(f, jacobian_f)
    self.update(measurement, h, jacobian_h) |> ignore
    states.push(self.state())
  }
  states
}