use std::collections::HashMap;
use std::fmt::Write;
#[derive(Debug, Clone)]
pub enum SamplingStrategy {
None,
Random { size: usize },
Reservoir { size: usize },
Stratified {
key_columns: Vec<String>,
samples_per_stratum: usize,
},
Progressive {
initial_size: usize,
confidence_level: f64,
max_size: usize,
},
Systematic { interval: usize },
Importance {
weight_column: String,
weight_threshold: f64,
},
MultiStage { stages: Vec<SamplingStrategy> },
}
#[derive(Debug, Default)]
pub struct SamplingState {
progressive_samples: usize,
stratum_samples: HashMap<String, usize>,
}
impl SamplingState {
pub fn new() -> Self {
Self::default()
}
pub fn progressive_taken(&self) -> usize {
self.progressive_samples
}
pub fn record_progressive(&mut self) {
self.progressive_samples += 1;
}
pub fn take_from_stratum(
&mut self,
row: super::sampler::RowView<'_>,
key_columns: &[String],
samples_per_stratum: usize,
) -> bool {
if key_columns.is_empty() {
return false;
}
let mut stratum = String::new();
for column in key_columns {
let Some(value) = row.get(column) else {
return false;
};
let _ = write!(stratum, "{}:{}", value.len(), value);
}
let taken = self.stratum_samples.entry(stratum).or_insert(0);
if *taken < samples_per_stratum {
*taken += 1;
true
} else {
false
}
}
pub fn strata_seen(&self) -> usize {
self.stratum_samples.len()
}
}
impl SamplingStrategy {
pub fn adaptive(total_rows: Option<usize>, file_size_mb: f64) -> Self {
match (total_rows, file_size_mb) {
(Some(rows), size_mb) if rows <= 10_000 && size_mb < 10.0 => SamplingStrategy::None,
(Some(rows), _) if rows <= 100_000 => SamplingStrategy::Random { size: 10_000 },
(Some(rows), _) if rows <= 1_000_000 => SamplingStrategy::Progressive {
initial_size: 10_000,
confidence_level: 0.95,
max_size: 50_000,
},
(_, size_mb) if size_mb > 1000.0 => SamplingStrategy::MultiStage {
stages: vec![
SamplingStrategy::Systematic { interval: 100 },
SamplingStrategy::Progressive {
initial_size: 5_000,
confidence_level: 0.99,
max_size: 25_000,
},
],
},
_ => SamplingStrategy::Reservoir { size: 100_000 },
}
}
pub fn stratified(key_columns: Vec<String>, samples_per_stratum: usize) -> Self {
Self::Stratified {
key_columns,
samples_per_stratum,
}
}
pub fn importance(weight_column: impl Into<String>, weight_threshold: f64) -> Self {
Self::Importance {
weight_column: weight_column.into(),
weight_threshold,
}
}
pub fn target_sample_size(&self) -> Option<usize> {
match self {
SamplingStrategy::None => None,
SamplingStrategy::Random { size } => Some(*size),
SamplingStrategy::Reservoir { size } => Some(*size),
SamplingStrategy::Stratified {
samples_per_stratum,
..
} => Some(*samples_per_stratum),
SamplingStrategy::Progressive { max_size, .. } => Some(*max_size),
SamplingStrategy::Systematic { .. } => None,
SamplingStrategy::Importance { .. } => None,
SamplingStrategy::MultiStage { stages } => {
stages.iter().filter_map(|s| s.target_sample_size()).min()
}
}
}
pub fn description(&self) -> String {
match self {
SamplingStrategy::None => "Full dataset analysis".to_string(),
SamplingStrategy::Random { size } => format!("Random sampling ({} records)", size),
SamplingStrategy::Reservoir { size } => {
format!("Reservoir sampling ({} records)", size)
}
SamplingStrategy::Stratified {
key_columns,
samples_per_stratum,
} => {
format!(
"Stratified by {} ({} per stratum)",
key_columns.join(", "),
samples_per_stratum
)
}
SamplingStrategy::Progressive {
initial_size,
confidence_level,
max_size,
} => {
format!(
"Progressive sampling ({}-{} records, {}% confidence)",
initial_size,
max_size,
(confidence_level * 100.0) as u8
)
}
SamplingStrategy::Systematic { interval } => {
format!("Systematic (every {}th record)", interval)
}
SamplingStrategy::Importance {
weight_column,
weight_threshold,
} => {
format!("Importance filter ({weight_column} >= {weight_threshold:.2})")
}
SamplingStrategy::MultiStage { stages } => {
format!("Multi-stage ({} stages)", stages.len())
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sampling::sampler::RowView;
#[test]
fn test_stratum_state_persists_across_rows() {
let headers = vec!["region".to_string()];
let mut state = SamplingState::new();
let keys = vec!["region".to_string()];
let north = vec!["north".to_string()];
let south = vec!["south".to_string()];
assert!(state.take_from_stratum(RowView::new(&headers, &north), &keys, 2));
assert!(state.take_from_stratum(RowView::new(&headers, &north), &keys, 2));
assert!(
!state.take_from_stratum(RowView::new(&headers, &north), &keys, 2),
"the third northern row exceeds the per-stratum cap"
);
assert!(
state.take_from_stratum(RowView::new(&headers, &south), &keys, 2),
"a different stratum has its own budget"
);
assert_eq!(state.strata_seen(), 2);
}
#[test]
fn test_stratum_requires_every_key_column() {
let headers = vec!["region".to_string()];
let values = vec!["north".to_string()];
let mut state = SamplingState::new();
let keys = vec!["region".to_string(), "segment".to_string()];
assert!(
!state.take_from_stratum(RowView::new(&headers, &values), &keys, 5),
"a row missing a key column belongs to no stratum"
);
assert_eq!(state.strata_seen(), 0);
}
#[test]
fn test_importance_constructor_names_its_column() {
let strategy = SamplingStrategy::importance("risk", 0.8);
match strategy {
SamplingStrategy::Importance {
ref weight_column,
weight_threshold,
} => {
assert_eq!(weight_column, "risk");
assert_eq!(weight_threshold, 0.8);
}
other => panic!("expected an importance filter, got {other:?}"),
}
assert!(strategy.description().contains("risk"));
}
#[test]
fn test_adaptive_strategy() {
let small = SamplingStrategy::adaptive(Some(5_000), 1.0);
matches!(small, SamplingStrategy::None);
let medium = SamplingStrategy::adaptive(Some(50_000), 10.0);
matches!(medium, SamplingStrategy::Random { .. });
let large = SamplingStrategy::adaptive(Some(10_000_000), 2000.0);
matches!(large, SamplingStrategy::MultiStage { .. });
}
}