1use crate::data_quality::{
7 QualityScope, QualitySourceContext, apply_quality_scope, prepare_source_quality_scan,
8};
9use crate::numfmt;
10use crate::statistics::{AnalysisRows, collect_lazy, sample_rank};
11use color_eyre::Result;
12use color_eyre::eyre::Report;
13use polars::prelude::*;
14use std::collections::HashMap;
15
16pub const DEFAULT_SAMPLE_ROWS: usize = 100_000;
18
19#[derive(Debug, Clone, Copy, PartialEq, Eq)]
21pub enum SizeError {
22 NotASize,
23 Zero,
24}
25
26impl SizeError {
27 pub fn short(self) -> &'static str {
29 match self {
30 Self::NotASize => "not a size (50k, 2m)",
31 Self::Zero => "at least 1 row",
32 }
33 }
34}
35
36impl std::fmt::Display for SizeError {
37 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
38 f.write_str(match self {
39 Self::NotASize => "Sample size is a number of rows, like 50000, 50k or 2m",
40 Self::Zero => "Sample size is at least 1 row",
41 })
42 }
43}
44
45pub fn parse_size(text: &str) -> Result<usize, SizeError> {
48 let refuse = || SizeError::NotASize;
49 let cleaned: String = text
50 .trim()
51 .chars()
52 .filter(|c| !matches!(c, ',' | '_'))
53 .collect::<String>()
54 .to_ascii_lowercase();
55 let (number, scale) = match cleaned.strip_suffix('k') {
56 Some(n) => (n, 1e3),
57 None => match cleaned.strip_suffix('m') {
58 Some(n) => (n, 1e6),
59 None => (cleaned.as_str(), 1.0),
60 },
61 };
62 if number.is_empty() || !number.chars().all(|c| c.is_ascii_digit() || c == '.') {
63 return Err(refuse());
64 }
65 let rows = if scale == 1.0 {
66 if number.contains('.') {
68 return Err(refuse());
69 }
70 number.parse::<usize>().unwrap_or(usize::MAX)
71 } else {
72 let value: f64 = number.parse().map_err(|_| refuse())?;
73 (value * scale).round() as usize
74 };
75 if rows == 0 {
76 return Err(SizeError::Zero);
77 }
78 Ok(rows)
79}
80
81const MAX_GROUPS: usize = 10_000;
85const MAX_GROUP_ROWS: usize = 2_000_000;
86
87const GROUP_POSITION: &str = "__datui_group_sample_position";
89
90pub(crate) const COUNT_KEY: &str = "__datui_count_key";
93
94pub const MAX_COUNTED_KEYS: usize = 1_000_000;
97
98#[derive(Debug, Clone, PartialEq)]
100pub enum Counted {
101 Totals(std::collections::BTreeMap<Option<String>, usize>),
104 TooMany,
106}
107
108#[derive(Debug)]
110pub(crate) struct KeyCounter {
111 totals: HashMap<Option<String>, usize>,
112 limit: usize,
113 too_many: bool,
114}
115
116impl Default for KeyCounter {
117 fn default() -> Self {
118 Self::with_limit(MAX_COUNTED_KEYS)
119 }
120}
121
122impl KeyCounter {
123 fn with_limit(limit: usize) -> Self {
124 Self {
125 totals: HashMap::new(),
126 limit,
127 too_many: false,
128 }
129 }
130
131 pub(crate) fn observe(&mut self, batch: &mut DataFrame) -> PolarsResult<()> {
134 if batch.column(COUNT_KEY).is_err() {
135 return Ok(());
136 }
137 let key = batch.drop_in_place(COUNT_KEY)?;
138 if self.too_many {
139 return Ok(());
140 }
141 let counts = key.as_materialized_series().value_counts(
144 false,
145 false,
146 "__datui_count_rows".into(),
147 false,
148 )?;
149 let keys = counts.column(COUNT_KEY)?;
150 let rows = counts.column("__datui_count_rows")?;
151 for row in 0..counts.height() {
152 let value = keys.get(row)?;
153 let key = (!value.is_null()).then(|| crate::exact::str_value(&value).into_owned());
154 let n = rows.get(row)?.extract::<usize>().unwrap_or(0);
155 *self.totals.entry(key).or_default() += n;
156 }
157 if self.totals.len() > self.limit {
158 self.too_many = true;
159 self.totals = HashMap::new();
160 }
161 Ok(())
162 }
163
164 pub(crate) fn finish(self) -> Counted {
165 if self.too_many {
166 Counted::TooMany
167 } else {
168 Counted::Totals(self.totals.into_iter().collect())
169 }
170 }
171}
172
173pub(crate) fn with_count_key(lf: LazyFrame, count: Option<&Expr>) -> LazyFrame {
175 match count {
176 Some(key) => lf.with_column(key.clone().alias(COUNT_KEY)),
177 None => lf,
178 }
179}
180
181pub const CANCELLED: &str = "Cancelled";
183
184#[derive(Debug, Clone, Default)]
190pub struct ReadWatch {
191 stop: std::sync::Arc<std::sync::atomic::AtomicBool>,
192 rows: std::sync::Arc<std::sync::atomic::AtomicUsize>,
193 counted: std::sync::Arc<std::sync::atomic::AtomicBool>,
196 held: Option<HeldCheck>,
199 memory: std::sync::Arc<std::sync::Mutex<Option<String>>>,
201}
202
203pub type HeldJudge = dyn Fn(u64, usize) -> Option<String> + Send + Sync;
205
206#[derive(Clone)]
207struct HeldCheck(std::sync::Arc<HeldJudge>);
208
209impl std::fmt::Debug for HeldCheck {
210 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
211 f.write_str("HeldCheck")
212 }
213}
214
215impl ReadWatch {
216 pub(crate) fn judging_held(judge: std::sync::Arc<HeldJudge>) -> Self {
219 Self {
220 held: Some(HeldCheck(judge)),
221 ..Self::default()
222 }
223 }
224
225 pub(crate) fn hold(&self, bytes: u64, rows: usize) {
227 let Some(HeldCheck(judge)) = &self.held else {
228 return;
229 };
230 if let Some(reason) = judge(bytes, rows) {
231 *self.memory.lock().unwrap_or_else(|e| e.into_inner()) = Some(reason);
232 self.stop();
233 }
234 }
235
236 pub(crate) fn memory_stopped(&self) -> Option<String> {
238 self.memory
239 .lock()
240 .unwrap_or_else(|e| e.into_inner())
241 .clone()
242 }
243
244 pub fn stop(&self) {
245 self.stop.store(true, std::sync::atomic::Ordering::Relaxed);
246 }
247
248 pub fn stopped(&self) -> bool {
249 self.stop.load(std::sync::atomic::Ordering::Relaxed)
250 }
251
252 pub fn rows_seen(&self) -> Option<usize> {
254 self.counted
255 .load(std::sync::atomic::Ordering::Relaxed)
256 .then(|| self.rows.load(std::sync::atomic::Ordering::Relaxed))
257 }
258
259 pub(crate) fn saw(&self, rows: usize) {
260 self.counted
261 .store(true, std::sync::atomic::Ordering::Relaxed);
262 self.rows
263 .fetch_add(rows, std::sync::atomic::Ordering::Relaxed);
264 }
265
266 pub(crate) fn restart(&self) -> Option<usize> {
269 let counted = self
270 .counted
271 .swap(false, std::sync::atomic::Ordering::Relaxed);
272 let rows = self.rows.swap(0, std::sync::atomic::Ordering::Relaxed);
273 counted.then_some(rows)
274 }
275
276 pub(crate) fn check(&self) -> Result<()> {
279 if self.stopped() && self.memory_stopped().is_none() {
280 Err(Report::msg(CANCELLED))
281 } else {
282 Ok(())
283 }
284 }
285}
286
287#[derive(Debug, Clone, PartialEq, Eq, Default)]
289pub enum SampleMethod {
290 #[default]
292 Spread,
293 PerPartition { column: String },
296 FirstRows,
298 EveryRow,
300}
301
302impl SampleMethod {
303 pub fn label(&self) -> String {
304 match self {
305 Self::Spread => "Random".to_string(),
306 Self::PerPartition { column } => format!("Equal per {column}"),
307 Self::FirstRows => "First rows".to_string(),
308 Self::EveryRow => "Every row".to_string(),
309 }
310 }
311}
312
313#[derive(Debug, Clone, PartialEq, Eq)]
315pub struct Sample {
316 pub scope: QualityScope,
317 pub method: SampleMethod,
318 pub rows: usize,
321 pub seed: u64,
322}
323
324impl Default for Sample {
325 fn default() -> Self {
326 Self {
327 scope: QualityScope::CurrentView,
328 method: SampleMethod::Spread,
329 rows: DEFAULT_SAMPLE_ROWS,
330 seed: 42_891,
331 }
332 }
333}
334
335impl Sample {
336 pub fn summary(&self) -> String {
338 let rows = numfmt::group_chrome(self.rows);
339 let how = match &self.method {
340 SampleMethod::Spread => format!("{rows} random rows"),
341 SampleMethod::PerPartition { column } => format!("{rows} rows per {column}"),
342 SampleMethod::FirstRows => format!("first {rows} rows"),
343 SampleMethod::EveryRow => "every row".to_string(),
344 };
345 let seeded = matches!(
346 self.method,
347 SampleMethod::Spread | SampleMethod::PerPartition { .. }
348 );
349 let middot = crate::glyphs::get().middot;
350 if seeded {
351 format!(
352 "{how} {middot} {} {middot} seed {}",
353 self.scope.label(),
354 self.seed
355 )
356 } else {
357 format!("{how} {middot} {}", self.scope.label())
358 }
359 }
360
361 pub fn summary_within(&self, rows: Option<usize>) -> String {
365 match (rows, &self.method) {
366 (Some(n), SampleMethod::Spread | SampleMethod::FirstRows) if n <= self.rows => {
367 let middot = crate::glyphs::get().middot;
368 format!(
369 "all {} rows {middot} {}",
370 numfmt::group_chrome(n),
371 self.scope.label()
372 )
373 }
374 _ => self.summary(),
375 }
376 }
377
378 pub fn outcome(
382 &self,
383 total_rows: usize,
384 sample_size: Option<usize>,
385 per_value: Option<usize>,
386 ) -> String {
387 let count = numfmt::group_chrome;
388 let read = match (&self.method, sample_size) {
389 (SampleMethod::FirstRows, Some(n)) => format!("first {} rows", count(n)),
390 (SampleMethod::PerPartition { column }, Some(n)) => {
391 let each = per_value.unwrap_or(self.rows).min(self.rows);
392 let lowered = if each < self.rows {
393 format!(" (lowered from {})", count(self.rows))
394 } else {
395 String::new()
396 };
397 format!(
398 "{} rows, up to {} per {column}{lowered}, of {}",
399 count(n),
400 count(each),
401 count(total_rows)
402 )
403 }
404 (_, Some(n)) => format!("sample of {} of {} rows", count(n), count(total_rows)),
405 (_, None) => format!("all {} rows", count(total_rows)),
406 };
407 if self.scope == QualityScope::CurrentView {
408 read
409 } else {
410 format!(
411 "{read} {} {}",
412 crate::glyphs::get().middot,
413 self.scope.label()
414 )
415 }
416 }
417}
418
419pub struct SampleSource {
423 lf: LazyFrame,
424 source: Option<QualitySourceContext>,
425 from_source: bool,
426}
427
428impl SampleSource {
429 pub fn view(lf: LazyFrame) -> Self {
431 Self {
432 lf,
433 source: None,
434 from_source: false,
435 }
436 }
437
438 pub fn loaded(lf: LazyFrame, source: Option<QualitySourceContext>) -> Self {
440 Self {
441 lf,
442 source,
443 from_source: true,
444 }
445 }
446
447 pub fn cut(self, scope: &QualityScope) -> Result<LazyFrame> {
450 let lf = if self.from_source {
451 prepare_source_quality_scan(self.lf, self.source.as_ref())?
452 } else {
453 self.lf
454 };
455 let lf = apply_quality_scope(lf, scope, self.source.as_ref())?;
456 let schema = lf.clone().collect_schema()?;
457 let helpers = [
458 crate::schema_union::DRIFT_COLUMN,
459 "__datui_quality_row",
460 self.source
461 .as_ref()
462 .map(|source| source.row_index_column.as_str())
463 .unwrap_or(""),
464 ];
465 let keep = schema
466 .iter_names()
467 .filter(|name| !helpers.contains(&name.as_str()))
468 .map(|name| col(name.clone()))
469 .collect::<Vec<_>>();
470 Ok(if keep.len() == schema.len() {
471 lf
472 } else {
473 lf.select(keep)
474 })
475 }
476}
477
478pub fn view_scope_rows(view_rows: Option<usize>, scope: &QualityScope) -> Option<usize> {
481 let rows = view_rows?;
482 match scope {
483 QualityScope::CurrentView => Some(rows),
484 QualityScope::FirstRows(limit) => Some(rows.min(*limit)),
485 QualityScope::ViewRows { start, end } => {
486 Some(rows.min(*end).saturating_sub(start.saturating_sub(1)))
487 }
488 _ => None,
489 }
490}
491
492pub fn read(
498 lf: &LazyFrame,
499 sample: &Sample,
500 known_total: Option<usize>,
501 polars_streaming: bool,
502) -> Result<AnalysisRows> {
503 let rows = read_rows(lf, sample, known_total, polars_streaming)?;
504 if rows.total_rows == 0 {
505 return Err(no_rows_error(&sample.scope));
506 }
507 Ok(rows)
508}
509
510pub fn no_rows_error(scope: &QualityScope) -> Report {
514 if *scope == QualityScope::CurrentView {
515 Report::msg("The table has no rows to sample")
516 } else {
517 Report::msg(format!(
518 "No rows match {}; change Rows from in the Sample form (s)",
519 scope.label()
520 ))
521 }
522}
523
524pub(crate) fn read_rows(
525 lf: &LazyFrame,
526 sample: &Sample,
527 known_total: Option<usize>,
528 polars_streaming: bool,
529) -> Result<AnalysisRows> {
530 read_rows_watched(lf, sample, known_total, polars_streaming, None)
531}
532
533pub(crate) fn read_rows_watched(
535 lf: &LazyFrame,
536 sample: &Sample,
537 known_total: Option<usize>,
538 polars_streaming: bool,
539 watch: Option<&ReadWatch>,
540) -> Result<AnalysisRows> {
541 acquire(lf, sample, known_total, polars_streaming, watch, None).map(|read| read.rows)
542}
543
544pub(crate) struct SampledRows {
546 pub rows: AnalysisRows,
547 pub positions: Vec<IdxSize>,
550 pub counted: Option<Counted>,
553}
554
555pub(crate) fn acquire(
559 lf: &LazyFrame,
560 sample: &Sample,
561 known_total: Option<usize>,
562 polars_streaming: bool,
563 watch: Option<&ReadWatch>,
564 count: Option<&Expr>,
565) -> Result<SampledRows> {
566 let n = sample.rows.max(1);
567 match &sample.method {
568 SampleMethod::EveryRow => crate::statistics::sample_rows_counting(
569 lf,
570 None,
571 known_total,
572 sample.seed,
573 polars_streaming,
574 watch,
575 count,
576 ),
577 SampleMethod::Spread => crate::statistics::sample_rows_counting(
578 lf,
579 Some(n),
580 known_total,
581 sample.seed,
582 polars_streaming,
583 watch,
584 count,
585 ),
586 SampleMethod::FirstRows => {
587 let df = collect_lazy(lf.clone().limit(n as IdxSize), polars_streaming)
588 .map_err(Report::from)?;
589 let height = df.height();
590 let sampled = match known_total {
592 Some(total) => total > height,
593 None => height == n,
594 };
595 Ok(SampledRows {
596 positions: (0..height as IdxSize).collect(),
597 rows: AnalysisRows {
598 df,
599 total_rows: known_total.unwrap_or(height),
600 sample_size: sampled.then_some(height),
601 per_value: None,
602 },
603 counted: None,
604 })
605 }
606 SampleMethod::PerPartition { column } => {
607 let read =
608 per_group_sample_within(lf, column, n, sample.seed, MAX_GROUP_ROWS, watch, count)?;
609 let sample_size = (read.seen > read.df.height()).then_some(read.df.height());
610 Ok(SampledRows {
611 rows: AnalysisRows {
612 df: read.df,
613 total_rows: read.seen,
614 sample_size,
615 per_value: Some(read.per_value),
616 },
617 positions: read.positions,
618 counted: read.counted,
619 })
620 }
621 }
622}
623
624#[derive(Debug, Clone, Default, PartialEq)]
626pub struct PerValue {
627 pub kept: usize,
630 pub totals: std::collections::BTreeMap<Option<String>, usize>,
634}
635
636struct GroupRead {
638 df: DataFrame,
639 seen: usize,
640 per_value: PerValue,
641 positions: Vec<IdxSize>,
642 counted: Option<Counted>,
643}
644
645fn per_group_sample_within(
655 lf: &LazyFrame,
656 column: &str,
657 n: usize,
658 seed: u64,
659 limit: usize,
660 watch: Option<&ReadWatch>,
661 count: Option<&Expr>,
662) -> Result<GroupRead> {
663 let schema = lf.clone().collect_schema()?;
664 if schema.get(column).is_none() {
665 return Err(Report::msg(format!(
666 "partition column {column:?} is not in the rows sampled; choose another"
667 )));
668 }
669 let state = std::sync::Arc::new(std::sync::Mutex::new(GroupState {
670 column: column.to_string(),
671 cap: n,
672 limit,
673 seed,
674 ..Default::default()
675 }));
676 let callback_state = std::sync::Arc::clone(&state);
677 let callback_watch = watch.cloned();
678 let sink = with_count_key(lf.clone(), count)
679 .with_row_index(GROUP_POSITION, None)
680 .sink_batches(
681 PlanCallback::new(move |batch: DataFrame| {
682 if let Some(watch) = &callback_watch {
684 if watch.stopped() {
685 return Ok(true);
686 }
687 watch.saw(batch.height());
688 }
689 let mut state = callback_state
690 .lock()
691 .map_err(|_| PolarsError::ComputeError("sampler lock failed".into()))?;
692 state.observe(batch)?;
693 if let Some(watch) = &callback_watch {
694 watch.hold(state.bytes(), state.held);
695 }
696 Ok(false)
697 }),
698 true,
699 None,
700 )?;
701 collect_lazy(sink, true).map_err(Report::from)?;
703 if let Some(watch) = watch {
704 watch.check()?;
705 }
706 let state = std::mem::take(
707 &mut *state
708 .lock()
709 .map_err(|_| Report::msg("sampler lock failed"))?,
710 );
711 let seen = state.seen;
712 let cap = state
715 .cap
716 .min((state.limit / state.groups.len().max(1)).max(1));
717 let mut totals = std::collections::BTreeMap::new();
718 let mut out: Option<DataFrame> = None;
719 for (key, mut group) in state.groups {
720 totals.insert(key, group.total);
721 group.trim(cap)?;
722 let Some(rows) = group.rows else {
723 continue;
724 };
725 out = Some(match out {
726 Some(frame) => frame.vstack(&rows)?,
727 None => rows,
728 });
729 }
730 let (df, positions) = match out {
731 Some(df) => {
732 let df = df.sort([GROUP_POSITION], SortMultipleOptions::default())?;
733 let positions = df
734 .column(GROUP_POSITION)?
735 .idx()?
736 .into_no_null_iter()
737 .collect();
738 (df.drop(GROUP_POSITION)?, positions)
739 }
740 None => (
741 collect_lazy(lf.clone().limit(0), true).map_err(Report::from)?,
742 Vec::new(),
743 ),
744 };
745 Ok(GroupRead {
746 df,
747 seen,
748 per_value: PerValue { kept: cap, totals },
749 positions,
750 counted: count.is_some().then(|| state.counter.finish()),
751 })
752}
753
754#[derive(Default)]
755struct GroupState {
756 column: String,
757 cap: usize,
760 limit: usize,
762 seed: u64,
763 seen: usize,
764 held: usize,
766 groups: HashMap<Option<String>, GroupSample>,
767 counter: KeyCounter,
768}
769
770#[derive(Default)]
771struct GroupSample {
772 rows: Option<DataFrame>,
773 ranks: Vec<u64>,
774 total: usize,
776}
777
778impl GroupSample {
779 fn trim(&mut self, cap: usize) -> PolarsResult<usize> {
782 if self.ranks.len() <= cap {
783 return Ok(0);
784 }
785 let mut order: Vec<usize> = (0..self.ranks.len()).collect();
786 order.sort_unstable_by_key(|i| self.ranks[*i]);
787 order.truncate(cap);
788 let take: Vec<IdxSize> = order.iter().map(|i| *i as IdxSize).collect();
789 if let Some(kept) = self.rows.take() {
790 self.rows = Some(kept.take(&IdxCa::from_vec("kept".into(), take))?);
791 }
792 let removed = self.ranks.len() - cap;
793 self.ranks = order.iter().map(|i| self.ranks[*i]).collect();
794 Ok(removed)
795 }
796}
797
798impl GroupState {
799 fn bytes(&self) -> u64 {
801 self.groups
802 .values()
803 .filter_map(|group| group.rows.as_ref())
804 .map(|rows| rows.estimated_size() as u64)
805 .sum()
806 }
807
808 fn observe(&mut self, mut batch: DataFrame) -> PolarsResult<()> {
809 self.counter.observe(&mut batch)?;
810 self.seen += batch.height();
811 let positions = batch.column(GROUP_POSITION)?.idx()?.clone();
812 let keys = batch.column(&self.column)?.as_materialized_series().clone();
813 let mut by_key: HashMap<Option<String>, (Vec<IdxSize>, Vec<u64>)> = HashMap::new();
814 for (index, (key, position)) in keys.iter().zip(positions.into_no_null_iter()).enumerate() {
815 let key = (!key.is_null()).then(|| crate::exact::str_value(&key).into_owned());
817 let entry = by_key.entry(key).or_default();
818 entry.0.push(index as IdxSize);
819 entry.1.push(sample_rank(self.seed, position as u64));
820 }
821 for (key, (indices, ranks)) in by_key {
822 if !self.groups.contains_key(&key) && self.groups.len() >= MAX_GROUPS {
823 return Err(PolarsError::ComputeError(
824 format!(
825 "more than {MAX_GROUPS} values of {}; sample per a coarser column",
826 self.column
827 )
828 .into(),
829 ));
830 }
831 let group = self.groups.entry(key).or_default();
832 group.total += indices.len();
833 self.held += indices.len();
834 let rows = batch.take(&IdxCa::from_vec("picked".into(), indices))?;
835 group.rows = Some(match group.rows.take() {
836 Some(kept) => kept.vstack(&rows)?,
837 None => rows,
838 });
839 group.ranks.extend(ranks);
840 self.held -= group.trim(self.cap)?;
841 }
842 if self.held > self.limit.saturating_add(self.limit / 4) {
846 self.cap = self.cap.min((self.limit / self.groups.len()).max(1));
847 for group in self.groups.values_mut() {
848 self.held -= group.trim(self.cap)?;
849 }
850 }
851 Ok(())
852 }
853}
854
855#[cfg(test)]
856mod tests {
857 #[test]
858 fn a_size_takes_shorthand() {
859 use super::parse_size;
860 for (text, rows) in [
861 ("50000", 50_000),
862 ("50,000", 50_000),
863 ("1_000", 1_000),
864 ("50k", 50_000),
865 ("250K", 250_000),
866 ("2m", 2_000_000),
867 ("2.5M", 2_500_000),
868 (" 7 ", 7),
869 ("99999999999999999999999", usize::MAX),
870 ] {
871 assert_eq!(parse_size(text), Ok(rows), "{text}");
872 }
873 for bad in ["", "k", "12x", "1.5", "-3", "1e6", "2mm"] {
874 assert_eq!(parse_size(bad), Err(super::SizeError::NotASize), "{bad}");
875 }
876 for zero in ["0", "0k", "0.0001k"] {
877 assert_eq!(parse_size(zero), Err(super::SizeError::Zero), "{zero}");
878 }
879 }
880
881 use super::*;
882
883 fn table() -> LazyFrame {
884 let sizes = [("a", 9_000usize), ("b", 900), ("c", 100)];
886 let mut part = Vec::new();
887 let mut value = Vec::new();
888 for (name, size) in sizes {
889 for row in 0..size {
890 part.push(name);
891 value.push(row as i64);
892 }
893 }
894 df!("part" => part, "value" => value).unwrap().lazy()
895 }
896
897 fn sample(method: SampleMethod, rows: usize) -> Sample {
898 Sample {
899 method,
900 rows,
901 ..Sample::default()
902 }
903 }
904
905 #[test]
908 fn per_partition_keeps_up_to_n_from_each_value() {
909 let method = SampleMethod::PerPartition {
910 column: "part".to_string(),
911 };
912 let rows = read(&table(), &sample(method, 200), None, false).unwrap();
913 assert_eq!(rows.total_rows, 10_000);
914 let counts = rows
915 .df
916 .column("part")
917 .unwrap()
918 .as_materialized_series()
919 .value_counts(true, true, "n".into(), false)
920 .unwrap();
921 let n = |part: &str| {
922 (0..counts.height())
923 .find(|row| {
924 counts
925 .column("part")
926 .unwrap()
927 .get(*row)
928 .unwrap()
929 .str_value()
930 == part
931 })
932 .map(|row| {
933 counts
934 .column("n")
935 .unwrap()
936 .get(row)
937 .unwrap()
938 .try_extract::<u32>()
939 .unwrap()
940 })
941 .unwrap()
942 };
943 assert_eq!((n("a"), n("b"), n("c")), (200, 200, 100));
944 assert_eq!(rows.sample_size, Some(500));
945 assert!(
946 rows.df.column(GROUP_POSITION).is_err(),
947 "no helper column leaks"
948 );
949 }
950
951 #[test]
955 fn per_partition_past_the_limit_keeps_fewer_of_each_value() {
956 let GroupRead {
957 df,
958 seen,
959 per_value,
960 ..
961 } = per_group_sample_within(&table(), "part", 500, 42_891, 999, None, None).unwrap();
962 assert_eq!(seen, 10_000);
963 assert_eq!(per_value.kept, 333);
964 assert_eq!(df.height(), 333 + 333 + 100);
965 assert_eq!(
966 per_value.totals,
967 [("a", 9_000), ("b", 900), ("c", 100)]
968 .into_iter()
969 .map(|(part, rows)| (Some(part.to_string()), rows))
970 .collect()
971 );
972 let asked =
973 per_group_sample_within(&table(), "part", 333, 42_891, usize::MAX, None, None).unwrap();
974 assert!(df.equals(&asked.df), "the rows a sample of 333 each keeps");
975
976 let lowered = Sample {
977 method: SampleMethod::PerPartition {
978 column: "part".into(),
979 },
980 rows: 500,
981 ..Sample::default()
982 };
983 assert_eq!(
984 lowered.outcome(10_000, Some(766), Some(333)),
985 "766 rows, up to 333 per part (lowered from 500), of 10,000"
986 );
987 }
988
989 #[test]
994 fn a_sample_says_where_its_rows_sat_and_a_stream_counts_on_the_way() {
995 let lf = table()
997 .with_column((col("value") % lit(4)).alias("quarter"))
998 .with_row_index("row", None);
999 let totals: std::collections::BTreeMap<_, _> = [("a", 9_000), ("b", 900), ("c", 100)]
1000 .into_iter()
1001 .map(|(part, rows)| (Some(part.to_string()), rows))
1002 .collect();
1003 for method in [
1004 SampleMethod::Spread,
1005 SampleMethod::PerPartition {
1006 column: "quarter".to_string(),
1007 },
1008 SampleMethod::FirstRows,
1009 ] {
1010 let read = acquire(
1011 &lf,
1012 &sample(method.clone(), 50),
1013 None,
1014 false,
1015 None,
1016 Some(&col("part")),
1017 )
1018 .unwrap();
1019 let rows: Vec<IdxSize> = read
1020 .rows
1021 .df
1022 .column("row")
1023 .unwrap()
1024 .idx()
1025 .unwrap()
1026 .into_no_null_iter()
1027 .collect();
1028 assert_eq!(rows, read.positions, "{method:?}");
1029 assert!(read.rows.df.column(COUNT_KEY).is_err(), "{method:?}");
1030 if method == SampleMethod::FirstRows {
1031 assert_eq!(read.counted, None);
1032 } else {
1033 assert_eq!(
1034 read.counted,
1035 Some(Counted::Totals(totals.clone())),
1036 "{method:?}"
1037 );
1038 }
1039 }
1040 }
1041
1042 #[test]
1045 fn a_count_past_its_limit_gives_up_and_says_so() {
1046 let mut counter = KeyCounter::with_limit(2);
1047 let mut batch = df!(COUNT_KEY => ["a", "b", "c"], "value" => [1, 2, 3]).unwrap();
1048 counter.observe(&mut batch).unwrap();
1049 assert_eq!(batch.get_column_names(), ["value"]);
1050 assert_eq!(counter.finish(), Counted::TooMany);
1051 }
1052
1053 #[test]
1054 fn first_rows_is_the_head_and_every_row_is_all_of_it() {
1055 let head = read(
1056 &table(),
1057 &sample(SampleMethod::FirstRows, 50),
1058 Some(10_000),
1059 false,
1060 )
1061 .unwrap();
1062 assert_eq!(head.df.height(), 50);
1063 assert_eq!(head.sample_size, Some(50));
1064 assert_eq!(
1065 head.df.column("value").unwrap().i64().unwrap().get(49),
1066 Some(49)
1067 );
1068 let all = read(&table(), &sample(SampleMethod::EveryRow, 50), None, false).unwrap();
1069 assert_eq!((all.df.height(), all.sample_size), (10_000, None));
1070 }
1071
1072 #[test]
1073 fn a_seeded_per_partition_sample_repeats() {
1074 let method = SampleMethod::PerPartition {
1075 column: "part".to_string(),
1076 };
1077 let one = read(&table(), &sample(method.clone(), 50), None, false).unwrap();
1078 let two = read(&table(), &sample(method.clone(), 50), None, false).unwrap();
1079 assert!(one.df.equals(&two.df));
1080 let other = read(
1081 &table(),
1082 &Sample {
1083 seed: 7,
1084 ..sample(method, 50)
1085 },
1086 None,
1087 false,
1088 )
1089 .unwrap();
1090 assert!(!one.df.equals(&other.df));
1091 }
1092
1093 #[test]
1096 fn a_scope_that_matches_nothing_is_an_error() {
1097 let scope = QualityScope::parse_command("partition part=zzz").unwrap();
1098 let lf = SampleSource::view(table()).cut(&scope).unwrap();
1099 let Err(error) = read(
1100 &lf,
1101 &Sample {
1102 scope,
1103 ..Sample::default()
1104 },
1105 None,
1106 false,
1107 ) else {
1108 panic!("a scope that matches nothing must not sample");
1109 };
1110 assert!(error.to_string().contains("No rows match"), "{error}");
1111 }
1112
1113 #[test]
1114 fn the_summary_says_what_will_be_read() {
1115 let middot = crate::glyphs::get().middot;
1116 assert_eq!(
1117 Sample::default().summary(),
1118 format!("100,000 random rows {middot} current view {middot} seed 42891")
1119 );
1120 assert_eq!(
1121 sample(SampleMethod::FirstRows, 1_000).summary(),
1122 format!("first 1,000 rows {middot} current view")
1123 );
1124 assert_eq!(
1126 Sample::default().summary_within(Some(1_000)),
1127 format!("all 1,000 rows {middot} current view")
1128 );
1129 assert_eq!(
1130 Sample::default().summary_within(Some(1_000_000)),
1131 Sample::default().summary()
1132 );
1133 assert_eq!(
1134 Sample::default().summary_within(None),
1135 Sample::default().summary()
1136 );
1137 }
1138}