///|
/// Policy for multi-sensor fusion and outlier handling.
pub struct FusionPolicy {
  gate_threshold : Double
  max_time_gap : Int
  covariance_inflation : Double
  reject_non_finite : Bool
  predict_on_missing : Bool
} derive(Debug)

///|
pub fn FusionPolicy::new(
  gate_threshold : Double,
  max_time_gap : Int,
  covariance_inflation : Double,
  reject_non_finite : Bool,
  predict_on_missing : Bool,
) -> FusionPolicy {
  {
    gate_threshold: if gate_threshold < 0.0 {
      0.0
    } else {
      gate_threshold
    },
    max_time_gap: if max_time_gap < 0 {
      0
    } else {
      max_time_gap
    },
    covariance_inflation: if covariance_inflation < 1.0 {
      1.0
    } else {
      covariance_inflation
    },
    reject_non_finite,
    predict_on_missing,
  }
}

///|
pub fn FusionPolicy::default() -> FusionPolicy {
  FusionPolicy::new(9.210340371976184, 100, 1.25, true, true)
}

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

///|
pub fn FusionPolicy::max_time_gap(self : FusionPolicy) -> Int {
  self.max_time_gap
}

///|
pub fn FusionPolicy::covariance_inflation(self : FusionPolicy) -> Double {
  self.covariance_inflation
}

///|
pub fn FusionPolicy::reject_non_finite(self : FusionPolicy) -> Bool {
  self.reject_non_finite
}

///|
pub fn FusionPolicy::predict_on_missing(self : FusionPolicy) -> Bool {
  self.predict_on_missing
}

///|
pub struct FusionEvent {
  timestamp : Int
  sensor : String
  result : UpdateResult
  state : Array[Double]
  covariance : Matrix
  nis : Double
} derive(Debug)

///|
pub fn FusionEvent::timestamp(self : FusionEvent) -> Int {
  self.timestamp
}

///|
pub fn FusionEvent::sensor(self : FusionEvent) -> String {
  self.sensor
}

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

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

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

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

///|
pub struct FusionStatistics {
  mut total_packets : Int
  mut accepted_packets : Int
  mut rejected_packets : Int
  mut missing_packets : Int
  mut last_timestamp : Int
} derive(Debug)

///|
pub fn FusionStatistics::total_packets(self : FusionStatistics) -> Int {
  self.total_packets
}

///|
pub fn FusionStatistics::accepted_packets(self : FusionStatistics) -> Int {
  self.accepted_packets
}

///|
pub fn FusionStatistics::rejected_packets(self : FusionStatistics) -> Int {
  self.rejected_packets
}

///|
pub fn FusionStatistics::missing_packets(self : FusionStatistics) -> Int {
  self.missing_packets
}

///|
pub fn FusionStatistics::last_timestamp(self : FusionStatistics) -> Int {
  self.last_timestamp
}

///|
pub struct SensorFusion {
  filter : KalmanND
  policy : FusionPolicy
  mut statistics : FusionStatistics
  mut consecutive_missing : Int
}

///|
pub fn SensorFusion::new(
  filter : KalmanND,
  policy : FusionPolicy,
) -> SensorFusion {
  {
    filter,
    policy,
    statistics: {
      total_packets: 0,
      accepted_packets: 0,
      rejected_packets: 0,
      missing_packets: 0,
      last_timestamp: 0,
    },
    consecutive_missing: 0,
  }
}

///|
pub fn SensorFusion::state(self : SensorFusion) -> Array[Double] {
  self.filter.state()
}

///|
pub fn SensorFusion::covariance(self : SensorFusion) -> Matrix {
  self.filter.covariance()
}

///|
pub fn SensorFusion::set_gate_threshold(
  self : SensorFusion,
  threshold : Double,
) -> Unit {
  self.filter.set_gate_threshold(threshold)
}

///|
pub fn SensorFusion::statistics(self : SensorFusion) -> FusionStatistics {
  self.statistics
}

///|
pub fn SensorFusion::predict(self : SensorFusion) -> Unit {
  self.filter.predict()
}

///|
pub fn SensorFusion::process(
  self : SensorFusion,
  packet : ObservationPacket,
) -> FusionEvent {
  self.statistics.total_packets = self.statistics.total_packets + 1
  let timestamp = packet.timestamp()
  let sensor = packet.sensor()
  if !packet.is_valid() ||
    (self.policy.reject_non_finite() && !vector_is_finite(packet.values())) {
    self.statistics.rejected_packets = self.statistics.rejected_packets + 1
    return {
      timestamp,
      sensor,
      result: InvalidMeasurement,
      state: self.filter.state(),
      covariance: self.filter.covariance(),
      nis: 0.0,
    }
  }
  if self.policy.predict_on_missing() {
    self.filter.predict()
  }
  let old_noise = self.filter.measurement_noise()
  let packet_noise = packet.covariance()
  if packet_noise.rows() == old_noise.rows() && packet_noise.is_square() {
    self.filter.set_measurement_noise(packet_noise) |> ignore
  }
  let result = self.filter.update_gated(
    packet.values(),
    self.policy.gate_threshold(),
  )
  self.filter.set_measurement_noise(old_noise) |> ignore
  match result {
    Accepted => {
      self.statistics.accepted_packets = self.statistics.accepted_packets + 1
      self.consecutive_missing = 0
    }
    RejectedByGate | InvalidMeasurement | SingularInnovation => {
      self.statistics.rejected_packets = self.statistics.rejected_packets + 1
      self.consecutive_missing = self.consecutive_missing + 1
    }
    MissingMeasurement => {
      self.statistics.missing_packets = self.statistics.missing_packets + 1
      self.consecutive_missing = self.consecutive_missing + 1
    }
  }
  self.statistics.last_timestamp = timestamp
  {
    timestamp,
    sensor,
    result,
    state: self.filter.state(),
    covariance: self.filter.covariance(),
    nis: self.filter.normalized_innovation_squared(),
  }
}

///|
pub fn SensorFusion::process_missing(
  self : SensorFusion,
  timestamp : Int,
  sensor : String,
) -> FusionEvent {
  self.statistics.total_packets = self.statistics.total_packets + 1
  self.statistics.missing_packets = self.statistics.missing_packets + 1
  self.consecutive_missing = self.consecutive_missing + 1
  if self.policy.predict_on_missing() {
    self.filter.predict()
  }
  if self.consecutive_missing > self.policy.max_time_gap() {
    self.filter.inflate_covariance(
      self.policy.covariance_inflation(),
      1000000.0,
    )
  }
  self.statistics.last_timestamp = timestamp
  {
    timestamp,
    sensor,
    result: MissingMeasurement,
    state: self.filter.state(),
    covariance: self.filter.covariance(),
    nis: 0.0,
  }
}

///|
pub fn SensorFusion::run(
  self : SensorFusion,
  packets : Array[ObservationPacket],
) -> Array[FusionEvent] {
  let events : Array[FusionEvent] = []
  for packet in packets {
    events.push(self.process(packet))
  }
  events
}

///|
pub fn SensorFusion::reset(self : SensorFusion) -> Unit {
  self.filter.reset()
  self.statistics = {
    total_packets: 0,
    accepted_packets: 0,
    rejected_packets: 0,
    missing_packets: 0,
    last_timestamp: 0,
  }
  self.consecutive_missing = 0
}

///|
/// A ready-to-use two-dimensional position tracker with state
/// `[x, y, vx, vy]` and a position measurement `[x, y]`.
pub struct ConstantVelocityTracker2D {
  filter : KalmanND
  acceleration_variance : Double
  position_variance : Double
  mut last_timestamp : Int
  mut initialized_timestamp : Bool
}

///|
pub fn ConstantVelocityTracker2D::new(
  initial_position : Array[Double],
  initial_velocity : Array[Double],
  initial_variance : Double,
  acceleration_variance : Double,
  position_variance : Double,
) -> ConstantVelocityTracker2D {
  let state = make_constant_velocity_state(initial_position, initial_velocity)
  let variance = if initial_variance < 0.0 { 0.0 } else { initial_variance }
  let model = SensorModelPair::constant_velocity(
    2, 1.0, acceleration_variance, position_variance, position_variance,
  )
  let filter = KalmanND::from_model(
    model.position,
    state,
    Matrix::diagonal(4, variance),
  )
  {
    filter,
    acceleration_variance,
    position_variance,
    last_timestamp: 0,
    initialized_timestamp: false,
  }
}

///|
pub fn ConstantVelocityTracker2D::filter(
  self : ConstantVelocityTracker2D,
) -> KalmanND {
  self.filter
}

///|
pub fn ConstantVelocityTracker2D::state(
  self : ConstantVelocityTracker2D,
) -> Array[Double] {
  self.filter.state()
}

///|
pub fn ConstantVelocityTracker2D::covariance(
  self : ConstantVelocityTracker2D,
) -> Matrix {
  self.filter.covariance()
}

///|
pub fn ConstantVelocityTracker2D::set_gate_threshold(
  self : ConstantVelocityTracker2D,
  threshold : Double,
) -> Unit {
  self.filter.set_gate_threshold(threshold)
}

///|
pub fn ConstantVelocityTracker2D::predict(
  self : ConstantVelocityTracker2D,
  dt : Double,
) -> Unit {
  self.filter.set_transition_model(constant_velocity_transition(2, dt))
  |> ignore
  self.filter.set_process_noise(
    constant_velocity_process_noise(2, dt, self.acceleration_variance),
  )
  |> ignore
  self.filter.predict()
}

///|
pub fn ConstantVelocityTracker2D::step(
  self : ConstantVelocityTracker2D,
  timestamp : Int,
  position : Array[Double],
  covariance : Matrix,
) -> UpdateResult {
  if position.length() != 2 || covariance.rows() != 2 || covariance.cols() != 2 {
    return InvalidMeasurement
  }
  let dt = if !self.initialized_timestamp {
    1.0
  } else {
    let difference = timestamp - self.last_timestamp
    if difference <= 0 {
      1.0
    } else {
      difference.to_double()
    }
  }
  let transition = constant_velocity_transition(2, dt)
  let process = constant_velocity_process_noise(
    2,
    dt,
    self.acceleration_variance,
  )
  self.filter.set_transition_model(transition) |> ignore
  self.filter.set_process_noise(process) |> ignore
  self.filter.predict()
  let old_noise = self.filter.measurement_noise()
  self.filter.set_measurement_noise(covariance) |> ignore
  let result = self.filter.update(position)
  self.filter.set_measurement_noise(old_noise) |> ignore
  self.last_timestamp = timestamp
  self.initialized_timestamp = true
  result
}

///|
pub fn ConstantVelocityTracker2D::step_position(
  self : ConstantVelocityTracker2D,
  timestamp : Int,
  x : Double,
  y : Double,
) -> UpdateResult {
  self.step(timestamp, [x, y], Matrix::diagonal(2, self.position_variance))
}

///|
pub fn ConstantVelocityTracker2D::reset(
  self : ConstantVelocityTracker2D,
) -> Unit {
  self.filter.reset()
  self.last_timestamp = 0
  self.initialized_timestamp = false
}