#![forbid(unsafe_code)]
use dsfb::DsfbState;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Regime {
Unknown,
Stable,
Drift,
Slew,
}
pub fn classify(state: DsfbState, slew_alpha: f64, drift_omega: f64) -> Regime {
if state.alpha.abs() > slew_alpha {
Regime::Slew
} else if state.omega.abs() > drift_omega {
Regime::Drift
} else {
Regime::Stable
}
}
#[derive(Debug, Clone, Copy)]
pub struct MeasurementTracker {
ema: f64,
delta_ema: f64,
samples: u64,
slew_window: u64,
pre_slew_ema: f64,
}
pub const SLEW_DELTA_THRESHOLD: f64 = 0.25;
pub const DRIFT_RATE_THRESHOLD: f64 = 0.004;
pub const SLEW_WINDOW: u64 = 8;
pub const SLEW_RECOVERY_EPS: f64 = 0.1;
impl Default for MeasurementTracker {
fn default() -> Self {
Self {
ema: 0.0,
delta_ema: 0.0,
samples: 0,
slew_window: 0,
pre_slew_ema: 0.0,
}
}
}
impl MeasurementTracker {
pub fn observe(&mut self, m: f64) -> Regime {
self.samples += 1;
if self.samples == 1 {
self.ema = m;
return Regime::Unknown;
}
let delta = m - self.ema;
self.delta_ema = 0.9 * self.delta_ema + 0.1 * delta;
self.ema = 0.8 * self.ema + 0.2 * m;
if self.slew_window > 0 {
self.slew_window -= 1;
if (m - self.pre_slew_ema).abs() < SLEW_RECOVERY_EPS {
self.slew_window = 0;
} else if self.slew_window > 0 {
return Regime::Slew;
}
}
if self.samples > 4 && delta.abs() > SLEW_DELTA_THRESHOLD {
self.slew_window = SLEW_WINDOW;
self.pre_slew_ema = self.ema;
return Regime::Slew;
}
if self.samples > 4 && self.delta_ema.abs() > DRIFT_RATE_THRESHOLD {
Regime::Drift
} else {
Regime::Stable
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Guidance {
KeepBasis,
UpdateResiduals,
NewBaseline,
}
impl Guidance {
pub const fn for_regime(regime: Regime) -> Guidance {
match regime {
Regime::Stable => Guidance::KeepBasis,
Regime::Drift => Guidance::UpdateResiduals,
Regime::Slew => Guidance::NewBaseline,
Regime::Unknown => Guidance::UpdateResiduals,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn classification() {
assert_eq!(
classify(DsfbState::new(0.0, 0.001, 0.0), 0.05, 0.02),
Regime::Stable
);
assert_eq!(
classify(DsfbState::new(0.0, 0.05, 0.0), 0.05, 0.02),
Regime::Drift
);
assert_eq!(
classify(DsfbState::new(0.0, 0.05, 0.1), 0.05, 0.02),
Regime::Slew
);
}
#[test]
fn tracker_stable_on_constant() {
let mut t = MeasurementTracker::default();
assert_eq!(t.observe(1.0), Regime::Unknown);
for _ in 0..20 {
assert_eq!(t.observe(1.0), Regime::Stable);
}
}
#[test]
fn tracker_drift_on_gradual_decline() {
let mut t = MeasurementTracker::default();
t.observe(1.0);
let mut v = 1.0;
let mut saw_drift = false;
for _ in 0..40 {
v -= 0.01;
if t.observe(v) == Regime::Drift {
saw_drift = true;
}
}
assert!(saw_drift);
}
#[test]
fn tracker_slew_on_jump_with_persistence() {
let mut t = MeasurementTracker::default();
t.observe(1.0);
for _ in 0..10 {
t.observe(1.0);
}
assert_eq!(t.observe(0.0), Regime::Slew);
for _ in 0..3 {
assert_eq!(t.observe(0.0), Regime::Slew);
}
let mut saw_exit = false;
for _ in 0..20 {
if t.observe(0.0) != Regime::Slew {
saw_exit = true;
}
}
assert!(saw_exit);
}
}