use crate::data_quality::{
QualityScope, QualitySourceContext, apply_quality_scope, prepare_source_quality_scan,
};
use crate::numfmt;
use crate::statistics::{AnalysisRows, collect_lazy, sample_rank};
use color_eyre::Result;
use color_eyre::eyre::Report;
use polars::prelude::*;
use std::collections::HashMap;
pub const DEFAULT_SAMPLE_ROWS: usize = 100_000;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SizeError {
NotASize,
Zero,
}
impl SizeError {
pub fn short(self) -> &'static str {
match self {
Self::NotASize => "not a size (50k, 2m)",
Self::Zero => "at least 1 row",
}
}
}
impl std::fmt::Display for SizeError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(match self {
Self::NotASize => "Sample size is a number of rows, like 50000, 50k or 2m",
Self::Zero => "Sample size is at least 1 row",
})
}
}
pub fn parse_size(text: &str) -> Result<usize, SizeError> {
let refuse = || SizeError::NotASize;
let cleaned: String = text
.trim()
.chars()
.filter(|c| !matches!(c, ',' | '_'))
.collect::<String>()
.to_ascii_lowercase();
let (number, scale) = match cleaned.strip_suffix('k') {
Some(n) => (n, 1e3),
None => match cleaned.strip_suffix('m') {
Some(n) => (n, 1e6),
None => (cleaned.as_str(), 1.0),
},
};
if number.is_empty() || !number.chars().all(|c| c.is_ascii_digit() || c == '.') {
return Err(refuse());
}
let rows = if scale == 1.0 {
if number.contains('.') {
return Err(refuse());
}
number.parse::<usize>().unwrap_or(usize::MAX)
} else {
let value: f64 = number.parse().map_err(|_| refuse())?;
(value * scale).round() as usize
};
if rows == 0 {
return Err(SizeError::Zero);
}
Ok(rows)
}
const MAX_GROUPS: usize = 10_000;
const MAX_GROUP_ROWS: usize = 2_000_000;
const GROUP_POSITION: &str = "__datui_group_sample_position";
pub(crate) const COUNT_KEY: &str = "__datui_count_key";
pub const MAX_COUNTED_KEYS: usize = 1_000_000;
#[derive(Debug, Clone, PartialEq)]
pub enum Counted {
Totals(std::collections::BTreeMap<Option<String>, usize>),
TooMany,
}
#[derive(Debug)]
pub(crate) struct KeyCounter {
totals: HashMap<Option<String>, usize>,
limit: usize,
too_many: bool,
}
impl Default for KeyCounter {
fn default() -> Self {
Self::with_limit(MAX_COUNTED_KEYS)
}
}
impl KeyCounter {
fn with_limit(limit: usize) -> Self {
Self {
totals: HashMap::new(),
limit,
too_many: false,
}
}
pub(crate) fn observe(&mut self, batch: &mut DataFrame) -> PolarsResult<()> {
if batch.column(COUNT_KEY).is_err() {
return Ok(());
}
let key = batch.drop_in_place(COUNT_KEY)?;
if self.too_many {
return Ok(());
}
let counts = key.as_materialized_series().value_counts(
false,
false,
"__datui_count_rows".into(),
false,
)?;
let keys = counts.column(COUNT_KEY)?;
let rows = counts.column("__datui_count_rows")?;
for row in 0..counts.height() {
let value = keys.get(row)?;
let key = (!value.is_null()).then(|| crate::exact::str_value(&value).into_owned());
let n = rows.get(row)?.extract::<usize>().unwrap_or(0);
*self.totals.entry(key).or_default() += n;
}
if self.totals.len() > self.limit {
self.too_many = true;
self.totals = HashMap::new();
}
Ok(())
}
pub(crate) fn finish(self) -> Counted {
if self.too_many {
Counted::TooMany
} else {
Counted::Totals(self.totals.into_iter().collect())
}
}
}
pub(crate) fn with_count_key(lf: LazyFrame, count: Option<&Expr>) -> LazyFrame {
match count {
Some(key) => lf.with_column(key.clone().alias(COUNT_KEY)),
None => lf,
}
}
pub const CANCELLED: &str = "Cancelled";
#[derive(Debug, Clone, Default)]
pub struct ReadWatch {
stop: std::sync::Arc<std::sync::atomic::AtomicBool>,
rows: std::sync::Arc<std::sync::atomic::AtomicUsize>,
counted: std::sync::Arc<std::sync::atomic::AtomicBool>,
held: Option<HeldCheck>,
memory: std::sync::Arc<std::sync::Mutex<Option<String>>>,
}
pub type HeldJudge = dyn Fn(u64, usize) -> Option<String> + Send + Sync;
#[derive(Clone)]
struct HeldCheck(std::sync::Arc<HeldJudge>);
impl std::fmt::Debug for HeldCheck {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("HeldCheck")
}
}
impl ReadWatch {
pub(crate) fn judging_held(judge: std::sync::Arc<HeldJudge>) -> Self {
Self {
held: Some(HeldCheck(judge)),
..Self::default()
}
}
pub(crate) fn hold(&self, bytes: u64, rows: usize) {
let Some(HeldCheck(judge)) = &self.held else {
return;
};
if let Some(reason) = judge(bytes, rows) {
*self.memory.lock().unwrap_or_else(|e| e.into_inner()) = Some(reason);
self.stop();
}
}
pub(crate) fn memory_stopped(&self) -> Option<String> {
self.memory
.lock()
.unwrap_or_else(|e| e.into_inner())
.clone()
}
pub fn stop(&self) {
self.stop.store(true, std::sync::atomic::Ordering::Relaxed);
}
pub fn stopped(&self) -> bool {
self.stop.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn rows_seen(&self) -> Option<usize> {
self.counted
.load(std::sync::atomic::Ordering::Relaxed)
.then(|| self.rows.load(std::sync::atomic::Ordering::Relaxed))
}
pub(crate) fn saw(&self, rows: usize) {
self.counted
.store(true, std::sync::atomic::Ordering::Relaxed);
self.rows
.fetch_add(rows, std::sync::atomic::Ordering::Relaxed);
}
pub(crate) fn restart(&self) -> Option<usize> {
let counted = self
.counted
.swap(false, std::sync::atomic::Ordering::Relaxed);
let rows = self.rows.swap(0, std::sync::atomic::Ordering::Relaxed);
counted.then_some(rows)
}
pub(crate) fn check(&self) -> Result<()> {
if self.stopped() && self.memory_stopped().is_none() {
Err(Report::msg(CANCELLED))
} else {
Ok(())
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub enum SampleMethod {
#[default]
Spread,
PerPartition { column: String },
FirstRows,
EveryRow,
}
impl SampleMethod {
pub fn label(&self) -> String {
match self {
Self::Spread => "Random".to_string(),
Self::PerPartition { column } => format!("Equal per {column}"),
Self::FirstRows => "First rows".to_string(),
Self::EveryRow => "Every row".to_string(),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Sample {
pub scope: QualityScope,
pub method: SampleMethod,
pub rows: usize,
pub seed: u64,
}
impl Default for Sample {
fn default() -> Self {
Self {
scope: QualityScope::CurrentView,
method: SampleMethod::Spread,
rows: DEFAULT_SAMPLE_ROWS,
seed: 42_891,
}
}
}
impl Sample {
pub fn summary(&self) -> String {
let rows = numfmt::group_chrome(self.rows);
let how = match &self.method {
SampleMethod::Spread => format!("{rows} random rows"),
SampleMethod::PerPartition { column } => format!("{rows} rows per {column}"),
SampleMethod::FirstRows => format!("first {rows} rows"),
SampleMethod::EveryRow => "every row".to_string(),
};
let seeded = matches!(
self.method,
SampleMethod::Spread | SampleMethod::PerPartition { .. }
);
let middot = crate::glyphs::get().middot;
if seeded {
format!(
"{how} {middot} {} {middot} seed {}",
self.scope.label(),
self.seed
)
} else {
format!("{how} {middot} {}", self.scope.label())
}
}
pub fn summary_within(&self, rows: Option<usize>) -> String {
match (rows, &self.method) {
(Some(n), SampleMethod::Spread | SampleMethod::FirstRows) if n <= self.rows => {
let middot = crate::glyphs::get().middot;
format!(
"all {} rows {middot} {}",
numfmt::group_chrome(n),
self.scope.label()
)
}
_ => self.summary(),
}
}
pub fn outcome(
&self,
total_rows: usize,
sample_size: Option<usize>,
per_value: Option<usize>,
) -> String {
let count = numfmt::group_chrome;
let read = match (&self.method, sample_size) {
(SampleMethod::FirstRows, Some(n)) => format!("first {} rows", count(n)),
(SampleMethod::PerPartition { column }, Some(n)) => {
let each = per_value.unwrap_or(self.rows).min(self.rows);
let lowered = if each < self.rows {
format!(" (lowered from {})", count(self.rows))
} else {
String::new()
};
format!(
"{} rows, up to {} per {column}{lowered}, of {}",
count(n),
count(each),
count(total_rows)
)
}
(_, Some(n)) => format!("sample of {} of {} rows", count(n), count(total_rows)),
(_, None) => format!("all {} rows", count(total_rows)),
};
if self.scope == QualityScope::CurrentView {
read
} else {
format!(
"{read} {} {}",
crate::glyphs::get().middot,
self.scope.label()
)
}
}
}
pub struct SampleSource {
lf: LazyFrame,
source: Option<QualitySourceContext>,
from_source: bool,
}
impl SampleSource {
pub fn view(lf: LazyFrame) -> Self {
Self {
lf,
source: None,
from_source: false,
}
}
pub fn loaded(lf: LazyFrame, source: Option<QualitySourceContext>) -> Self {
Self {
lf,
source,
from_source: true,
}
}
pub fn cut(self, scope: &QualityScope) -> Result<LazyFrame> {
let lf = if self.from_source {
prepare_source_quality_scan(self.lf, self.source.as_ref())?
} else {
self.lf
};
let lf = apply_quality_scope(lf, scope, self.source.as_ref())?;
let schema = lf.clone().collect_schema()?;
let helpers = [
crate::schema_union::DRIFT_COLUMN,
"__datui_quality_row",
self.source
.as_ref()
.map(|source| source.row_index_column.as_str())
.unwrap_or(""),
];
let keep = schema
.iter_names()
.filter(|name| !helpers.contains(&name.as_str()))
.map(|name| col(name.clone()))
.collect::<Vec<_>>();
Ok(if keep.len() == schema.len() {
lf
} else {
lf.select(keep)
})
}
}
pub fn view_scope_rows(view_rows: Option<usize>, scope: &QualityScope) -> Option<usize> {
let rows = view_rows?;
match scope {
QualityScope::CurrentView => Some(rows),
QualityScope::FirstRows(limit) => Some(rows.min(*limit)),
QualityScope::ViewRows { start, end } => {
Some(rows.min(*end).saturating_sub(start.saturating_sub(1)))
}
_ => None,
}
}
pub fn read(
lf: &LazyFrame,
sample: &Sample,
known_total: Option<usize>,
polars_streaming: bool,
) -> Result<AnalysisRows> {
let rows = read_rows(lf, sample, known_total, polars_streaming)?;
if rows.total_rows == 0 {
return Err(no_rows_error(&sample.scope));
}
Ok(rows)
}
pub fn no_rows_error(scope: &QualityScope) -> Report {
if *scope == QualityScope::CurrentView {
Report::msg("The table has no rows to sample")
} else {
Report::msg(format!(
"No rows match {}; change Rows from in the Sample form (s)",
scope.label()
))
}
}
pub(crate) fn read_rows(
lf: &LazyFrame,
sample: &Sample,
known_total: Option<usize>,
polars_streaming: bool,
) -> Result<AnalysisRows> {
read_rows_watched(lf, sample, known_total, polars_streaming, None)
}
pub(crate) fn read_rows_watched(
lf: &LazyFrame,
sample: &Sample,
known_total: Option<usize>,
polars_streaming: bool,
watch: Option<&ReadWatch>,
) -> Result<AnalysisRows> {
acquire(lf, sample, known_total, polars_streaming, watch, None).map(|read| read.rows)
}
pub(crate) struct SampledRows {
pub rows: AnalysisRows,
pub positions: Vec<IdxSize>,
pub counted: Option<Counted>,
}
pub(crate) fn acquire(
lf: &LazyFrame,
sample: &Sample,
known_total: Option<usize>,
polars_streaming: bool,
watch: Option<&ReadWatch>,
count: Option<&Expr>,
) -> Result<SampledRows> {
let n = sample.rows.max(1);
match &sample.method {
SampleMethod::EveryRow => crate::statistics::sample_rows_counting(
lf,
None,
known_total,
sample.seed,
polars_streaming,
watch,
count,
),
SampleMethod::Spread => crate::statistics::sample_rows_counting(
lf,
Some(n),
known_total,
sample.seed,
polars_streaming,
watch,
count,
),
SampleMethod::FirstRows => {
let df = collect_lazy(lf.clone().limit(n as IdxSize), polars_streaming)
.map_err(Report::from)?;
let height = df.height();
let sampled = match known_total {
Some(total) => total > height,
None => height == n,
};
Ok(SampledRows {
positions: (0..height as IdxSize).collect(),
rows: AnalysisRows {
df,
total_rows: known_total.unwrap_or(height),
sample_size: sampled.then_some(height),
per_value: None,
},
counted: None,
})
}
SampleMethod::PerPartition { column } => {
let read =
per_group_sample_within(lf, column, n, sample.seed, MAX_GROUP_ROWS, watch, count)?;
let sample_size = (read.seen > read.df.height()).then_some(read.df.height());
Ok(SampledRows {
rows: AnalysisRows {
df: read.df,
total_rows: read.seen,
sample_size,
per_value: Some(read.per_value),
},
positions: read.positions,
counted: read.counted,
})
}
}
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct PerValue {
pub kept: usize,
pub totals: std::collections::BTreeMap<Option<String>, usize>,
}
struct GroupRead {
df: DataFrame,
seen: usize,
per_value: PerValue,
positions: Vec<IdxSize>,
counted: Option<Counted>,
}
fn per_group_sample_within(
lf: &LazyFrame,
column: &str,
n: usize,
seed: u64,
limit: usize,
watch: Option<&ReadWatch>,
count: Option<&Expr>,
) -> Result<GroupRead> {
let schema = lf.clone().collect_schema()?;
if schema.get(column).is_none() {
return Err(Report::msg(format!(
"partition column {column:?} is not in the rows sampled; choose another"
)));
}
let state = std::sync::Arc::new(std::sync::Mutex::new(GroupState {
column: column.to_string(),
cap: n,
limit,
seed,
..Default::default()
}));
let callback_state = std::sync::Arc::clone(&state);
let callback_watch = watch.cloned();
let sink = with_count_key(lf.clone(), count)
.with_row_index(GROUP_POSITION, None)
.sink_batches(
PlanCallback::new(move |batch: DataFrame| {
if let Some(watch) = &callback_watch {
if watch.stopped() {
return Ok(true);
}
watch.saw(batch.height());
}
let mut state = callback_state
.lock()
.map_err(|_| PolarsError::ComputeError("sampler lock failed".into()))?;
state.observe(batch)?;
if let Some(watch) = &callback_watch {
watch.hold(state.bytes(), state.held);
}
Ok(false)
}),
true,
None,
)?;
collect_lazy(sink, true).map_err(Report::from)?;
if let Some(watch) = watch {
watch.check()?;
}
let state = std::mem::take(
&mut *state
.lock()
.map_err(|_| Report::msg("sampler lock failed"))?,
);
let seen = state.seen;
let cap = state
.cap
.min((state.limit / state.groups.len().max(1)).max(1));
let mut totals = std::collections::BTreeMap::new();
let mut out: Option<DataFrame> = None;
for (key, mut group) in state.groups {
totals.insert(key, group.total);
group.trim(cap)?;
let Some(rows) = group.rows else {
continue;
};
out = Some(match out {
Some(frame) => frame.vstack(&rows)?,
None => rows,
});
}
let (df, positions) = match out {
Some(df) => {
let df = df.sort([GROUP_POSITION], SortMultipleOptions::default())?;
let positions = df
.column(GROUP_POSITION)?
.idx()?
.into_no_null_iter()
.collect();
(df.drop(GROUP_POSITION)?, positions)
}
None => (
collect_lazy(lf.clone().limit(0), true).map_err(Report::from)?,
Vec::new(),
),
};
Ok(GroupRead {
df,
seen,
per_value: PerValue { kept: cap, totals },
positions,
counted: count.is_some().then(|| state.counter.finish()),
})
}
#[derive(Default)]
struct GroupState {
column: String,
cap: usize,
limit: usize,
seed: u64,
seen: usize,
held: usize,
groups: HashMap<Option<String>, GroupSample>,
counter: KeyCounter,
}
#[derive(Default)]
struct GroupSample {
rows: Option<DataFrame>,
ranks: Vec<u64>,
total: usize,
}
impl GroupSample {
fn trim(&mut self, cap: usize) -> PolarsResult<usize> {
if self.ranks.len() <= cap {
return Ok(0);
}
let mut order: Vec<usize> = (0..self.ranks.len()).collect();
order.sort_unstable_by_key(|i| self.ranks[*i]);
order.truncate(cap);
let take: Vec<IdxSize> = order.iter().map(|i| *i as IdxSize).collect();
if let Some(kept) = self.rows.take() {
self.rows = Some(kept.take(&IdxCa::from_vec("kept".into(), take))?);
}
let removed = self.ranks.len() - cap;
self.ranks = order.iter().map(|i| self.ranks[*i]).collect();
Ok(removed)
}
}
impl GroupState {
fn bytes(&self) -> u64 {
self.groups
.values()
.filter_map(|group| group.rows.as_ref())
.map(|rows| rows.estimated_size() as u64)
.sum()
}
fn observe(&mut self, mut batch: DataFrame) -> PolarsResult<()> {
self.counter.observe(&mut batch)?;
self.seen += batch.height();
let positions = batch.column(GROUP_POSITION)?.idx()?.clone();
let keys = batch.column(&self.column)?.as_materialized_series().clone();
let mut by_key: HashMap<Option<String>, (Vec<IdxSize>, Vec<u64>)> = HashMap::new();
for (index, (key, position)) in keys.iter().zip(positions.into_no_null_iter()).enumerate() {
let key = (!key.is_null()).then(|| crate::exact::str_value(&key).into_owned());
let entry = by_key.entry(key).or_default();
entry.0.push(index as IdxSize);
entry.1.push(sample_rank(self.seed, position as u64));
}
for (key, (indices, ranks)) in by_key {
if !self.groups.contains_key(&key) && self.groups.len() >= MAX_GROUPS {
return Err(PolarsError::ComputeError(
format!(
"more than {MAX_GROUPS} values of {}; sample per a coarser column",
self.column
)
.into(),
));
}
let group = self.groups.entry(key).or_default();
group.total += indices.len();
self.held += indices.len();
let rows = batch.take(&IdxCa::from_vec("picked".into(), indices))?;
group.rows = Some(match group.rows.take() {
Some(kept) => kept.vstack(&rows)?,
None => rows,
});
group.ranks.extend(ranks);
self.held -= group.trim(self.cap)?;
}
if self.held > self.limit.saturating_add(self.limit / 4) {
self.cap = self.cap.min((self.limit / self.groups.len()).max(1));
for group in self.groups.values_mut() {
self.held -= group.trim(self.cap)?;
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
#[test]
fn a_size_takes_shorthand() {
use super::parse_size;
for (text, rows) in [
("50000", 50_000),
("50,000", 50_000),
("1_000", 1_000),
("50k", 50_000),
("250K", 250_000),
("2m", 2_000_000),
("2.5M", 2_500_000),
(" 7 ", 7),
("99999999999999999999999", usize::MAX),
] {
assert_eq!(parse_size(text), Ok(rows), "{text}");
}
for bad in ["", "k", "12x", "1.5", "-3", "1e6", "2mm"] {
assert_eq!(parse_size(bad), Err(super::SizeError::NotASize), "{bad}");
}
for zero in ["0", "0k", "0.0001k"] {
assert_eq!(parse_size(zero), Err(super::SizeError::Zero), "{zero}");
}
}
use super::*;
fn table() -> LazyFrame {
let sizes = [("a", 9_000usize), ("b", 900), ("c", 100)];
let mut part = Vec::new();
let mut value = Vec::new();
for (name, size) in sizes {
for row in 0..size {
part.push(name);
value.push(row as i64);
}
}
df!("part" => part, "value" => value).unwrap().lazy()
}
fn sample(method: SampleMethod, rows: usize) -> Sample {
Sample {
method,
rows,
..Sample::default()
}
}
#[test]
fn per_partition_keeps_up_to_n_from_each_value() {
let method = SampleMethod::PerPartition {
column: "part".to_string(),
};
let rows = read(&table(), &sample(method, 200), None, false).unwrap();
assert_eq!(rows.total_rows, 10_000);
let counts = rows
.df
.column("part")
.unwrap()
.as_materialized_series()
.value_counts(true, true, "n".into(), false)
.unwrap();
let n = |part: &str| {
(0..counts.height())
.find(|row| {
counts
.column("part")
.unwrap()
.get(*row)
.unwrap()
.str_value()
== part
})
.map(|row| {
counts
.column("n")
.unwrap()
.get(row)
.unwrap()
.try_extract::<u32>()
.unwrap()
})
.unwrap()
};
assert_eq!((n("a"), n("b"), n("c")), (200, 200, 100));
assert_eq!(rows.sample_size, Some(500));
assert!(
rows.df.column(GROUP_POSITION).is_err(),
"no helper column leaks"
);
}
#[test]
fn per_partition_past_the_limit_keeps_fewer_of_each_value() {
let GroupRead {
df,
seen,
per_value,
..
} = per_group_sample_within(&table(), "part", 500, 42_891, 999, None, None).unwrap();
assert_eq!(seen, 10_000);
assert_eq!(per_value.kept, 333);
assert_eq!(df.height(), 333 + 333 + 100);
assert_eq!(
per_value.totals,
[("a", 9_000), ("b", 900), ("c", 100)]
.into_iter()
.map(|(part, rows)| (Some(part.to_string()), rows))
.collect()
);
let asked =
per_group_sample_within(&table(), "part", 333, 42_891, usize::MAX, None, None).unwrap();
assert!(df.equals(&asked.df), "the rows a sample of 333 each keeps");
let lowered = Sample {
method: SampleMethod::PerPartition {
column: "part".into(),
},
rows: 500,
..Sample::default()
};
assert_eq!(
lowered.outcome(10_000, Some(766), Some(333)),
"766 rows, up to 333 per part (lowered from 500), of 10,000"
);
}
#[test]
fn a_sample_says_where_its_rows_sat_and_a_stream_counts_on_the_way() {
let lf = table()
.with_column((col("value") % lit(4)).alias("quarter"))
.with_row_index("row", None);
let totals: std::collections::BTreeMap<_, _> = [("a", 9_000), ("b", 900), ("c", 100)]
.into_iter()
.map(|(part, rows)| (Some(part.to_string()), rows))
.collect();
for method in [
SampleMethod::Spread,
SampleMethod::PerPartition {
column: "quarter".to_string(),
},
SampleMethod::FirstRows,
] {
let read = acquire(
&lf,
&sample(method.clone(), 50),
None,
false,
None,
Some(&col("part")),
)
.unwrap();
let rows: Vec<IdxSize> = read
.rows
.df
.column("row")
.unwrap()
.idx()
.unwrap()
.into_no_null_iter()
.collect();
assert_eq!(rows, read.positions, "{method:?}");
assert!(read.rows.df.column(COUNT_KEY).is_err(), "{method:?}");
if method == SampleMethod::FirstRows {
assert_eq!(read.counted, None);
} else {
assert_eq!(
read.counted,
Some(Counted::Totals(totals.clone())),
"{method:?}"
);
}
}
}
#[test]
fn a_count_past_its_limit_gives_up_and_says_so() {
let mut counter = KeyCounter::with_limit(2);
let mut batch = df!(COUNT_KEY => ["a", "b", "c"], "value" => [1, 2, 3]).unwrap();
counter.observe(&mut batch).unwrap();
assert_eq!(batch.get_column_names(), ["value"]);
assert_eq!(counter.finish(), Counted::TooMany);
}
#[test]
fn first_rows_is_the_head_and_every_row_is_all_of_it() {
let head = read(
&table(),
&sample(SampleMethod::FirstRows, 50),
Some(10_000),
false,
)
.unwrap();
assert_eq!(head.df.height(), 50);
assert_eq!(head.sample_size, Some(50));
assert_eq!(
head.df.column("value").unwrap().i64().unwrap().get(49),
Some(49)
);
let all = read(&table(), &sample(SampleMethod::EveryRow, 50), None, false).unwrap();
assert_eq!((all.df.height(), all.sample_size), (10_000, None));
}
#[test]
fn a_seeded_per_partition_sample_repeats() {
let method = SampleMethod::PerPartition {
column: "part".to_string(),
};
let one = read(&table(), &sample(method.clone(), 50), None, false).unwrap();
let two = read(&table(), &sample(method.clone(), 50), None, false).unwrap();
assert!(one.df.equals(&two.df));
let other = read(
&table(),
&Sample {
seed: 7,
..sample(method, 50)
},
None,
false,
)
.unwrap();
assert!(!one.df.equals(&other.df));
}
#[test]
fn a_scope_that_matches_nothing_is_an_error() {
let scope = QualityScope::parse_command("partition part=zzz").unwrap();
let lf = SampleSource::view(table()).cut(&scope).unwrap();
let Err(error) = read(
&lf,
&Sample {
scope,
..Sample::default()
},
None,
false,
) else {
panic!("a scope that matches nothing must not sample");
};
assert!(error.to_string().contains("No rows match"), "{error}");
}
#[test]
fn the_summary_says_what_will_be_read() {
let middot = crate::glyphs::get().middot;
assert_eq!(
Sample::default().summary(),
format!("100,000 random rows {middot} current view {middot} seed 42891")
);
assert_eq!(
sample(SampleMethod::FirstRows, 1_000).summary(),
format!("first 1,000 rows {middot} current view")
);
assert_eq!(
Sample::default().summary_within(Some(1_000)),
format!("all 1,000 rows {middot} current view")
);
assert_eq!(
Sample::default().summary_within(Some(1_000_000)),
Sample::default().summary()
);
assert_eq!(
Sample::default().summary_within(None),
Sample::default().summary()
);
}
}