1use std::collections::HashMap;
2use std::fmt::Write;
3
4#[derive(Debug, Clone)]
15pub enum SamplingStrategy {
16 None,
18
19 Random { size: usize },
27
28 Reservoir { size: usize },
34
35 Stratified {
41 key_columns: Vec<String>,
42 samples_per_stratum: usize,
43 },
44
45 Progressive {
57 initial_size: usize,
58 confidence_level: f64,
59 max_size: usize,
60 },
61
62 Systematic { interval: usize },
64
65 Importance {
74 weight_column: String,
75 weight_threshold: f64,
76 },
77
78 MultiStage { stages: Vec<SamplingStrategy> },
84}
85
86#[derive(Debug, Default)]
93pub struct SamplingState {
94 progressive_samples: usize,
96
97 stratum_samples: HashMap<String, usize>,
99}
100
101impl SamplingState {
102 pub fn new() -> Self {
103 Self::default()
104 }
105
106 pub fn progressive_taken(&self) -> usize {
108 self.progressive_samples
109 }
110
111 pub fn record_progressive(&mut self) {
113 self.progressive_samples += 1;
114 }
115
116 pub fn take_from_stratum(
122 &mut self,
123 row: super::sampler::RowView<'_>,
124 key_columns: &[String],
125 samples_per_stratum: usize,
126 ) -> bool {
127 if key_columns.is_empty() {
128 return false;
129 }
130
131 let mut stratum = String::new();
135 for column in key_columns {
136 let Some(value) = row.get(column) else {
137 return false;
138 };
139 let _ = write!(stratum, "{}:{}", value.len(), value);
140 }
141
142 let taken = self.stratum_samples.entry(stratum).or_insert(0);
143 if *taken < samples_per_stratum {
144 *taken += 1;
145 true
146 } else {
147 false
148 }
149 }
150
151 pub fn strata_seen(&self) -> usize {
153 self.stratum_samples.len()
154 }
155}
156
157impl SamplingStrategy {
158 pub fn adaptive(total_rows: Option<usize>, file_size_mb: f64) -> Self {
160 match (total_rows, file_size_mb) {
161 (Some(rows), size_mb) if rows <= 10_000 && size_mb < 10.0 => SamplingStrategy::None,
162 (Some(rows), _) if rows <= 100_000 => SamplingStrategy::Random { size: 10_000 },
163 (Some(rows), _) if rows <= 1_000_000 => SamplingStrategy::Progressive {
164 initial_size: 10_000,
165 confidence_level: 0.95,
166 max_size: 50_000,
167 },
168 (_, size_mb) if size_mb > 1000.0 => SamplingStrategy::MultiStage {
169 stages: vec![
170 SamplingStrategy::Systematic { interval: 100 },
171 SamplingStrategy::Progressive {
172 initial_size: 5_000,
173 confidence_level: 0.99,
174 max_size: 25_000,
175 },
176 ],
177 },
178 _ => SamplingStrategy::Reservoir { size: 100_000 },
179 }
180 }
181
182 pub fn stratified(key_columns: Vec<String>, samples_per_stratum: usize) -> Self {
184 Self::Stratified {
185 key_columns,
186 samples_per_stratum,
187 }
188 }
189
190 pub fn importance(weight_column: impl Into<String>, weight_threshold: f64) -> Self {
192 Self::Importance {
193 weight_column: weight_column.into(),
194 weight_threshold,
195 }
196 }
197
198 pub fn target_sample_size(&self) -> Option<usize> {
199 match self {
200 SamplingStrategy::None => None,
201 SamplingStrategy::Random { size } => Some(*size),
202 SamplingStrategy::Reservoir { size } => Some(*size),
203 SamplingStrategy::Stratified {
204 samples_per_stratum,
205 ..
206 } => Some(*samples_per_stratum),
207 SamplingStrategy::Progressive { max_size, .. } => Some(*max_size),
208 SamplingStrategy::Systematic { .. } => None,
209 SamplingStrategy::Importance { .. } => None,
210 SamplingStrategy::MultiStage { stages } => {
211 stages.iter().filter_map(|s| s.target_sample_size()).min()
213 }
214 }
215 }
216
217 pub fn description(&self) -> String {
219 match self {
220 SamplingStrategy::None => "Full dataset analysis".to_string(),
221 SamplingStrategy::Random { size } => format!("Random sampling ({} records)", size),
222 SamplingStrategy::Reservoir { size } => {
223 format!("Reservoir sampling ({} records)", size)
224 }
225 SamplingStrategy::Stratified {
226 key_columns,
227 samples_per_stratum,
228 } => {
229 format!(
230 "Stratified by {} ({} per stratum)",
231 key_columns.join(", "),
232 samples_per_stratum
233 )
234 }
235 SamplingStrategy::Progressive {
236 initial_size,
237 confidence_level,
238 max_size,
239 } => {
240 format!(
241 "Progressive sampling ({}-{} records, {}% confidence)",
242 initial_size,
243 max_size,
244 (confidence_level * 100.0) as u8
245 )
246 }
247 SamplingStrategy::Systematic { interval } => {
248 format!("Systematic (every {}th record)", interval)
249 }
250 SamplingStrategy::Importance {
251 weight_column,
252 weight_threshold,
253 } => {
254 format!("Importance filter ({weight_column} >= {weight_threshold:.2})")
255 }
256 SamplingStrategy::MultiStage { stages } => {
257 format!("Multi-stage ({} stages)", stages.len())
258 }
259 }
260 }
261}
262
263#[cfg(test)]
264mod tests {
265 use super::*;
266 use crate::sampling::sampler::RowView;
267
268 #[test]
272 fn test_stratum_state_persists_across_rows() {
273 let headers = vec!["region".to_string()];
274 let mut state = SamplingState::new();
275 let keys = vec!["region".to_string()];
276
277 let north = vec!["north".to_string()];
278 let south = vec!["south".to_string()];
279
280 assert!(state.take_from_stratum(RowView::new(&headers, &north), &keys, 2));
281 assert!(state.take_from_stratum(RowView::new(&headers, &north), &keys, 2));
282 assert!(
283 !state.take_from_stratum(RowView::new(&headers, &north), &keys, 2),
284 "the third northern row exceeds the per-stratum cap"
285 );
286 assert!(
287 state.take_from_stratum(RowView::new(&headers, &south), &keys, 2),
288 "a different stratum has its own budget"
289 );
290 assert_eq!(state.strata_seen(), 2);
291 }
292
293 #[test]
294 fn test_stratum_requires_every_key_column() {
295 let headers = vec!["region".to_string()];
296 let values = vec!["north".to_string()];
297 let mut state = SamplingState::new();
298 let keys = vec!["region".to_string(), "segment".to_string()];
299
300 assert!(
301 !state.take_from_stratum(RowView::new(&headers, &values), &keys, 5),
302 "a row missing a key column belongs to no stratum"
303 );
304 assert_eq!(state.strata_seen(), 0);
305 }
306
307 #[test]
308 fn test_importance_constructor_names_its_column() {
309 let strategy = SamplingStrategy::importance("risk", 0.8);
310 match strategy {
311 SamplingStrategy::Importance {
312 ref weight_column,
313 weight_threshold,
314 } => {
315 assert_eq!(weight_column, "risk");
316 assert_eq!(weight_threshold, 0.8);
317 }
318 other => panic!("expected an importance filter, got {other:?}"),
319 }
320 assert!(strategy.description().contains("risk"));
321 }
322
323 #[test]
324 fn test_adaptive_strategy() {
325 let small = SamplingStrategy::adaptive(Some(5_000), 1.0);
327 matches!(small, SamplingStrategy::None);
328
329 let medium = SamplingStrategy::adaptive(Some(50_000), 10.0);
331 matches!(medium, SamplingStrategy::Random { .. });
332
333 let large = SamplingStrategy::adaptive(Some(10_000_000), 2000.0);
335 matches!(large, SamplingStrategy::MultiStage { .. });
336 }
337}