#![forbid(unsafe_code)]
use std::collections::HashMap;
use std::sync::Mutex;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use crate::dsfb::drift::{MeasurementTracker, Regime};
use crate::dsfb::features::{Channel, ChunkKey};
use crate::dsfb::selection::{SearchPlan, SearchStrategy};
pub const DSFB_SHARDS: usize = 16;
#[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; Channel::ALL.len()],
weights: [f64; Channel::ALL.len()],
last_y: [f64; Channel::ALL.len()],
regime: Regime,
samples: u64,
winner: Channel,
baseline: Option<crate::core::extent::ChunkId>,
}
pub struct ShardedStorageObserver {
params: StorageDsfbParams,
shards: [Mutex<HashMap<ChunkKey, ChunkObserver>>; DSFB_SHARDS],
tracked: AtomicUsize,
steps: AtomicU64,
drift_events: AtomicU64,
slew_events: AtomicU64,
narrowed_searches: AtomicU64,
evict_cursor: AtomicUsize,
prior: Mutex<crate::dsfb::semantics::SemanticPrior>,
}
#[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 ShardedStorageObserver {
fn default() -> Self {
Self::new(StorageDsfbParams::default())
}
}
const FNV_OFFSET_BASIS: u64 = 0xcbf2_9ce4_8422_2325;
const FNV_PRIME: u64 = 0x0000_0100_0000_01b3;
impl ShardedStorageObserver {
pub fn new(params: StorageDsfbParams) -> Self {
Self {
params,
shards: std::array::from_fn(|_| Mutex::new(HashMap::new())),
tracked: AtomicUsize::new(0),
steps: AtomicU64::new(0),
drift_events: AtomicU64::new(0),
slew_events: AtomicU64::new(0),
narrowed_searches: AtomicU64::new(0),
evict_cursor: AtomicUsize::new(0),
prior: Mutex::new(crate::dsfb::semantics::SemanticPrior::default()),
}
}
pub fn shard_of(&self, key: &ChunkKey) -> usize {
let mut h = FNV_OFFSET_BASIS;
for b in key
.ino
.to_le_bytes()
.iter()
.chain(key.index.to_le_bytes().iter())
{
h ^= *b as u64;
h = h.wrapping_mul(FNV_PRIME);
}
for &b in key.content_id.as_bytes() {
h ^= b as u64;
h = h.wrapping_mul(FNV_PRIME);
}
(h % DSFB_SHARDS as u64) as usize
}
pub fn observe(
&self,
key: ChunkKey,
measurements: &[(Channel, f64)],
winner: Channel,
outcome_quality: f64,
semantic: Option<crate::dsfb::semantics::SemanticContext>,
mode: crate::dsfb::semantics::SemanticMode,
) -> Regime {
self.steps.fetch_add(1, Ordering::Relaxed);
let mut m = [0.0f64; Channel::ALL.len()];
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 shard = self.shard_of(&key);
let mut map = self.shards[shard].lock().expect("dsfb shard poisoned");
if !map.contains_key(&key) {
self.tracked.fetch_add(1, Ordering::Relaxed);
}
let entry = map.entry(key).or_insert_with(|| ChunkObserver {
inner: dsfb::DsfbObserver::new(self.params.dsfb_params(), Channel::ALL.len()),
tracker: MeasurementTracker::default(),
ema: [1.0; Channel::ALL.len()],
weights: [0.125; Channel::ALL.len()],
last_y: [0.0; Channel::ALL.len()],
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; Channel::ALL.len()];
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.drift_events.fetch_add(1, Ordering::Relaxed);
}
Regime::Slew => {
self.slew_events.fetch_add(1, Ordering::Relaxed);
}
Regime::Stable | Regime::Unknown => {}
}
entry.regime = regime;
if let Some(pkey) = semantic.and_then(|s| s.key_for(mode)) {
self.prior
.lock()
.expect("dsfb prior poisoned")
.observe(pkey, winner);
}
regime
}
pub fn trust(&self, key: &ChunkKey, channel: Channel) -> f64 {
let shard = self.shard_of(key);
self.shards[shard]
.lock()
.expect("dsfb shard poisoned")
.get(key)
.map(|c| c.weights[channel as usize])
.unwrap_or(0.125)
}
pub fn plan(
&self,
key: &ChunkKey,
semantic: Option<crate::dsfb::semantics::SemanticContext>,
mode: crate::dsfb::semantics::SemanticMode,
) -> SearchPlan {
let shard = self.shard_of(key);
let (regime, trust) = {
let map = self.shards[shard].lock().expect("dsfb shard poisoned");
let regime = map.get(key).map(|c| c.regime).unwrap_or(Regime::Unknown);
let mut trust: Vec<(Channel, f64)> = Channel::ALL
.iter()
.map(|&c| (c, self.trust_locked(&map, key, c)))
.collect();
drop(map);
if let Some(pkey) = semantic.and_then(|s| s.key_for(mode)) {
let prior = self.prior.lock().expect("dsfb prior poisoned");
for (c, t) in trust.iter_mut() {
*t += crate::dsfb::semantics::SEMANTIC_WEIGHT * prior.prior(pkey, *c);
}
}
trust.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
(regime, trust)
};
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(),
}
}
fn trust_locked(
&self,
map: &HashMap<ChunkKey, ChunkObserver>,
key: &ChunkKey,
channel: Channel,
) -> f64 {
map.get(key)
.map(|c| c.weights[channel as usize])
.unwrap_or(0.125)
}
pub fn class_rans_share(
&self,
semantic: Option<crate::dsfb::semantics::SemanticContext>,
mode: crate::dsfb::semantics::SemanticMode,
) -> Option<(u64, f64)> {
let pkey = semantic?.key_for(mode)?;
let prior = self.prior.lock().expect("dsfb prior poisoned");
let count = prior.count(pkey);
let share = prior.prior(pkey, Channel::Rans);
Some((count, share))
}
pub fn forget(&self, key: &ChunkKey) {
let shard = self.shard_of(key);
let mut map = self.shards[shard].lock().expect("dsfb shard poisoned");
if map.remove(key).is_some() {
self.tracked.fetch_sub(1, Ordering::Relaxed);
}
}
pub fn evict_one_from(&self, shard: usize) {
let mut map = self.shards[shard].lock().expect("dsfb shard poisoned");
if let Some(k) = map.keys().next().copied() {
map.remove(&k);
self.tracked.fetch_sub(1, Ordering::Relaxed);
}
}
pub fn evict_one(&self) {
let shard = self.evict_cursor.fetch_add(1, Ordering::Relaxed) % DSFB_SHARDS;
self.evict_one_from(shard);
}
pub fn len(&self) -> usize {
self.tracked.load(Ordering::Relaxed)
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn stats(&self) -> ObserverStats {
ObserverStats {
tracked_chunks: self.tracked.load(Ordering::Relaxed),
drift_events: self.drift_events.load(Ordering::Relaxed),
slew_events: self.slew_events.load(Ordering::Relaxed),
steps: self.steps.load(Ordering::Relaxed),
narrowed_searches: self.narrowed_searches.load(Ordering::Relaxed),
}
}
}
impl std::fmt::Debug for ShardedStorageObserver {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ShardedStorageObserver")
.field("params", &self.params)
.field("shards", &DSFB_SHARDS)
.field("tracked_chunks", &self.len())
.field("stats", &self.stats())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::extent::ChunkId;
use crate::dsfb::semantics::{SemanticContext, SemanticMode};
fn key(ino: u64, index: u64) -> ChunkKey {
ChunkKey::new(ino, index, ChunkId::of(&[ino as u8, index as u8]))
}
fn no_sem() -> (Option<SemanticContext>, SemanticMode) {
(None, SemanticMode::None)
}
#[test]
fn stable_evidence_keeps_trust_high() {
let obs = ShardedStorageObserver::default();
let k = key(1, 0);
let (sem, mode) = no_sem();
let mut regime = Regime::Unknown;
for _ in 0..50 {
regime = obs.observe(
k,
&[(Channel::PrevVersion, 1.0)],
Channel::PrevVersion,
1.0,
sem,
mode,
);
}
assert_eq!(regime, Regime::Stable);
let p0 = obs.trust(&k, Channel::PrevVersion);
assert!(p0 > 0.5, "p0 trust {p0}");
let plan = obs.plan(&k, sem, mode);
assert_eq!(plan.ordered_channels[0], Channel::PrevVersion);
assert_eq!(plan.strategy, SearchStrategy::Narrow);
assert_eq!(obs.len(), 1);
assert_eq!(obs.stats().steps, 50);
}
#[test]
fn slew_broadens_search() {
let obs = ShardedStorageObserver::default();
let k = key(2, 0);
let (sem, mode) = no_sem();
for _ in 0..10 {
obs.observe(
k,
&[(Channel::PrevVersion, 1.0)],
Channel::PrevVersion,
1.0,
sem,
mode,
);
}
assert_eq!(
obs.observe(
k,
&[(Channel::PrevVersion, 0.0)],
Channel::Raw,
0.0,
sem,
mode
),
Regime::Slew
);
let plan = obs.plan(&k, sem, mode);
assert_eq!(plan.strategy, SearchStrategy::Broad);
for _ in 0..20 {
obs.observe(
k,
&[(Channel::PrevVersion, 0.0)],
Channel::Raw,
0.5,
sem,
mode,
);
}
let final_plan = obs.plan(&k, sem, mode);
assert_ne!(final_plan.strategy, SearchStrategy::Narrow);
assert!(obs.stats().slew_events > 0);
}
#[test]
fn drift_keeps_narrow() {
let obs = ShardedStorageObserver::default();
let k = key(3, 0);
let (sem, mode) = no_sem();
let mut v = 1.0f64;
for _ in 0..40 {
obs.observe(
k,
&[(Channel::PrevVersion, v)],
Channel::PrevVersion,
v,
sem,
mode,
);
v = (v - 0.01).max(0.8);
}
assert!(obs.stats().drift_events > 0);
}
#[test]
fn eviction_bounds_state() {
let obs = ShardedStorageObserver::default();
let cap = 90usize;
let (sem, mode) = no_sem();
for i in 0..cap {
obs.observe(key(i as u64, 0), &[], Channel::Raw, 0.5, sem, mode);
}
assert_eq!(obs.len(), cap);
for i in cap..500 {
let k = key(i as u64, 0);
obs.observe(k, &[], Channel::Raw, 0.5, sem, mode);
if obs.len() > cap {
obs.evict_one_from(obs.shard_of(&k));
}
assert!(
obs.len() <= cap,
"targeted eviction must keep the total at the cap (got {})",
obs.len()
);
}
assert_eq!(obs.len(), cap);
}
#[test]
fn evict_one_rotates_until_drained() {
let obs = ShardedStorageObserver::default();
let (sem, mode) = no_sem();
for i in 0..16 {
obs.observe(key(i as u64, 0), &[], Channel::Raw, 0.5, sem, mode);
}
assert_eq!(obs.len(), 16);
let mut guard = 0;
while !obs.is_empty() {
obs.evict_one();
guard += 1;
assert!(guard <= 16 * 16, "rotating eviction failed to drain");
}
obs.evict_one();
assert!(obs.is_empty());
assert_eq!(obs.stats().tracked_chunks, 0);
}
#[test]
fn shard_of_is_stable_and_bounded() {
let obs = ShardedStorageObserver::default();
let k = key(7, 3);
for _ in 0..10 {
assert_eq!(obs.shard_of(&k), obs.shard_of(&k));
}
assert!(obs.shard_of(&k) < DSFB_SHARDS);
let mut seen = [false; DSFB_SHARDS];
for i in 0..4096u64 {
seen[obs.shard_of(&key(i, 0))] = true;
}
assert!(
seen.iter().all(|&s| s),
"shard hash must spread across all {DSFB_SHARDS} shards"
);
}
#[test]
fn forget_and_targeted_eviction_keep_count_exact() {
let obs = ShardedStorageObserver::default();
let (sem, mode) = no_sem();
for i in 0..64 {
obs.observe(key(i as u64, 0), &[], Channel::Raw, 0.5, sem, mode);
}
assert_eq!(obs.len(), 64);
obs.forget(&key(0, 0));
assert_eq!(obs.len(), 63);
obs.forget(&key(0, 0));
assert_eq!(obs.len(), 63);
let k = key(30, 0);
let shard = obs.shard_of(&k);
obs.evict_one_from(shard);
assert_eq!(obs.len(), 62);
obs.observe(key(0, 0), &[], Channel::Raw, 0.5, sem, mode);
assert_eq!(obs.len(), 63);
}
#[test]
fn concurrent_observation_is_safe() {
let obs = ShardedStorageObserver::default();
let (sem, mode) = no_sem();
std::thread::scope(|s| {
for t in 0..16usize {
let obs = &obs;
s.spawn(move || {
for i in 0..256usize {
let ino = (t * 256 + i) as u64;
obs.observe(key(ino, 0), &[], Channel::Raw, 0.5, sem, mode);
let _ = obs.plan(&key(ino, 0), sem, mode);
let _ = obs.trust(&key(ino, 0), Channel::Raw);
}
});
}
});
assert_eq!(obs.len(), 4096, "exact count under concurrency");
assert_eq!(obs.stats().steps, 4096);
let k = key(0, 0);
obs.observe(k, &[], Channel::Raw, 0.5, sem, mode);
assert_eq!(obs.len(), 4096);
assert_eq!(obs.stats().steps, 4097);
}
}