///|
/// A robust residual policy used before a measurement enters a filter.
pub(all) enum ResidualAction {
  Accept
  Downweight(Double)
  Reject
} derive(Debug, Eq)

///|
pub struct ResidualPolicy {
  soft_limit : Double
  hard_limit : Double
  minimum_weight : Double
  mut accepted : Int
  mut downweighted : Int
  mut rejected : Int
}

///|
pub fn ResidualPolicy::new(
  soft_limit : Double,
  hard_limit : Double,
  minimum_weight : Double,
) -> ResidualPolicy {
  let safe_soft = if soft_limit <= 0.0 { 1.0 } else { soft_limit }
  let safe_hard = if hard_limit <= safe_soft {
    safe_soft * 2.0
  } else {
    hard_limit
  }
  let safe_minimum = if minimum_weight <= 0.0 {
    0.01
  } else if minimum_weight > 1.0 {
    1.0
  } else {
    minimum_weight
  }
  {
    soft_limit: safe_soft,
    hard_limit: safe_hard,
    minimum_weight: safe_minimum,
    accepted: 0,
    downweighted: 0,
    rejected: 0,
  }
}

///|
pub fn ResidualPolicy::classify(
  self : ResidualPolicy,
  residual : Double,
) -> ResidualAction {
  if residual.is_nan() || residual.is_inf() {
    self.rejected = self.rejected + 1
    return Reject
  }
  let magnitude = residual.abs()
  if magnitude <= self.soft_limit {
    self.accepted = self.accepted + 1
    Accept
  } else if magnitude >= self.hard_limit {
    self.rejected = self.rejected + 1
    Reject
  } else {
    let weight = self.soft_limit / magnitude
    let safe_weight = if weight < self.minimum_weight {
      self.minimum_weight
    } else {
      weight
    }
    self.downweighted = self.downweighted + 1
    Downweight(safe_weight)
  }
}

///|
pub fn ResidualPolicy::soft_limit(self : ResidualPolicy) -> Double {
  self.soft_limit
}

///|
pub fn ResidualPolicy::hard_limit(self : ResidualPolicy) -> Double {
  self.hard_limit
}

///|
pub fn ResidualPolicy::minimum_weight(self : ResidualPolicy) -> Double {
  self.minimum_weight
}

///|
pub fn ResidualPolicy::accepted(self : ResidualPolicy) -> Int {
  self.accepted
}

///|
pub fn ResidualPolicy::downweighted(self : ResidualPolicy) -> Int {
  self.downweighted
}

///|
pub fn ResidualPolicy::rejected(self : ResidualPolicy) -> Int {
  self.rejected
}

///|
pub fn ResidualPolicy::reset(self : ResidualPolicy) -> Unit {
  self.accepted = 0
  self.downweighted = 0
  self.rejected = 0
}

///|
/// Apply a scalar policy to a residual vector.  The smallest component weight
/// is used for the complete observation, which is conservative for correlated
/// sensors and avoids accidentally trusting a partially corrupted packet.
pub fn classify_residual_vector(
  policy : ResidualPolicy,
  residual : Array[Double],
) -> ResidualAction {
  if residual.length() == 0 {
    return Reject
  }
  let mut strongest = 0.0
  for value in residual {
    if value.is_nan() || value.is_inf() {
      return policy.classify(value)
    }
    if value.abs() > strongest {
      strongest = value.abs()
    }
  }
  policy.classify(strongest)
}

///|
pub fn inflate_for_residual(
  covariance : Matrix,
  residual : Array[Double],
  policy : ResidualPolicy,
) -> Matrix {
  let action = classify_residual_vector(policy, residual)
  match action {
    Accept => covariance.copy()
    Downweight(weight) => covariance.scale(1.0 / (weight * weight))
    Reject => covariance.add_diagonal(1000000.0)
  }
}

///|
/// A slowly adapting process-noise schedule.  It reacts to sustained large
/// innovations while decaying toward the nominal factor after recovery.
pub struct NoiseSchedule {
  nominal : Double
  minimum : Double
  maximum : Double
  growth : Double
  decay : Double
  mut factor : Double
  mut stressed_steps : Int
}

///|
pub fn NoiseSchedule::new(
  nominal : Double,
  minimum : Double,
  maximum : Double,
  growth : Double,
  decay : Double,
) -> NoiseSchedule {
  let safe_nominal = if nominal <= 0.0 { 1.0 } else { nominal }
  let safe_minimum = if minimum <= 0.0 { 0.1 } else { minimum }
  let safe_maximum = if maximum < safe_nominal { safe_nominal } else { maximum }
  {
    nominal: safe_nominal,
    minimum: if safe_minimum > safe_nominal {
      safe_nominal
    } else {
      safe_minimum
    },
    maximum: safe_maximum,
    growth: if growth <= 0.0 {
      0.25
    } else {
      growth
    },
    decay: if decay <= 0.0 || decay > 1.0 {
      0.1
    } else {
      decay
    },
    factor: safe_nominal,
    stressed_steps: 0,
  }
}

///|
pub fn NoiseSchedule::observe(
  self : NoiseSchedule,
  nis : Double,
  threshold : Double,
) -> Double {
  let safe_threshold = if threshold <= 0.0 { 1.0 } else { threshold }
  if nis.is_nan() || nis.is_inf() || nis > safe_threshold {
    self.factor = self.factor * (1.0 + self.growth)
    if self.factor > self.maximum {
      self.factor = self.maximum
    }
    self.stressed_steps = self.stressed_steps + 1
  } else {
    self.factor = self.factor + self.decay * (self.nominal - self.factor)
    if self.factor < self.minimum {
      self.factor = self.minimum
    }
    if self.stressed_steps > 0 {
      self.stressed_steps = self.stressed_steps - 1
    }
  }
  self.factor
}

///|
pub fn NoiseSchedule::factor(self : NoiseSchedule) -> Double {
  self.factor
}

///|
pub fn NoiseSchedule::stressed_steps(self : NoiseSchedule) -> Int {
  self.stressed_steps
}

///|
pub fn NoiseSchedule::reset(self : NoiseSchedule) -> Unit {
  self.factor = self.nominal
  self.stressed_steps = 0
}

///|
/// A small batch gate that turns individual residual decisions into one
/// packet-level decision and exposes a stable covariance inflation factor.
pub struct BatchGate {
  policy : ResidualPolicy
  mut inspected : Int
  mut accepted : Int
  mut downweighted : Int
  mut rejected : Int
}

///|
pub fn BatchGate::new(policy : ResidualPolicy) -> BatchGate {
  { policy, inspected: 0, accepted: 0, downweighted: 0, rejected: 0 }
}

///|
pub fn BatchGate::inspect(
  self : BatchGate,
  residual : Array[Double],
) -> ResidualAction {
  self.inspected = self.inspected + 1
  let action = classify_residual_vector(self.policy, residual)
  match action {
    Accept => self.accepted = self.accepted + 1
    Downweight(_) => self.downweighted = self.downweighted + 1
    Reject => self.rejected = self.rejected + 1
  }
  action
}

///|
pub fn BatchGate::inflation(self : BatchGate) -> Double {
  if self.rejected > 0 {
    1000000.0
  } else if self.downweighted > 0 {
    4.0
  } else {
    1.0
  }
}

///|
pub fn BatchGate::inspected(self : BatchGate) -> Int {
  self.inspected
}

///|
pub fn BatchGate::accepted(self : BatchGate) -> Int {
  self.accepted
}

///|
pub fn BatchGate::downweighted(self : BatchGate) -> Int {
  self.downweighted
}

///|
pub fn BatchGate::rejected(self : BatchGate) -> Int {
  self.rejected
}

///|
pub fn BatchGate::reset(self : BatchGate) -> Unit {
  self.inspected = 0
  self.accepted = 0
  self.downweighted = 0
  self.rejected = 0
}