Skip to main content

datui_lib/
sampling.rs

1//! The one sampler every analysis tool reads through: which rows (a scope), how they
2//! are picked (a method), how many, and the seed. Describe, Distribution, Correlation
3//! and Data Quality all take their rows from [`read`], so a sample means the same
4//! thing whichever tool shows it.
5
6use 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
16/// The default sample size, before `[analysis] sample_rows` says otherwise.
17pub const DEFAULT_SAMPLE_ROWS: usize = 100_000;
18
19/// Why a typed sample size cannot be read.
20#[derive(Debug, Clone, Copy, PartialEq, Eq)]
21pub enum SizeError {
22    NotASize,
23    Zero,
24}
25
26impl SizeError {
27    /// A few words, for a row with little room.
28    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
45/// A typed sample size: `50000`, `50,000`, `50_000`, `50k`, `2m`, `2.5M`. A size of
46/// no rows is refused; past `usize` it saturates and the caller clamps.
47pub 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        // Whole digits: exact, however long.
67        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
81/// A per-partition sample keeps at most this many partitions and rows in memory. Past
82/// the rows it keeps fewer of each value; past the partitions it is refused, which a
83/// column with that many values meets within its first few batches.
84const MAX_GROUPS: usize = 10_000;
85const MAX_GROUP_ROWS: usize = 2_000_000;
86
87/// The row index the per-partition sampler ranks rows by, dropped before anyone sees it.
88const GROUP_POSITION: &str = "__datui_group_sample_position";
89
90/// The key a streamed pass counts rows by, computed beside the rows and taken off
91/// each batch before the sampler keeps any of it.
92pub(crate) const COUNT_KEY: &str = "__datui_count_key";
93
94/// Distinct keys a pass counts before it gives up counting. Past this the grain is
95/// finer than a report can show, and the map would grow with the table.
96pub const MAX_COUNTED_KEYS: usize = 1_000_000;
97
98/// What a streamed pass counted beside its sample.
99#[derive(Debug, Clone, PartialEq)]
100pub enum Counted {
101    /// Every row of the scope by its key, named as a segment names its value
102    /// (`AnyValue::str_value`), `None` for null.
103    Totals(std::collections::BTreeMap<Option<String>, usize>),
104    /// More than [`MAX_COUNTED_KEYS`] keys: the count was dropped, the sample kept.
105    TooMany,
106}
107
108/// Rows by [`COUNT_KEY`], a batch at a time, bounded by [`MAX_COUNTED_KEYS`].
109#[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    /// Count `batch`'s keys and take the key off it, so the rows kept are the
132    /// table's own. A batch without the key is left as it is.
133    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        // Grouped in Polars first, so a key is turned into text once per batch
142        // rather than once per row.
143        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
173/// `lf` with `count`'s key beside its rows, for a pass that counts as it samples.
174pub(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
181/// What a read that was stopped says. Its work is dropped, never shown as a result.
182pub const CANCELLED: &str = "Cancelled";
183
184/// A read's line to the screen: told to stop, and telling how many rows it has seen.
185///
186/// Shared with the UI thread, which sets `stop` on a cancel and reads the count as
187/// it draws. A streamed read checks `stop` between batches and a seeded block read
188/// between blocks; a single collect cannot be stopped partway and runs to its end.
189#[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    /// Whether anything has counted rows yet: a read that cannot observe its batches
194    /// has no count to show, which is not a count of zero.
195    counted: std::sync::Arc<std::sync::atomic::AtomicBool>,
196    /// Judges the rows a sampler holds, in bytes, as it reads: the reason to stop
197    /// when they would not fit.
198    held: Option<HeldCheck>,
199    /// Why the held rows stopped the read, once they did.
200    memory: std::sync::Arc<std::sync::Mutex<Option<String>>>,
201}
202
203/// What [`ReadWatch::hold`] asks of the bytes a sampler holds and the rows they are.
204pub 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    /// A watch whose sampler stops, keeping what it holds, once `judge` says the
217    /// bytes it holds will not fit.
218    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    /// A sampler holds `bytes` in `rows` rows now: past what fits, the read stops.
226    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    /// Why memory stopped the read, if it did: its rows so far are kept.
237    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    /// Rows the read has seen so far, once it has counted any.
253    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    /// Start counting again for the next read, handing back what the last one
267    /// counted, if it counted anything.
268    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    /// Stopped: the read's partial rows are not a sample, so it fails instead. Not
277    /// when memory stopped it: what it holds is kept.
278    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/// How the rows of a scope are picked.
288#[derive(Debug, Clone, PartialEq, Eq, Default)]
289pub enum SampleMethod {
290    /// A seeded random sample spread across the whole scope.
291    #[default]
292    Spread,
293    /// Up to the sample size from each value of a column, so a small partition is
294    /// represented beside a large one.
295    PerPartition { column: String },
296    /// The first rows of the scope, in order: the quickest read, and only the head.
297    FirstRows,
298    /// Every row: no sampling.
299    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/// Which rows an analysis reads and how it picks them.
314#[derive(Debug, Clone, PartialEq, Eq)]
315pub struct Sample {
316    pub scope: QualityScope,
317    pub method: SampleMethod,
318    /// Rows to keep: in all, or per partition for [`SampleMethod::PerPartition`].
319    /// Ignored by [`SampleMethod::EveryRow`].
320    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    /// The sample as one short line: `100,000 spread · current view · seed 42891`.
337    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    /// The summary, against the `rows` the scope is known to hold: a sample of at
362    /// least that many reads every one of them, and says so, `all 1,000 rows`,
363    /// rather than promising 100,000 from a table of 1,000.
364    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    /// What was read, once it was: `sample of 100,000 of 36,839,175 rows`, then the
379    /// scope when it is not simply the table as shown. `per_value` is how many rows
380    /// an equal-per-value sample kept of each, when that was fewer than asked.
381    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
419/// Where a tool's rows come from before its scope cuts them: the table as shown, or
420/// the loaded source with what its footers said. Built on the UI thread; cut in the
421/// worker, since preparing a source scan can read its schema.
422pub struct SampleSource {
423    lf: LazyFrame,
424    source: Option<QualitySourceContext>,
425    from_source: bool,
426}
427
428impl SampleSource {
429    /// The table as it is shown (query and filters applied).
430    pub fn view(lf: LazyFrame) -> Self {
431        Self {
432            lf,
433            source: None,
434            from_source: false,
435        }
436    }
437
438    /// The loaded source, before any query or filter.
439    pub fn loaded(lf: LazyFrame, source: Option<QualitySourceContext>) -> Self {
440        Self {
441            lf,
442            source,
443            from_source: true,
444        }
445    }
446
447    /// The frame cut to `scope`, with only the table's own columns: the provenance
448    /// index a source scope needs to find its files is dropped once it has.
449    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
478/// How many rows a view scope holds, from the view's row count; `None` for a source
479/// scope, whose size only a read can tell.
480pub 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
492/// Read the rows `sample` asks for from a frame already cut to its scope.
493///
494/// `known_total` saves a count. The first-rows method reports the rows it read as the
495/// total when none is known, and says so through `sample_size: None`, rather than pay
496/// for a count the method exists to avoid.
497pub 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
510/// A chosen set of rows that matches nothing is a mistake to say, not an empty
511/// sample to analyze: a typed value that is not in the data, a range past its end.
512/// The table as shown may simply be empty, and says so itself.
513pub 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
533/// [`read_rows`], stopping when `watch` says to and counting the rows it streams.
534pub(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
544/// The rows a sample kept, where each sat in the frame, and what its pass counted.
545pub(crate) struct SampledRows {
546    pub rows: AnalysisRows,
547    /// Each kept row's position in the frame read, in the order of `rows.df`: what
548    /// cuts a sample into row chunks without reading it again.
549    pub positions: Vec<IdxSize>,
550    /// Rows by `count`'s key, when the pass that read the sample saw every row.
551    /// `None` when it did not (seeded runs, the head) or nothing was asked.
552    pub counted: Option<Counted>,
553}
554
555/// [`read_rows_watched`], keeping each row's position, and counting every row by
556/// `count` when the read is one streamed pass over all of them: the pass is being
557/// paid for anyway, and a second read of the key is what this saves.
558pub(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            // Without a count, a full head means there may be more: call it a sample.
591            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/// What an equal-per-value sample learned beside its rows, from the same pass.
625#[derive(Debug, Clone, Default, PartialEq)]
626pub struct PerValue {
627    /// Rows kept of each value: the size asked for, or fewer when that many of every
628    /// value would pass [`MAX_GROUP_ROWS`].
629    pub kept: usize,
630    /// Every row of the scope, counted by value as it streamed past. Keyed as a
631    /// segment names a value (`AnyValue::str_value`), `None` for null, so a segment
632    /// by the same column finds its count here instead of in a second read.
633    pub totals: std::collections::BTreeMap<Option<String>, usize>,
634}
635
636/// What [`per_group_sample_within`] read.
637struct GroupRead {
638    df: DataFrame,
639    seen: usize,
640    per_value: PerValue,
641    positions: Vec<IdxSize>,
642    counted: Option<Counted>,
643}
644
645/// Up to `n` seeded rows from each value of `column`, from one streamed pass, in table
646/// order, and how many rows there were, holding at most `limit` rows.
647///
648/// Never refused for keeping too many rows. Whether `n` of every value fits is only
649/// known once every value has been seen, which is the end of the read, and a read
650/// that ends in a refusal has been paid for and thrown away. So the size per value
651/// comes down as values arrive, to what [`MAX_GROUP_ROWS`] holds for all of them; each
652/// value keeps its lowest-ranked rows, which is a seeded uniform sample of it at any
653/// size, and [`PerValue::kept`] says what the size came down to.
654fn 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                // True stops the sink: a cancel ends the read at the next batch.
683                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    // Streaming whatever the setting: holding the table is what this is here to avoid.
702    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    // Every value the same size: the cap may have come down after a value was last
713    // trimmed, and one that arrived late was only ever held to the cap of its time.
714    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    /// Rows each value may keep: the size asked for until the values seen so far
758    /// would not all fit, then what does.
759    cap: usize,
760    /// Rows held across every value at most, once trimmed.
761    limit: usize,
762    seed: u64,
763    seen: usize,
764    /// Rows held across every value.
765    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    /// Rows of this value seen, kept or not.
775    total: usize,
776}
777
778impl GroupSample {
779    /// Keep the `cap` lowest-ranked rows: a seeded uniform sample of the value,
780    /// whatever order its rows arrived in. Returns how many went.
781    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    /// Bytes the rows held take.
800    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            // Named as a segment names its value, so the counts line up with segments.
816            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        // Past the row limit, every value's share comes down to what fits. A quarter
843        // over before trimming, so a run of new values costs a trim now and then
844        // rather than one per value.
845        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        // Three partitions of very different sizes, in order.
885        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    /// Per partition keeps up to n from each, so the small one is all there and the
906    /// large one does not crowd it out.
907    #[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    /// Past the row limit, every value keeps fewer rows rather than the read being
952    /// refused at its end: the same rows a sample asking for that many from the start
953    /// keeps, and every value's rows counted on the way.
954    #[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    /// Every sampler says where each kept row sat, which is what cuts a sample into
990    /// row chunks later without a read. A streamed pass counts every row by the key
991    /// it is given while it samples, the counts a read of the key would give; the head
992    /// does not see every row, so it counts nothing.
993    #[test]
994    fn a_sample_says_where_its_rows_sat_and_a_stream_counts_on_the_way() {
995        // Sampled per one column and counted by another.
996        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    /// Past its limit a count stops counting and says so; the batch it saw still
1043    /// loses its key, so the rows the sampler keeps are the table's.
1044    #[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    /// Rows chosen that match nothing are an error that names them, not an empty
1094    /// sample every tool would analyze as if it were the data.
1095    #[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        // A table smaller than the sample is read whole, and the line says so.
1125        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}