1use crate::analysis::data_quality::{
7 QualityScope, QualitySourceContext, apply_quality_scope, prepare_source_quality_scan,
8};
9use crate::analysis::statistics::collect_lazy;
10use crate::numfmt;
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;
84const MAX_GROUP_ROWS: usize = 2_000_000;
85
86const GROUP_POSITION: &str = "__datui_group_sample_position";
88
89pub(crate) const COUNT_KEY: &str = "__datui_count_key";
92
93pub const MAX_COUNTED_KEYS: usize = 1_000_000;
96
97#[derive(Debug, Clone, PartialEq)]
99pub enum Counted {
100 Totals(std::collections::BTreeMap<Option<String>, usize>),
103 TooMany,
105}
106
107#[derive(Debug)]
109pub(crate) struct KeyCounter {
110 totals: HashMap<Option<String>, usize>,
111 limit: usize,
112 too_many: bool,
113}
114
115impl Default for KeyCounter {
116 fn default() -> Self {
117 Self::with_limit(MAX_COUNTED_KEYS)
118 }
119}
120
121impl KeyCounter {
122 fn with_limit(limit: usize) -> Self {
123 Self {
124 totals: HashMap::new(),
125 limit,
126 too_many: false,
127 }
128 }
129
130 pub(crate) fn observe(&mut self, batch: &mut DataFrame) -> PolarsResult<()> {
133 if batch.column(COUNT_KEY).is_err() {
134 return Ok(());
135 }
136 let key = batch.drop_in_place(COUNT_KEY)?;
137 if self.too_many {
138 return Ok(());
139 }
140 let counts = key.as_materialized_series().value_counts(
143 false,
144 false,
145 "__datui_count_rows".into(),
146 false,
147 )?;
148 let keys = counts.column(COUNT_KEY)?;
149 let rows = counts.column("__datui_count_rows")?;
150 for row in 0..counts.height() {
151 let value = keys.get(row)?;
152 let key = (!value.is_null()).then(|| crate::exact::str_value(&value).into_owned());
153 let n = rows.get(row)?.extract::<usize>().unwrap_or(0);
154 *self.totals.entry(key).or_default() += n;
155 }
156 if self.totals.len() > self.limit {
157 self.too_many = true;
158 self.totals = HashMap::new();
159 }
160 Ok(())
161 }
162
163 pub(crate) fn finish(self) -> Counted {
164 if self.too_many {
165 Counted::TooMany
166 } else {
167 Counted::Totals(self.totals.into_iter().collect())
168 }
169 }
170}
171
172pub(crate) fn with_count_key(lf: LazyFrame, count: Option<&Expr>) -> LazyFrame {
174 match count {
175 Some(key) => lf.with_column(key.clone().alias(COUNT_KEY)),
176 None => lf,
177 }
178}
179
180pub const CANCELLED: &str = "Cancelled";
182
183#[derive(Debug, Clone, Default)]
187pub struct ReadWatch {
188 stop: std::sync::Arc<std::sync::atomic::AtomicBool>,
189 rows: std::sync::Arc<std::sync::atomic::AtomicUsize>,
190 counted: std::sync::Arc<std::sync::atomic::AtomicBool>,
193 held: Option<HeldCheck>,
196 memory: std::sync::Arc<std::sync::Mutex<Option<String>>>,
198}
199
200pub type HeldJudge = dyn Fn(u64, usize) -> Option<String> + Send + Sync;
202
203#[derive(Clone)]
204struct HeldCheck(std::sync::Arc<HeldJudge>);
205
206impl std::fmt::Debug for HeldCheck {
207 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
208 f.write_str("HeldCheck")
209 }
210}
211
212impl ReadWatch {
213 pub(crate) fn judging_held(judge: std::sync::Arc<HeldJudge>) -> Self {
216 Self {
217 held: Some(HeldCheck(judge)),
218 ..Self::default()
219 }
220 }
221
222 pub(crate) fn hold(&self, bytes: u64, rows: usize) {
224 let Some(HeldCheck(judge)) = &self.held else {
225 return;
226 };
227 if let Some(reason) = judge(bytes, rows) {
228 *self.memory.lock().unwrap_or_else(|e| e.into_inner()) = Some(reason);
229 self.stop();
230 }
231 }
232
233 pub(crate) fn memory_stopped(&self) -> Option<String> {
235 self.memory
236 .lock()
237 .unwrap_or_else(|e| e.into_inner())
238 .clone()
239 }
240
241 pub fn stop(&self) {
242 self.stop.store(true, std::sync::atomic::Ordering::Relaxed);
243 }
244
245 pub fn stopped(&self) -> bool {
246 self.stop.load(std::sync::atomic::Ordering::Relaxed)
247 }
248
249 pub fn rows_seen(&self) -> Option<usize> {
251 self.counted
252 .load(std::sync::atomic::Ordering::Relaxed)
253 .then(|| self.rows.load(std::sync::atomic::Ordering::Relaxed))
254 }
255
256 pub(crate) fn saw(&self, rows: usize) {
257 self.counted
258 .store(true, std::sync::atomic::Ordering::Relaxed);
259 self.rows
260 .fetch_add(rows, std::sync::atomic::Ordering::Relaxed);
261 }
262
263 pub(crate) fn restart(&self) -> Option<usize> {
266 let counted = self
267 .counted
268 .swap(false, std::sync::atomic::Ordering::Relaxed);
269 let rows = self.rows.swap(0, std::sync::atomic::Ordering::Relaxed);
270 counted.then_some(rows)
271 }
272
273 pub(crate) fn check(&self) -> Result<()> {
276 if self.stopped() && self.memory_stopped().is_none() {
277 Err(Report::msg(CANCELLED))
278 } else {
279 Ok(())
280 }
281 }
282}
283
284#[derive(Debug, Clone, PartialEq, Eq, Default)]
286pub enum SampleMethod {
287 #[default]
289 Spread,
290 PerPartition { column: String },
293 FirstRows,
295 EveryRow,
297}
298
299impl SampleMethod {
300 pub fn name(&self) -> &'static str {
302 match self {
303 Self::Spread => "Random",
304 Self::PerPartition { .. } => "Equal per value",
305 Self::FirstRows => "First rows",
306 Self::EveryRow => "Every row",
307 }
308 }
309
310 pub fn label(&self) -> String {
312 match self {
313 Self::PerPartition { column } => format!("Equal per {column}"),
314 method => method.name().to_string(),
315 }
316 }
317}
318
319#[derive(Debug, Clone, PartialEq, Eq)]
321pub struct Sample {
322 pub scope: QualityScope,
323 pub method: SampleMethod,
324 pub rows: usize,
327 pub seed: u64,
328}
329
330impl Default for Sample {
331 fn default() -> Self {
332 Self {
333 scope: QualityScope::CurrentView,
334 method: SampleMethod::Spread,
335 rows: DEFAULT_SAMPLE_ROWS,
336 seed: 42_891,
337 }
338 }
339}
340
341impl Sample {
342 pub fn summary(&self) -> String {
344 let rows = numfmt::group_chrome(self.rows);
345 let how = match &self.method {
346 SampleMethod::Spread => format!("{rows} random rows"),
347 SampleMethod::PerPartition { column } => format!("{rows} rows per {column}"),
348 SampleMethod::FirstRows => format!("first {rows} rows"),
349 SampleMethod::EveryRow => "every row".to_string(),
350 };
351 let seeded = matches!(
352 self.method,
353 SampleMethod::Spread | SampleMethod::PerPartition { .. }
354 );
355 let middot = crate::glyphs::get().middot;
356 if seeded {
357 format!(
358 "{how} {middot} {} {middot} seed {}",
359 self.scope.label(),
360 self.seed
361 )
362 } else {
363 format!("{how} {middot} {}", self.scope.label())
364 }
365 }
366
367 pub fn summary_within(&self, rows: Option<usize>) -> String {
370 match (rows, &self.method) {
371 (Some(n), SampleMethod::Spread | SampleMethod::FirstRows) if n <= self.rows => {
372 let middot = crate::glyphs::get().middot;
373 format!(
374 "all {} rows {middot} {}",
375 numfmt::group_chrome(n),
376 self.scope.label()
377 )
378 }
379 _ => self.summary(),
380 }
381 }
382
383 pub fn outcome(
387 &self,
388 total_rows: usize,
389 sample_size: Option<usize>,
390 per_value: Option<usize>,
391 ) -> String {
392 let count = numfmt::group_chrome;
393 let read = match (&self.method, sample_size) {
394 (SampleMethod::FirstRows, Some(n)) => format!("first {} rows", count(n)),
395 (SampleMethod::PerPartition { column }, Some(n)) => {
396 let each = per_value.unwrap_or(self.rows).min(self.rows);
397 let lowered = if each < self.rows {
398 format!(" (lowered from {})", count(self.rows))
399 } else {
400 String::new()
401 };
402 format!(
403 "{} rows, up to {} per {column}{lowered}, of {}",
404 count(n),
405 count(each),
406 count(total_rows)
407 )
408 }
409 (_, Some(n)) => format!("sample of {} of {} rows", count(n), count(total_rows)),
410 (_, None) => format!("all {} rows", count(total_rows)),
411 };
412 if self.scope == QualityScope::CurrentView {
413 read
414 } else {
415 format!(
416 "{read} {} {}",
417 crate::glyphs::get().middot,
418 self.scope.label()
419 )
420 }
421 }
422}
423
424pub struct SampleSource {
428 lf: LazyFrame,
429 source: Option<QualitySourceContext>,
430 from_source: bool,
431}
432
433impl SampleSource {
434 pub fn view(lf: LazyFrame) -> Self {
436 Self {
437 lf,
438 source: None,
439 from_source: false,
440 }
441 }
442
443 pub fn loaded(lf: LazyFrame, source: Option<QualitySourceContext>) -> Self {
445 Self {
446 lf,
447 source,
448 from_source: true,
449 }
450 }
451
452 pub fn cut(self, scope: &QualityScope) -> Result<LazyFrame> {
455 let lf = if self.from_source {
456 prepare_source_quality_scan(self.lf, self.source.as_ref())?
457 } else {
458 self.lf
459 };
460 let lf = apply_quality_scope(lf, scope, self.source.as_ref())?;
461 let schema = lf.clone().collect_schema()?;
462 let helpers = [
463 crate::formats::schema_union::DRIFT_COLUMN,
464 "__datui_quality_row",
465 self.source
466 .as_ref()
467 .map(|source| source.row_index_column.as_str())
468 .unwrap_or(""),
469 ];
470 let keep = schema
471 .iter_names()
472 .filter(|name| !helpers.contains(&name.as_str()))
473 .map(|name| col(name.clone()))
474 .collect::<Vec<_>>();
475 Ok(if keep.len() == schema.len() {
476 lf
477 } else {
478 lf.select(keep)
479 })
480 }
481}
482
483pub fn view_scope_rows(view_rows: Option<usize>, scope: &QualityScope) -> Option<usize> {
486 let rows = view_rows?;
487 match scope {
488 QualityScope::CurrentView => Some(rows),
489 QualityScope::FirstRows(limit) => Some(rows.min(*limit)),
490 QualityScope::ViewRows { start, end } => {
491 Some(rows.min(*end).saturating_sub(start.saturating_sub(1)))
492 }
493 _ => None,
494 }
495}
496
497pub fn read(
501 lf: &LazyFrame,
502 sample: &Sample,
503 known_total: Option<usize>,
504 polars_streaming: bool,
505) -> Result<AnalysisRows> {
506 let rows = read_rows(lf, sample, known_total, polars_streaming)?;
507 if rows.total_rows == 0 {
508 return Err(no_rows_error(&sample.scope));
509 }
510 Ok(rows)
511}
512
513pub fn no_rows_error(scope: &QualityScope) -> Report {
516 if *scope == QualityScope::CurrentView {
517 Report::msg("The table has no rows to sample")
518 } else {
519 Report::msg(format!(
520 "No rows match {}; change Rows from in the Sample form (s)",
521 scope.label()
522 ))
523 }
524}
525
526pub(crate) fn read_rows(
527 lf: &LazyFrame,
528 sample: &Sample,
529 known_total: Option<usize>,
530 polars_streaming: bool,
531) -> Result<AnalysisRows> {
532 read_rows_watched(lf, sample, known_total, polars_streaming, None)
533}
534
535pub(crate) fn read_rows_watched(
537 lf: &LazyFrame,
538 sample: &Sample,
539 known_total: Option<usize>,
540 polars_streaming: bool,
541 watch: Option<&ReadWatch>,
542) -> Result<AnalysisRows> {
543 acquire(lf, sample, known_total, polars_streaming, watch, None).map(|read| read.rows)
544}
545
546pub(crate) struct SampledRows {
548 pub rows: AnalysisRows,
549 pub positions: Vec<IdxSize>,
552 pub counted: Option<Counted>,
555}
556
557pub(crate) fn acquire(
560 lf: &LazyFrame,
561 sample: &Sample,
562 known_total: Option<usize>,
563 polars_streaming: bool,
564 watch: Option<&ReadWatch>,
565 count: Option<&Expr>,
566) -> Result<SampledRows> {
567 let n = sample.rows.max(1);
568 match &sample.method {
569 SampleMethod::EveryRow => sample_rows_counting(
570 lf,
571 None,
572 known_total,
573 sample.seed,
574 polars_streaming,
575 watch,
576 count,
577 ),
578 SampleMethod::Spread => sample_rows_counting(
579 lf,
580 Some(n),
581 known_total,
582 sample.seed,
583 polars_streaming,
584 watch,
585 count,
586 ),
587 SampleMethod::FirstRows => {
588 let df = collect_lazy(lf.clone().limit(n as IdxSize), polars_streaming)
589 .map_err(Report::from)?;
590 let height = df.height();
591 let sampled = match known_total {
593 Some(total) => total > height,
594 None => height == n,
595 };
596 Ok(SampledRows {
597 positions: (0..height as IdxSize).collect(),
598 rows: AnalysisRows {
599 df,
600 total_rows: known_total.unwrap_or(height),
601 sample_size: sampled.then_some(height),
602 per_value: None,
603 },
604 counted: None,
605 })
606 }
607 SampleMethod::PerPartition { column } => {
608 let read =
609 per_group_sample_within(lf, column, n, sample.seed, MAX_GROUP_ROWS, watch, count)?;
610 let sample_size = (read.seen > read.df.height()).then_some(read.df.height());
611 Ok(SampledRows {
612 rows: AnalysisRows {
613 df: read.df,
614 total_rows: read.seen,
615 sample_size,
616 per_value: Some(read.per_value),
617 },
618 positions: read.positions,
619 counted: read.counted,
620 })
621 }
622 }
623}
624
625#[derive(Debug, Clone, Default, PartialEq)]
627pub struct PerValue {
628 pub kept: usize,
631 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(
651 lf: &LazyFrame,
652 column: &str,
653 n: usize,
654 seed: u64,
655 limit: usize,
656 watch: Option<&ReadWatch>,
657 count: Option<&Expr>,
658) -> Result<GroupRead> {
659 let schema = lf.clone().collect_schema()?;
660 if schema.get(column).is_none() {
661 return Err(Report::msg(format!(
662 "partition column {column:?} is not in the rows sampled; choose another"
663 )));
664 }
665 let held = watch.cloned();
666 let state = stream_fold(
667 with_count_key(lf.clone(), count).with_row_index(GROUP_POSITION, None),
668 watch,
669 true,
670 GroupState {
671 column: column.to_string(),
672 cap: n,
673 limit,
674 seed,
675 ..Default::default()
676 },
677 move |state, batch| {
678 state.observe(batch)?;
679 if let Some(watch) = &held {
680 watch.hold(state.bytes(), state.held);
681 }
682 Ok(false)
683 },
684 )?;
685 if let Some(watch) = watch {
686 watch.check()?;
687 }
688 let seen = state.seen;
689 let cap = state
692 .cap
693 .min((state.limit / state.groups.len().max(1)).max(1));
694 let mut totals = std::collections::BTreeMap::new();
695 let mut out: Option<DataFrame> = None;
696 for (key, mut group) in state.groups {
697 totals.insert(key, group.total);
698 group.trim(cap)?;
699 let Some(rows) = group.rows else {
700 continue;
701 };
702 out = Some(match out {
703 Some(frame) => frame.vstack(&rows)?,
704 None => rows,
705 });
706 }
707 let (df, positions) = match out {
708 Some(df) => {
709 let df = df.sort([GROUP_POSITION], SortMultipleOptions::default())?;
710 let positions = df
711 .column(GROUP_POSITION)?
712 .idx()?
713 .into_no_null_iter()
714 .collect();
715 (df.drop(GROUP_POSITION)?, positions)
716 }
717 None => (
718 collect_lazy(lf.clone().limit(0), true).map_err(Report::from)?,
719 Vec::new(),
720 ),
721 };
722 Ok(GroupRead {
723 df,
724 seen,
725 per_value: PerValue { kept: cap, totals },
726 positions,
727 counted: count.is_some().then(|| state.counter.finish()),
728 })
729}
730
731#[derive(Default)]
732struct GroupState {
733 column: String,
734 cap: usize,
737 limit: usize,
739 seed: u64,
740 seen: usize,
741 held: usize,
743 groups: HashMap<Option<String>, GroupSample>,
744 counter: KeyCounter,
745}
746
747#[derive(Default)]
748struct GroupSample {
749 rows: Option<DataFrame>,
750 ranks: Vec<u64>,
751 total: usize,
753}
754
755impl GroupSample {
756 fn trim(&mut self, cap: usize) -> PolarsResult<usize> {
759 if self.ranks.len() <= cap {
760 return Ok(0);
761 }
762 let mut order: Vec<usize> = (0..self.ranks.len()).collect();
763 order.sort_unstable_by_key(|i| self.ranks[*i]);
764 order.truncate(cap);
765 let take: Vec<IdxSize> = order.iter().map(|i| *i as IdxSize).collect();
766 if let Some(kept) = self.rows.take() {
767 self.rows = Some(kept.take(&IdxCa::from_vec("kept".into(), take))?);
768 }
769 let removed = self.ranks.len() - cap;
770 self.ranks = order.iter().map(|i| self.ranks[*i]).collect();
771 Ok(removed)
772 }
773}
774
775impl GroupState {
776 fn bytes(&self) -> u64 {
778 self.groups
779 .values()
780 .filter_map(|group| group.rows.as_ref())
781 .map(|rows| rows.estimated_size() as u64)
782 .sum()
783 }
784
785 fn observe(&mut self, mut batch: DataFrame) -> PolarsResult<()> {
786 self.counter.observe(&mut batch)?;
787 self.seen += batch.height();
788 let positions = batch.column(GROUP_POSITION)?.idx()?.clone();
789 let keys = batch.column(&self.column)?.as_materialized_series().clone();
790 let mut by_key: HashMap<Option<String>, (Vec<IdxSize>, Vec<u64>)> = HashMap::new();
791 for (index, (key, position)) in keys.iter().zip(positions.into_no_null_iter()).enumerate() {
792 let key = (!key.is_null()).then(|| crate::exact::str_value(&key).into_owned());
794 let entry = by_key.entry(key).or_default();
795 entry.0.push(index as IdxSize);
796 entry.1.push(sample_rank(self.seed, position as u64));
797 }
798 for (key, (indices, ranks)) in by_key {
799 if !self.groups.contains_key(&key) && self.groups.len() >= MAX_GROUPS {
800 return Err(PolarsError::ComputeError(
801 format!(
802 "more than {MAX_GROUPS} values of {}; sample per a coarser column",
803 self.column
804 )
805 .into(),
806 ));
807 }
808 let group = self.groups.entry(key).or_default();
809 group.total += indices.len();
810 self.held += indices.len();
811 let rows = batch.take(&IdxCa::from_vec("picked".into(), indices))?;
812 group.rows = Some(match group.rows.take() {
813 Some(kept) => kept.vstack(&rows)?,
814 None => rows,
815 });
816 group.ranks.extend(ranks);
817 self.held -= group.trim(self.cap)?;
818 }
819 if self.held > self.limit.saturating_add(self.limit / 4) {
822 self.cap = self.cap.min((self.limit / self.groups.len()).max(1));
823 for group in self.groups.values_mut() {
824 self.held -= group.trim(self.cap)?;
825 }
826 }
827 Ok(())
828 }
829}
830
831pub struct AnalysisRows {
833 pub df: DataFrame,
834 pub total_rows: usize,
835 pub sample_size: Option<usize>,
837 pub per_value: Option<PerValue>,
839}
840
841const SAMPLE_BLOCKS: usize = 50;
844
845const SAMPLE_READERS: usize = 8;
847
848const SAMPLE_POSITION: &str = "__datui_sample_position";
850
851pub fn count_rows(lf: &LazyFrame, polars_streaming: bool) -> Result<usize> {
853 let count_df =
854 collect_lazy(crate::table::row_count_lf(lf), polars_streaming).map_err(Report::from)?;
855 Ok(match count_df.get(0).and_then(|row| row.first().cloned()) {
856 Some(AnyValue::UInt64(n)) => n as usize,
857 Some(AnyValue::UInt32(n)) => n as usize,
858 _ => 0,
859 })
860}
861
862pub fn analysis_rows(
871 lf: &LazyFrame,
872 sample_rows: Option<usize>,
873 known_total: Option<usize>,
874 seed: u64,
875 polars_streaming: bool,
876) -> Result<AnalysisRows> {
877 analysis_rows_watched(lf, sample_rows, known_total, seed, polars_streaming, None)
878}
879
880pub(crate) fn analysis_rows_watched(
883 lf: &LazyFrame,
884 sample_rows: Option<usize>,
885 known_total: Option<usize>,
886 seed: u64,
887 polars_streaming: bool,
888 watch: Option<&ReadWatch>,
889) -> Result<AnalysisRows> {
890 sample_rows_counting(
891 lf,
892 sample_rows,
893 known_total,
894 seed,
895 polars_streaming,
896 watch,
897 None,
898 )
899 .map(|read| read.rows)
900}
901
902pub(crate) fn sample_rows_counting(
906 lf: &LazyFrame,
907 sample_rows: Option<usize>,
908 known_total: Option<usize>,
909 seed: u64,
910 polars_streaming: bool,
911 watch: Option<&ReadWatch>,
912 count: Option<&Expr>,
913) -> Result<SampledRows> {
914 let whole = |df: DataFrame, total_rows: usize| SampledRows {
915 positions: (0..df.height() as IdxSize).collect(),
916 rows: AnalysisRows {
917 df,
918 total_rows,
919 sample_size: None,
920 per_value: None,
921 },
922 counted: None,
923 };
924 let Some(n) = sample_rows.filter(|n| *n > 0) else {
925 let df = collect_lazy(lf.clone(), polars_streaming).map_err(Report::from)?;
926 let total_rows = df.height();
927 return Ok(whole(df, total_rows));
928 };
929 if !slices_reach_into_the_scan(lf) {
930 let read = stream_sample(lf, n, seed, watch, count)?;
931 let sample_size = (read.seen > n).then_some(read.df.height());
932 return Ok(SampledRows {
933 rows: AnalysisRows {
934 df: read.df,
935 total_rows: read.seen,
936 sample_size,
937 per_value: None,
938 },
939 positions: read.positions,
940 counted: read.counted,
941 });
942 }
943 let total_rows = match known_total {
944 Some(total) => total,
945 None => count_rows(lf, polars_streaming)?,
946 };
947 if total_rows <= n {
948 let df = collect_lazy(lf.clone(), polars_streaming).map_err(Report::from)?;
949 return Ok(whole(df, total_rows));
950 }
951 let along = Along {
952 watch,
953 count,
954 on_run: None,
955 };
956 let read = block_sample(lf, total_rows, n, seed, polars_streaming, along)?;
957 Ok(SampledRows {
958 rows: AnalysisRows {
959 sample_size: Some(read.df.height()),
960 df: read.df,
961 total_rows,
962 per_value: None,
963 },
964 positions: read.positions,
965 counted: read.counted,
966 })
967}
968
969pub fn slices_reach_into_the_scan(lf: &LazyFrame) -> bool {
976 let Ok(plan) = lf.clone().slice(1, 1).describe_optimized_plan() else {
977 return false;
978 };
979 let scans: Vec<&str> = plan
980 .lines()
981 .map(str::trim)
982 .filter(|l| l.starts_with("Parquet SCAN") || l.starts_with("IPC SCAN"))
983 .collect();
984 let [scan] = scans.as_slice() else {
985 return false;
986 };
987 let one_source = !scan.contains("other sources") && !scan.contains(", ");
988 let total_scans = plan.matches(" SCAN").count();
989 one_source && total_scans == 1 && plan.contains("SLICE: Positive") && !plan.contains("SLICE[")
990}
991
992pub(crate) fn block_sample_live(
996 lf: &LazyFrame,
997 total_rows: usize,
998 n: usize,
999 seed: u64,
1000 polars_streaming: bool,
1001 watch: &ReadWatch,
1002 on_run: &OnRun<'_>,
1003) -> Result<Option<DataFrame>> {
1004 let along = Along {
1005 watch: Some(watch),
1006 count: None,
1007 on_run: Some(on_run),
1008 };
1009 let read = block_sample(lf, total_rows, n, seed, polars_streaming, along)?;
1010 Ok((total_rows < 2 * n).then_some(read.df))
1011}
1012
1013struct Along<'a> {
1016 watch: Option<&'a ReadWatch>,
1017 count: Option<&'a Expr>,
1018 on_run: Option<&'a OnRun<'a>>,
1019}
1020
1021pub(crate) type OnRun<'a> = dyn Fn(usize, &DataFrame) + Sync + 'a;
1023
1024fn block_sample(
1028 lf: &LazyFrame,
1029 total_rows: usize,
1030 n: usize,
1031 seed: u64,
1032 polars_streaming: bool,
1033 along: Along<'_>,
1034) -> Result<StreamRead> {
1035 let Along {
1036 watch,
1037 count,
1038 on_run,
1039 } = along;
1040 if total_rows < 2 * n {
1043 let df = collect_lazy(lf.clone(), polars_streaming).map_err(Report::from)?;
1044 let counted = match count {
1045 Some(key) => {
1046 let mut keys = df
1047 .clone()
1048 .lazy()
1049 .select([key.clone().alias(COUNT_KEY)])
1050 .collect()?;
1051 let mut counter = KeyCounter::default();
1052 counter.observe(&mut keys)?;
1053 Some(counter.finish())
1054 }
1055 None => None,
1056 };
1057 let mut ranked: Vec<(u64, IdxSize)> = (0..df.height())
1058 .map(|i| (sample_rank(seed, i as u64), i as IdxSize))
1059 .collect();
1060 ranked.sort_unstable();
1061 let mut keep: Vec<IdxSize> = ranked.into_iter().take(n).map(|(_, i)| i).collect();
1062 keep.sort_unstable();
1063 let df = df.take(&IdxCa::from_vec("sample".into(), keep.clone()))?;
1064 return Ok(StreamRead {
1065 df,
1066 seen: total_rows,
1067 positions: keep,
1068 counted,
1069 });
1070 }
1071 let blocks = SAMPLE_BLOCKS.min(n).max(1);
1072 let stride = total_rows / blocks;
1075 let runs: Vec<(usize, usize)> = (0..blocks)
1076 .map(|block| {
1077 let run = (block + 1) * n / blocks - block * n / blocks;
1078 let room = stride.saturating_sub(run) as u64;
1079 let offset = block * stride + (sample_rank(seed, block as u64) % (room + 1)) as usize;
1080 (offset, run)
1081 })
1082 .collect();
1083 let next = std::sync::atomic::AtomicUsize::new(0);
1084 let read: Vec<Result<(usize, DataFrame)>> = std::thread::scope(|scope| {
1085 let workers: Vec<_> = (0..SAMPLE_READERS.min(blocks))
1086 .map(|_| {
1087 scope.spawn(|| {
1088 let mut read = Vec::new();
1089 loop {
1090 if watch.is_some_and(|watch| watch.stopped()) {
1091 break;
1092 }
1093 let block = next.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1094 let Some((offset, run)) = runs.get(block) else {
1095 break;
1096 };
1097 let rows = collect_lazy(
1098 lf.clone().slice(*offset as i64, *run as IdxSize),
1099 polars_streaming,
1100 )
1101 .map(|df| {
1102 if let Some(watch) = watch {
1103 watch.saw(df.height());
1104 }
1105 if let Some(on_run) = on_run {
1106 on_run(*offset, &df);
1107 }
1108 (block, df)
1109 })
1110 .map_err(Report::from);
1111 read.push(rows);
1112 }
1113 read
1114 })
1115 })
1116 .collect();
1117 workers
1118 .into_iter()
1119 .flat_map(|worker| match worker.join() {
1120 Ok(read) => read,
1121 Err(_) => vec![Err(Report::msg("a sample reader failed"))],
1123 })
1124 .collect()
1125 });
1126 if let Some(watch) = watch {
1127 watch.check()?;
1128 }
1129 let mut read = read.into_iter().collect::<Result<Vec<_>>>()?;
1130 read.sort_by_key(|(block, _)| *block);
1131 let mut out: Option<DataFrame> = None;
1132 let mut positions = Vec::with_capacity(n);
1133 for (block, rows) in read {
1134 let offset = runs[block].0;
1135 positions.extend((0..rows.height()).map(|row| (offset + row) as IdxSize));
1136 out = Some(match out {
1137 Some(frame) => frame.vstack(&rows)?,
1138 None => rows,
1139 });
1140 }
1141 Ok(StreamRead {
1142 df: out.unwrap_or_default(),
1143 seen: total_rows,
1144 positions,
1145 counted: None,
1146 })
1147}
1148
1149struct StreamRead {
1151 df: DataFrame,
1152 seen: usize,
1154 positions: Vec<IdxSize>,
1155 counted: Option<Counted>,
1156}
1157
1158pub(crate) fn stream_batches(
1162 lf: LazyFrame,
1163 watch: Option<&ReadWatch>,
1164 maintain_order: bool,
1165 on_batch: impl Fn(DataFrame) -> PolarsResult<bool> + Send + Sync + 'static,
1166) -> Result<()> {
1167 let watch = watch.cloned();
1168 let sink = lf.sink_batches(
1169 PlanCallback::new(move |batch: DataFrame| {
1170 if let Some(watch) = &watch {
1171 if watch.stopped() {
1172 return Ok(true);
1173 }
1174 watch.saw(batch.height());
1175 }
1176 on_batch(batch)
1177 }),
1178 maintain_order,
1179 None,
1180 )?;
1181 collect_lazy(sink, true).map_err(Report::from)?;
1182 Ok(())
1183}
1184
1185pub(crate) fn stream_fold<S: Send + 'static>(
1188 lf: LazyFrame,
1189 watch: Option<&ReadWatch>,
1190 maintain_order: bool,
1191 state: S,
1192 observe: impl Fn(&mut S, DataFrame) -> PolarsResult<bool> + Send + Sync + 'static,
1193) -> Result<S> {
1194 let shared = std::sync::Arc::new(std::sync::Mutex::new(Some(state)));
1195 let held = std::sync::Arc::clone(&shared);
1196 stream_batches(lf, watch, maintain_order, move |batch| {
1197 let mut state = held
1198 .lock()
1199 .map_err(|_| PolarsError::ComputeError("a streamed read failed".into()))?;
1200 match state.as_mut() {
1201 Some(state) => observe(state, batch),
1202 None => Ok(true),
1203 }
1204 })?;
1205 let state = shared
1206 .lock()
1207 .map_err(|_| Report::msg("a streamed read failed"))?
1208 .take();
1209 state.ok_or_else(|| Report::msg("a streamed read failed"))
1210}
1211
1212fn stream_sample(
1215 lf: &LazyFrame,
1216 n: usize,
1217 seed: u64,
1218 watch: Option<&ReadWatch>,
1219 count: Option<&Expr>,
1220) -> Result<StreamRead> {
1221 let held = watch.cloned();
1222 let mut reservoir = stream_fold(
1223 with_count_key(lf.clone(), count).with_row_index(SAMPLE_POSITION, None),
1224 watch,
1225 true,
1226 Reservoir::new(n, seed),
1227 move |reservoir, batch| {
1228 reservoir.observe(batch)?;
1229 if let Some(watch) = &held {
1230 let kept = reservoir.kept.as_ref();
1231 watch.hold(
1232 kept.map_or(0, |kept| kept.estimated_size() as u64),
1233 kept.map_or(0, DataFrame::height),
1234 );
1235 }
1236 Ok(false)
1237 },
1238 )?;
1239 if let Some(watch) = watch {
1240 watch.check()?;
1241 }
1242 let seen = reservoir.seen;
1243 let counted = count
1244 .is_some()
1245 .then(|| std::mem::take(&mut reservoir.counter).finish());
1246 let (df, positions) = match reservoir.finish()? {
1247 Some(kept) => kept,
1248 None => (
1250 collect_lazy(lf.clone().limit(0), true).map_err(Report::from)?,
1251 Vec::new(),
1252 ),
1253 };
1254 Ok(StreamRead {
1255 df,
1256 seen,
1257 positions,
1258 counted,
1259 })
1260}
1261
1262#[derive(Default)]
1265struct Reservoir {
1266 n: usize,
1267 seed: u64,
1268 seen: usize,
1269 kept: Option<DataFrame>,
1270 ranks: Vec<u64>,
1271 bar: u64,
1273 counter: KeyCounter,
1274}
1275
1276impl Reservoir {
1277 fn new(n: usize, seed: u64) -> Self {
1278 Self {
1279 n,
1280 seed,
1281 bar: u64::MAX,
1282 ..Default::default()
1283 }
1284 }
1285
1286 fn observe(&mut self, mut batch: DataFrame) -> PolarsResult<()> {
1287 self.counter.observe(&mut batch)?;
1288 self.seen += batch.height();
1289 let positions = batch.column(SAMPLE_POSITION)?.idx()?;
1290 let mut picked = Vec::new();
1291 let mut ranks = Vec::new();
1292 for (index, position) in positions.into_no_null_iter().enumerate() {
1293 let rank = sample_rank(self.seed, position as u64);
1294 if rank < self.bar {
1295 picked.push(index as IdxSize);
1296 ranks.push(rank);
1297 }
1298 }
1299 if picked.is_empty() {
1300 return Ok(());
1301 }
1302 let rows = batch.take(&IdxCa::from_vec("picked".into(), picked))?;
1303 self.kept = Some(match self.kept.take() {
1304 Some(kept) => kept.vstack(&rows)?,
1305 None => rows,
1306 });
1307 self.ranks.extend(ranks);
1308 if self.ranks.len() > 2 * self.n {
1309 self.prune()?;
1310 }
1311 Ok(())
1312 }
1313
1314 fn prune(&mut self) -> PolarsResult<()> {
1316 let Some(kept) = self.kept.take() else {
1317 return Ok(());
1318 };
1319 let mut order: Vec<usize> = (0..self.ranks.len()).collect();
1320 order.sort_unstable_by_key(|i| self.ranks[*i]);
1321 order.truncate(self.n);
1322 let take: Vec<IdxSize> = order.iter().map(|i| *i as IdxSize).collect();
1323 self.kept = Some(kept.take(&IdxCa::from_vec("kept".into(), take))?);
1324 self.ranks = order.iter().map(|i| self.ranks[*i]).collect();
1325 if self.ranks.len() == self.n {
1326 self.bar = self.ranks.iter().copied().max().unwrap_or(u64::MAX);
1327 }
1328 Ok(())
1329 }
1330
1331 fn finish(mut self) -> PolarsResult<Option<(DataFrame, Vec<IdxSize>)>> {
1334 self.prune()?;
1335 let Some(kept) = self.kept else {
1336 return Ok(None);
1337 };
1338 let sorted = kept.sort([SAMPLE_POSITION], SortMultipleOptions::default())?;
1339 let positions = sorted
1340 .column(SAMPLE_POSITION)?
1341 .idx()?
1342 .into_no_null_iter()
1343 .collect();
1344 Ok(Some((sorted.drop(SAMPLE_POSITION)?, positions)))
1345 }
1346}
1347
1348pub(crate) fn sample_rank(seed: u64, position: u64) -> u64 {
1351 let mut value = seed ^ position.wrapping_mul(0x9e37_79b9_7f4a_7c15);
1352 value = (value ^ (value >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
1353 value = (value ^ (value >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
1354 value ^ (value >> 31)
1355}
1356
1357#[cfg(test)]
1358mod sampler_tests;
1359
1360#[cfg(test)]
1361mod tests {
1362 #[test]
1363 fn a_size_takes_shorthand() {
1364 use super::parse_size;
1365 for (text, rows) in [
1366 ("50000", 50_000),
1367 ("50,000", 50_000),
1368 ("1_000", 1_000),
1369 ("50k", 50_000),
1370 ("250K", 250_000),
1371 ("2m", 2_000_000),
1372 ("2.5M", 2_500_000),
1373 (" 7 ", 7),
1374 ("99999999999999999999999", usize::MAX),
1375 ] {
1376 assert_eq!(parse_size(text), Ok(rows), "{text}");
1377 }
1378 for bad in ["", "k", "12x", "1.5", "-3", "1e6", "2mm"] {
1379 assert_eq!(parse_size(bad), Err(super::SizeError::NotASize), "{bad}");
1380 }
1381 for zero in ["0", "0k", "0.0001k"] {
1382 assert_eq!(parse_size(zero), Err(super::SizeError::Zero), "{zero}");
1383 }
1384 }
1385
1386 use super::*;
1387
1388 fn table() -> LazyFrame {
1389 let sizes = [("a", 9_000usize), ("b", 900), ("c", 100)];
1391 let mut part = Vec::new();
1392 let mut value = Vec::new();
1393 for (name, size) in sizes {
1394 for row in 0..size {
1395 part.push(name);
1396 value.push(row as i64);
1397 }
1398 }
1399 df!("part" => part, "value" => value).unwrap().lazy()
1400 }
1401
1402 fn sample(method: SampleMethod, rows: usize) -> Sample {
1403 Sample {
1404 method,
1405 rows,
1406 ..Sample::default()
1407 }
1408 }
1409
1410 #[test]
1413 fn per_partition_keeps_up_to_n_from_each_value() {
1414 let method = SampleMethod::PerPartition {
1415 column: "part".to_string(),
1416 };
1417 let rows = read(&table(), &sample(method, 200), None, false).unwrap();
1418 assert_eq!(rows.total_rows, 10_000);
1419 let counts = rows
1420 .df
1421 .column("part")
1422 .unwrap()
1423 .as_materialized_series()
1424 .value_counts(true, true, "n".into(), false)
1425 .unwrap();
1426 let n = |part: &str| {
1427 (0..counts.height())
1428 .find(|row| {
1429 counts
1430 .column("part")
1431 .unwrap()
1432 .get(*row)
1433 .unwrap()
1434 .str_value()
1435 == part
1436 })
1437 .map(|row| {
1438 counts
1439 .column("n")
1440 .unwrap()
1441 .get(row)
1442 .unwrap()
1443 .try_extract::<u32>()
1444 .unwrap()
1445 })
1446 .unwrap()
1447 };
1448 assert_eq!((n("a"), n("b"), n("c")), (200, 200, 100));
1449 assert_eq!(rows.sample_size, Some(500));
1450 assert!(
1451 rows.df.column(GROUP_POSITION).is_err(),
1452 "no helper column leaks"
1453 );
1454 }
1455
1456 #[test]
1460 fn per_partition_past_the_limit_keeps_fewer_of_each_value() {
1461 let GroupRead {
1462 df,
1463 seen,
1464 per_value,
1465 ..
1466 } = per_group_sample_within(&table(), "part", 500, 42_891, 999, None, None).unwrap();
1467 assert_eq!(seen, 10_000);
1468 assert_eq!(per_value.kept, 333);
1469 assert_eq!(df.height(), 333 + 333 + 100);
1470 assert_eq!(
1471 per_value.totals,
1472 [("a", 9_000), ("b", 900), ("c", 100)]
1473 .into_iter()
1474 .map(|(part, rows)| (Some(part.to_string()), rows))
1475 .collect()
1476 );
1477 let asked =
1478 per_group_sample_within(&table(), "part", 333, 42_891, usize::MAX, None, None).unwrap();
1479 assert!(df.equals(&asked.df), "the rows a sample of 333 each keeps");
1480
1481 let lowered = Sample {
1482 method: SampleMethod::PerPartition {
1483 column: "part".into(),
1484 },
1485 rows: 500,
1486 ..Sample::default()
1487 };
1488 assert_eq!(
1489 lowered.outcome(10_000, Some(766), Some(333)),
1490 "766 rows, up to 333 per part (lowered from 500), of 10,000"
1491 );
1492 }
1493
1494 #[test]
1499 fn a_sample_says_where_its_rows_sat_and_a_stream_counts_on_the_way() {
1500 let lf = table()
1502 .with_column((col("value") % lit(4)).alias("quarter"))
1503 .with_row_index("row", None);
1504 let totals: std::collections::BTreeMap<_, _> = [("a", 9_000), ("b", 900), ("c", 100)]
1505 .into_iter()
1506 .map(|(part, rows)| (Some(part.to_string()), rows))
1507 .collect();
1508 for method in [
1509 SampleMethod::Spread,
1510 SampleMethod::PerPartition {
1511 column: "quarter".to_string(),
1512 },
1513 SampleMethod::FirstRows,
1514 ] {
1515 let read = acquire(
1516 &lf,
1517 &sample(method.clone(), 50),
1518 None,
1519 false,
1520 None,
1521 Some(&col("part")),
1522 )
1523 .unwrap();
1524 let rows: Vec<IdxSize> = read
1525 .rows
1526 .df
1527 .column("row")
1528 .unwrap()
1529 .idx()
1530 .unwrap()
1531 .into_no_null_iter()
1532 .collect();
1533 assert_eq!(rows, read.positions, "{method:?}");
1534 assert!(read.rows.df.column(COUNT_KEY).is_err(), "{method:?}");
1535 if method == SampleMethod::FirstRows {
1536 assert_eq!(read.counted, None);
1537 } else {
1538 assert_eq!(
1539 read.counted,
1540 Some(Counted::Totals(totals.clone())),
1541 "{method:?}"
1542 );
1543 }
1544 }
1545 }
1546
1547 #[test]
1550 fn a_count_past_its_limit_gives_up_and_says_so() {
1551 let mut counter = KeyCounter::with_limit(2);
1552 let mut batch = df!(COUNT_KEY => ["a", "b", "c"], "value" => [1, 2, 3]).unwrap();
1553 counter.observe(&mut batch).unwrap();
1554 assert_eq!(batch.get_column_names(), ["value"]);
1555 assert_eq!(counter.finish(), Counted::TooMany);
1556 }
1557
1558 #[test]
1559 fn first_rows_is_the_head_and_every_row_is_all_of_it() {
1560 let head = read(
1561 &table(),
1562 &sample(SampleMethod::FirstRows, 50),
1563 Some(10_000),
1564 false,
1565 )
1566 .unwrap();
1567 assert_eq!(head.df.height(), 50);
1568 assert_eq!(head.sample_size, Some(50));
1569 assert_eq!(
1570 head.df.column("value").unwrap().i64().unwrap().get(49),
1571 Some(49)
1572 );
1573 let all = read(&table(), &sample(SampleMethod::EveryRow, 50), None, false).unwrap();
1574 assert_eq!((all.df.height(), all.sample_size), (10_000, None));
1575 }
1576
1577 #[test]
1578 fn a_seeded_per_partition_sample_repeats() {
1579 let method = SampleMethod::PerPartition {
1580 column: "part".to_string(),
1581 };
1582 let one = read(&table(), &sample(method.clone(), 50), None, false).unwrap();
1583 let two = read(&table(), &sample(method.clone(), 50), None, false).unwrap();
1584 assert!(one.df.equals(&two.df));
1585 let other = read(
1586 &table(),
1587 &Sample {
1588 seed: 7,
1589 ..sample(method, 50)
1590 },
1591 None,
1592 false,
1593 )
1594 .unwrap();
1595 assert!(!one.df.equals(&other.df));
1596 }
1597
1598 #[test]
1601 fn a_scope_that_matches_nothing_is_an_error() {
1602 let scope = QualityScope::parse_command("partition part=zzz").unwrap();
1603 let lf = SampleSource::view(table()).cut(&scope).unwrap();
1604 let Err(error) = read(
1605 &lf,
1606 &Sample {
1607 scope,
1608 ..Sample::default()
1609 },
1610 None,
1611 false,
1612 ) else {
1613 panic!("a scope that matches nothing must not sample");
1614 };
1615 assert!(error.to_string().contains("No rows match"), "{error}");
1616 }
1617
1618 #[test]
1619 fn the_summary_says_what_will_be_read() {
1620 let middot = crate::glyphs::get().middot;
1621 assert_eq!(
1622 Sample::default().summary(),
1623 format!("100,000 random rows {middot} current view {middot} seed 42891")
1624 );
1625 assert_eq!(
1626 sample(SampleMethod::FirstRows, 1_000).summary(),
1627 format!("first 1,000 rows {middot} current view")
1628 );
1629 assert_eq!(
1631 Sample::default().summary_within(Some(1_000)),
1632 format!("all 1,000 rows {middot} current view")
1633 );
1634 assert_eq!(
1635 Sample::default().summary_within(Some(1_000_000)),
1636 Sample::default().summary()
1637 );
1638 assert_eq!(
1639 Sample::default().summary_within(None),
1640 Sample::default().summary()
1641 );
1642 }
1643}