use super::{DriftReport, DriftVerdict, FeatureDrift, MetricKind, MetricScore};
use crate::alert::{Alerter, DriftAlertEvent, NopAlerter};
use crate::distribution::{Comparison, FeatureKind, LiveFeature, ReferenceDistribution};
use crate::error::{DriftError, Result};
use crate::metrics::{chi_square_test, js_divergence, kl_divergence, ks_test, psi};
pub const DEFAULT_DATASET_FRACTION_THRESHOLD: f64 = 0.5;
pub const DEFAULT_PSI_THRESHOLD: f64 = 0.25;
pub const DEFAULT_ALPHA: f64 = 0.05;
#[derive(Clone, Debug)]
pub struct FeatureConfig {
pub primary: MetricKind,
pub threshold: f64,
pub additional: Vec<MetricKind>,
}
impl FeatureConfig {
pub fn default_for(kind: FeatureKind) -> Self {
match kind {
FeatureKind::Continuous => FeatureConfig {
primary: MetricKind::Psi,
threshold: DEFAULT_PSI_THRESHOLD,
additional: vec![MetricKind::Ks],
},
FeatureKind::Categorical => FeatureConfig {
primary: MetricKind::Psi,
threshold: DEFAULT_PSI_THRESHOLD,
additional: vec![MetricKind::ChiSquare],
},
}
}
fn metrics(&self) -> Vec<MetricKind> {
let mut all = vec![self.primary];
for &m in &self.additional {
if !all.contains(&m) {
all.push(m);
}
}
all
}
}
struct FeatureEntry {
reference: ReferenceDistribution,
config: FeatureConfig,
}
pub struct DatasetMonitor {
features: Vec<FeatureEntry>,
dataset_fraction_threshold: f64,
min_live_samples: usize,
alerter: Box<dyn Alerter>,
}
impl Default for DatasetMonitor {
fn default() -> Self {
Self::new()
}
}
impl DatasetMonitor {
pub fn new() -> Self {
Self {
features: Vec::new(),
dataset_fraction_threshold: DEFAULT_DATASET_FRACTION_THRESHOLD,
min_live_samples: 2,
alerter: Box::new(NopAlerter),
}
}
pub fn add_feature(&mut self, reference: ReferenceDistribution) -> &mut Self {
let config = FeatureConfig::default_for(reference.kind());
self.features.push(FeatureEntry { reference, config });
self
}
pub fn add_feature_with_config(
&mut self,
reference: ReferenceDistribution,
config: FeatureConfig,
) -> &mut Self {
self.features.push(FeatureEntry { reference, config });
self
}
pub fn with_dataset_fraction_threshold(mut self, fraction: f64) -> Self {
self.dataset_fraction_threshold = fraction;
self
}
pub fn with_min_live_samples(mut self, min: usize) -> Self {
self.min_live_samples = min;
self
}
pub fn with_alerter(mut self, alerter: Box<dyn Alerter>) -> Self {
self.alerter = alerter;
self
}
pub fn len(&self) -> usize {
self.features.len()
}
pub fn is_empty(&self) -> bool {
self.features.is_empty()
}
pub fn check(&self, live: &[(&str, LiveFeature)]) -> Result<DriftReport> {
let mut feature_reports = Vec::with_capacity(self.features.len());
for entry in &self.features {
let name = entry.reference.name();
let live_feature = live
.iter()
.find(|(n, _)| *n == name)
.map(|(_, f)| *f)
.ok_or_else(|| DriftError::UnknownFeature(name.to_string()))?;
let n_samples = match live_feature {
LiveFeature::Continuous(s) => s.len(),
LiveFeature::Categorical(s) => s.len(),
};
if n_samples < self.min_live_samples {
return Err(DriftError::SampleTooSmall {
kind: "drift check",
minimum: self.min_live_samples,
actual: n_samples,
});
}
let comparison = entry.reference.compare(live_feature)?;
let scores = self.compute_scores(&entry.config, &comparison)?;
let verdict = feature_verdict(&entry.config, &scores);
feature_reports.push(FeatureDrift {
feature: name.to_string(),
kind: comparison.kind,
scores,
primary: entry.config.primary,
threshold: entry.config.threshold,
verdict,
reference_histogram: comparison.reference_hist,
live_histogram: comparison.live_hist,
});
}
let report = DriftReport {
features: feature_reports,
dataset_fraction_threshold: self.dataset_fraction_threshold,
};
#[cfg(feature = "prometheus-export")]
crate::export::record_report(&report);
if report.dataset_drift_detected() {
self.alerter.alert(&DriftAlertEvent::from_report(&report));
}
Ok(report)
}
fn compute_scores(
&self,
config: &FeatureConfig,
comparison: &Comparison,
) -> Result<Vec<MetricScore>> {
let mut scores = Vec::new();
for kind in config.metrics() {
match compute_metric(kind, comparison) {
Ok(score) => scores.push(score),
Err(DriftError::SampleTooSmall { .. }) if kind != config.primary => {}
Err(e) => return Err(e),
}
}
Ok(scores)
}
}
fn compute_metric(kind: MetricKind, comparison: &Comparison) -> Result<MetricScore> {
let (statistic, p_value) = match kind {
MetricKind::Psi => (
psi(&comparison.reference_hist, &comparison.live_hist)?,
None,
),
MetricKind::Kl => (
kl_divergence(&comparison.reference_hist, &comparison.live_hist)?,
None,
),
MetricKind::Js => (
js_divergence(&comparison.reference_hist, &comparison.live_hist)?,
None,
),
MetricKind::Ks => {
let (reference, live) = comparison.raw_samples.ok_or_else(|| {
DriftError::InvalidConfig(
"KS test requires continuous raw samples; not available for categorical features"
.into(),
)
})?;
let r = ks_test(reference, live)?;
(r.statistic, Some(r.p_value))
}
MetricKind::ChiSquare => {
let r = chi_square_test(&comparison.reference_hist, &comparison.live_hist)?;
(r.statistic, Some(r.p_value))
}
};
Ok(MetricScore {
kind,
statistic,
p_value,
})
}
fn feature_verdict(config: &FeatureConfig, scores: &[MetricScore]) -> DriftVerdict {
let primary = scores.iter().find(|s| s.kind == config.primary);
let drifted = match primary {
Some(score) => {
if config.primary.higher_is_more_drift() {
score.statistic > config.threshold
} else {
score.p_value.map(|p| p < config.threshold).unwrap_or(false)
}
}
None => false,
};
if drifted {
DriftVerdict::Drifted
} else {
DriftVerdict::Stable
}
}