#![forbid(unsafe_code)]
use std::collections::HashMap;
use crate::dsfb::drift::{MeasurementTracker, Regime};
use crate::dsfb::features::{Channel, ChunkKey};
use crate::dsfb::selection::{SearchPlan, SearchStrategy};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct StorageDsfbParams {
pub k_phi: f64,
pub k_omega: f64,
pub k_alpha: f64,
pub rho: f64,
pub sigma0: f64,
pub dt: f64,
pub slew_alpha_threshold: f64,
pub drift_omega_threshold: f64,
}
impl Default for StorageDsfbParams {
fn default() -> Self {
Self {
k_phi: 0.5,
k_omega: 0.1,
k_alpha: 0.02,
rho: 0.9,
sigma0: 0.1,
dt: 1.0,
slew_alpha_threshold: 0.05,
drift_omega_threshold: 0.02,
}
}
}
impl StorageDsfbParams {
pub fn dsfb_params(&self) -> dsfb::DsfbParams {
dsfb::DsfbParams::new(
self.k_phi,
self.k_omega,
self.k_alpha,
self.rho,
self.sigma0,
)
}
}
struct ChunkObserver {
inner: dsfb::DsfbObserver,
tracker: MeasurementTracker,
ema: [f64; 8],
weights: [f64; 8],
last_y: [f64; 8],
regime: Regime,
samples: u64,
winner: Channel,
baseline: Option<crate::core::extent::ChunkId>,
}
pub struct StorageObserver {
params: StorageDsfbParams,
chunks: HashMap<ChunkKey, ChunkObserver>,
pub stats: ObserverStats,
}
impl std::fmt::Debug for StorageObserver {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("StorageObserver")
.field("params", &self.params)
.field("tracked_chunks", &self.chunks.len())
.field("stats", &self.stats)
.finish()
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq)]
pub struct ObserverStats {
pub tracked_chunks: usize,
pub drift_events: u64,
pub slew_events: u64,
pub steps: u64,
pub narrowed_searches: u64,
}
impl Default for StorageObserver {
fn default() -> Self {
Self::new(StorageDsfbParams::default())
}
}
impl StorageObserver {
pub fn new(params: StorageDsfbParams) -> Self {
Self {
params,
chunks: HashMap::new(),
stats: ObserverStats::default(),
}
}
pub fn observe(
&mut self,
key: ChunkKey,
measurements: &[(Channel, f64)],
winner: Channel,
outcome_quality: f64,
) -> Regime {
self.stats.steps += 1;
let mut m = [0.0f64; 8];
for &(c, v) in measurements {
m[c as usize] = v.clamp(0.0, 1.0);
}
let winner_measurement = outcome_quality.clamp(0.0, 1.0);
let entry = self.chunks.entry(key).or_insert_with(|| ChunkObserver {
inner: dsfb::DsfbObserver::new(self.params.dsfb_params(), 8),
tracker: MeasurementTracker::default(),
ema: [1.0; 8],
weights: [0.125; 8],
last_y: [0.0; 8],
regime: Regime::Unknown,
samples: 0,
winner: Channel::Raw,
baseline: None,
});
entry.samples += 1;
entry.winner = winner;
if entry.regime == Regime::Slew || entry.baseline.is_none() {
entry.baseline = Some(key.content_id);
}
let rho = self.params.rho;
for &(c, v) in measurements {
let idx = c as usize;
entry.last_y[idx] = v.clamp(0.0, 1.0);
let err = (1.0 - entry.last_y[idx]).abs();
entry.ema[idx] = rho * entry.ema[idx] + (1.0 - rho) * err;
}
let residuals = entry.ema;
let mut scratch = [0.0f64; 8];
let w =
dsfb::trust::calculate_trust_weights(&residuals, &mut scratch, 0.0, self.params.sigma0);
for (k, &wk) in w.iter().enumerate() {
entry.weights[k] = wk;
}
entry.inner.step(&entry.last_y, self.params.dt);
let regime = entry.tracker.observe(winner_measurement);
match regime {
Regime::Drift => self.stats.drift_events += 1,
Regime::Slew => self.stats.slew_events += 1,
Regime::Stable | Regime::Unknown => {}
}
entry.regime = regime;
self.stats.tracked_chunks = self.chunks.len();
regime
}
pub fn trust(&self, key: &ChunkKey, channel: Channel) -> f64 {
self.chunks
.get(key)
.map(|c| c.weights[channel as usize])
.unwrap_or(0.125)
}
pub fn plan(&self, key: &ChunkKey) -> SearchPlan {
let regime = self
.chunks
.get(key)
.map(|c| c.regime)
.unwrap_or(Regime::Unknown);
let mut trust: Vec<(Channel, f64)> = Channel::ALL
.iter()
.map(|&c| (c, self.trust(key, c)))
.collect();
trust.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
let strategy = match regime {
Regime::Slew => SearchStrategy::Broad,
Regime::Drift => SearchStrategy::Balanced,
Regime::Stable | Regime::Unknown => SearchStrategy::Narrow,
};
SearchPlan {
ordered_channels: trust.into_iter().map(|(c, _)| c).collect(),
strategy,
budget: strategy.budget(),
}
}
pub fn forget(&mut self, key: &ChunkKey) {
self.chunks.remove(key);
self.stats.tracked_chunks = self.chunks.len();
}
pub fn evict_one(&mut self) {
if let Some(k) = self.chunks.keys().next().copied() {
self.chunks.remove(&k);
self.stats.tracked_chunks = self.chunks.len();
}
}
pub fn len(&self) -> usize {
self.chunks.len()
}
pub fn is_empty(&self) -> bool {
self.chunks.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::extent::ChunkId;
fn key(ino: u64, index: u64) -> ChunkKey {
ChunkKey::new(ino, index, ChunkId::of(&[ino as u8, index as u8]))
}
#[test]
fn stable_evidence_keeps_trust_high() {
let mut obs = StorageObserver::default();
let k = key(1, 0);
let mut regime = Regime::Unknown;
for _ in 0..50 {
regime = obs.observe(k, &[(Channel::PrevVersion, 1.0)], Channel::PrevVersion, 1.0);
}
assert_eq!(regime, Regime::Stable);
let p0 = obs.trust(&k, Channel::PrevVersion);
assert!(p0 > 0.5, "p0 trust {p0}");
let plan = obs.plan(&k);
assert_eq!(plan.ordered_channels[0], Channel::PrevVersion);
assert_eq!(plan.strategy, SearchStrategy::Narrow);
}
#[test]
fn slew_broadens_search() {
let mut obs = StorageObserver::default();
let k = key(2, 0);
for _ in 0..10 {
obs.observe(k, &[(Channel::PrevVersion, 1.0)], Channel::PrevVersion, 1.0);
}
assert_eq!(
obs.observe(k, &[(Channel::PrevVersion, 0.0)], Channel::Raw, 0.0),
Regime::Slew
);
let plan = obs.plan(&k);
assert_eq!(plan.strategy, SearchStrategy::Broad);
for _ in 0..20 {
obs.observe(k, &[(Channel::PrevVersion, 0.0)], Channel::Raw, 0.5);
}
let final_plan = obs.plan(&k);
assert_ne!(final_plan.strategy, SearchStrategy::Narrow);
assert!(obs.stats.slew_events > 0);
}
#[test]
fn drift_keeps_narrow() {
let mut obs = StorageObserver::default();
let k = key(3, 0);
let mut v = 1.0f64;
for _ in 0..40 {
obs.observe(k, &[(Channel::PrevVersion, v)], Channel::PrevVersion, v);
v = (v - 0.01).max(0.8);
}
assert!(obs.stats.drift_events > 0);
}
#[test]
fn eviction_bounds_state() {
let mut obs = StorageObserver::default();
for i in 0..100 {
obs.observe(key(i, 0), &[], Channel::Raw, 0.5);
}
assert_eq!(obs.len(), 100);
for _ in 0..100 {
obs.evict_one();
}
assert!(obs.is_empty());
}
}