Skip to main content

datui_lib/
statistics.rs

1use color_eyre::Result;
2use color_eyre::eyre::Report;
3use polars::polars_compute::rolling::QuantileMethod;
4use polars::prelude::*;
5use std::collections::HashMap;
6use std::ops::Range;
7
8/// Collects a LazyFrame into a DataFrame.
9///
10/// When the `streaming` feature is enabled and `use_streaming` is true, uses the Polars
11/// streaming engine (batch processing, lower memory). Otherwise collects normally.
12/// Returns `PolarsError` so callers that need to display or store the error (e.g.
13/// `DataTableState::error`) can do so without converting.
14pub fn collect_lazy(
15    lf: LazyFrame,
16    use_streaming: bool,
17) -> std::result::Result<DataFrame, PolarsError> {
18    #[cfg(feature = "streaming")]
19    {
20        if may_stream(&lf, use_streaming) && !sorts_by_one_wide_key(&lf) {
21            // A plain collect is always one frame; `Multiple` only comes from sink_multiple.
22            lf.collect_with_engine(Engine::Streaming)
23                .map(|result| result.unwrap_single())
24        } else {
25            lf.collect()
26        }
27    }
28    #[cfg(not(feature = "streaming"))]
29    {
30        let _ = use_streaming; // ignored when streaming feature is disabled
31        lf.collect()
32    }
33}
34
35/// Whether a query over `lf` may use the streaming engine: asked for, and possible.
36/// Polars 0.55's streaming engine cannot run an anonymous scan (a SQLite table): it
37/// stops at a `todo!`.
38pub fn may_stream(lf: &LazyFrame, wanted: bool) -> bool {
39    use polars::lazy::dsl::{DslPlan, FileScanDsl};
40    wanted
41        && !lf.logical_plan.into_iter().any(|node| match node {
42            DslPlan::Scan { scan_type, .. } => {
43                matches!(scan_type.as_ref(), FileScanDsl::Anonymous { .. })
44            }
45            _ => false,
46        })
47}
48
49/// Whether `lf` takes the first rows of an unstable sort by a single Decimal or Int128
50/// key. Polars 0.55's streaming engine runs that as a top-k, which panics on those
51/// dtypes ("not implemented for dtype Int128"); the in-memory engine sorts them. A
52/// format spec's `scale` reads as Decimal, so a query's `by price`, whose group sort
53/// is unstable, takes this for its first page. A stable sort (the table's own) carries
54/// a row index as a second key, and a full sort or a slice further in is no top-k:
55/// both stream.
56#[cfg(feature = "streaming")]
57fn sorts_by_one_wide_key(lf: &LazyFrame) -> bool {
58    use polars::lazy::dsl::DslPlan;
59    let wide_sort = |node: &DslPlan| match node {
60        DslPlan::Sort {
61            input,
62            by_column,
63            sort_options,
64            ..
65        } if by_column.len() == 1 && !sort_options.maintain_order => {
66            LazyFrame::from((**input).clone())
67                .select([by_column[0].clone()])
68                .collect_schema()
69                .ok()
70                .and_then(|schema| schema.get_at_index(0).map(|(_, dtype)| dtype.clone()))
71                .is_some_and(|dtype| dtype.is_decimal() || dtype == DataType::Int128)
72        }
73        _ => false,
74    };
75    lf.logical_plan.into_iter().any(|node| match node {
76        DslPlan::Slice {
77            input, offset: 0, ..
78        } => input.into_iter().any(wide_sort),
79        DslPlan::Sort { sort_options, .. } => sort_options.limit.is_some() && wide_sort(node),
80        _ => false,
81    })
82}
83
84/// Default sampling threshold: datasets >= this size are sampled.
85/// Used as fallback when sample_size is None. App uses config value.
86pub const SAMPLING_THRESHOLD: usize = 10_000;
87
88#[derive(Clone)]
89pub struct ColumnStatistics {
90    pub name: String,
91    pub dtype: DataType,
92    pub count: usize,
93    pub null_count: usize,
94    pub numeric_stats: Option<NumericStatistics>,
95    pub categorical_stats: Option<CategoricalStatistics>,
96    pub temporal_stats: Option<TemporalStatistics>,
97    pub distribution_info: Option<DistributionInfo>,
98}
99
100#[derive(Clone)]
101pub struct NumericStatistics {
102    pub mean: f64,
103    pub std: f64,
104    pub min: f64,
105    pub max: f64,
106    pub median: f64,
107    pub q25: f64,
108    pub q75: f64,
109    pub percentiles: HashMap<u8, f64>, // 1, 5, 25, 50, 75, 95, 99
110    pub skewness: f64,
111    pub kurtosis: f64,
112    pub outliers_iqr: usize,
113    pub outliers_zscore: usize,
114}
115
116#[derive(Clone)]
117pub struct CategoricalStatistics {
118    pub unique_count: usize,
119    pub mode: Option<String>,
120    pub top_values: Vec<(String, usize)>,
121    pub min: Option<String>, // Lexicographically smallest string
122    pub max: Option<String>, // Lexicographically largest string
123}
124
125/// Describe for a Date, Datetime, Time or Duration column: each statistic a value of the
126/// column's own type, written as the table writes it; `None` for a null.
127#[derive(Clone, Default)]
128pub struct TemporalStatistics {
129    pub mean: Option<String>,
130    pub min: Option<String>,
131    pub q25: Option<String>,
132    pub median: Option<String>,
133    pub q75: Option<String>,
134    pub max: Option<String>,
135}
136
137#[derive(Clone)]
138pub struct DistributionInfo {
139    pub distribution_type: DistributionType,
140    pub confidence: f64,
141    pub sample_size: usize,
142    pub is_sampled: bool,
143    pub fit_quality: Option<f64>, // 0.0-1.0, how well data fits detected type
144    /// Every family's fit and test, or why it does not apply.
145    pub fits: Vec<(DistributionType, crate::distribution_fit::FitOutcome)>,
146}
147
148#[derive(Clone)]
149pub struct DistributionAnalysis {
150    pub column_name: String,
151    pub distribution_type: DistributionType,
152    pub confidence: f64,  // 0.0-1.0
153    pub fit_quality: f64, // 0.0-1.0, how well data fits detected type
154    pub characteristics: DistributionCharacteristics,
155    pub outliers: OutlierAnalysis,
156    pub percentiles: PercentileBreakdown,
157    pub sorted_sample_values: Vec<f64>, // Sorted data values for Q-Q plot (all data if < threshold, sampled if >= threshold)
158    pub is_sampled: bool,               // Whether data was sampled
159    pub sample_size: usize,             // Actual number of values used
160    /// Every family's fit and test, or why it does not apply.
161    pub fits: Vec<(DistributionType, crate::distribution_fit::FitOutcome)>,
162    /// Each fitted family's quantiles at the plotting positions of
163    /// `sorted_sample_values`, for its Q-Q plot: computed with the fit, not per frame.
164    pub qq: Vec<(DistributionType, Vec<f64>)>,
165}
166
167impl DistributionAnalysis {
168    pub fn fit(&self, family: DistributionType) -> Option<&crate::distribution_fit::FitOutcome> {
169        self.fits
170            .iter()
171            .find(|(fitted, _)| *fitted == family)
172            .map(|(_, outcome)| outcome)
173    }
174
175    pub fn qq(&self, family: DistributionType) -> Option<&[f64]> {
176        self.qq
177            .iter()
178            .find(|(fitted, _)| *fitted == family)
179            .map(|(_, quantiles)| quantiles.as_slice())
180    }
181}
182
183#[derive(Clone)]
184pub struct DistributionCharacteristics {
185    pub shapiro_wilk_stat: Option<f64>,
186    pub shapiro_wilk_pvalue: Option<f64>,
187    pub skewness: f64,
188    pub kurtosis: f64,
189    pub mean: f64,
190    pub median: f64,
191    pub std_dev: f64,
192    pub variance: f64,
193    pub coefficient_of_variation: f64,
194    pub mode: Option<f64>, // For unimodal distributions
195}
196
197#[derive(Clone)]
198pub struct OutlierAnalysis {
199    pub total_count: usize,
200    pub percentage: f64,
201    pub iqr_count: usize,
202    pub zscore_count: usize,
203    pub outlier_rows: Vec<OutlierRow>, // Limited to top N for performance
204}
205
206#[derive(Clone)]
207pub struct OutlierRow {
208    pub row_index: usize,
209    pub column_value: f64,
210    pub context_data: HashMap<String, String>, // Other column values for context
211    pub detection_method: OutlierMethod,
212    pub z_score: Option<f64>,
213    pub iqr_position: Option<IqrPosition>, // Below Q1-1.5*IQR or above Q3+1.5*IQR
214}
215
216#[derive(Clone, Debug)]
217pub enum OutlierMethod {
218    IQR,
219    ZScore,
220    Both,
221}
222
223#[derive(Clone, Debug)]
224pub enum IqrPosition {
225    BelowLowerFence,
226    AboveUpperFence,
227}
228
229#[derive(Clone)]
230pub struct PercentileBreakdown {
231    pub p1: f64,
232    pub p5: f64,
233    pub p25: f64,
234    pub p50: f64,
235    pub p75: f64,
236    pub p95: f64,
237    pub p99: f64,
238}
239
240/// Which coefficient the correlation matrix shows.
241#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
242pub enum CorrelationMethod {
243    /// Pearson's r: how close the pairs fall to a line.
244    #[default]
245    Pearson,
246    /// Spearman's ρ: Pearson's r of the pairs' ranks, for any monotone relation.
247    Spearman,
248}
249
250impl CorrelationMethod {
251    pub fn toggled(self) -> Self {
252        match self {
253            Self::Pearson => Self::Spearman,
254            Self::Spearman => Self::Pearson,
255        }
256    }
257}
258
259// Correlation matrix structures
260#[derive(Clone)]
261pub struct CorrelationMatrix {
262    pub columns: Vec<String>,            // Numeric column names
263    pub correlations: Vec<Vec<f64>>,     // Square matrix of Pearson correlations
264    pub p_values: Option<Vec<Vec<f64>>>, // Statistical significance (optional)
265    pub sample_sizes: Vec<Vec<usize>>,   // Sample size for each pair
266    /// Spearman's ρ for each pair, over the same pairs as Pearson's r; `None` when
267    /// the rows read hold more values than [`RANK_VALUES`].
268    pub rank_correlations: Option<Vec<Vec<f64>>>,
269    pub rank_p_values: Option<Vec<Vec<f64>>>,
270}
271
272/// The most values Spearman's ρ ranks: eight bytes each, so 512 MiB beside the rows
273/// read. A sample's 100,000 rows rank up to 671 columns; a read of every row of a
274/// large table can be past it, and the matrix then has Pearson's r only.
275pub const RANK_VALUES: usize = 64 * 1024 * 1024;
276
277impl CorrelationMatrix {
278    /// The pair's coefficient by `method`; NaN where there is none.
279    pub fn coefficient(&self, method: CorrelationMethod, row: usize, col: usize) -> f64 {
280        let matrix = match method {
281            CorrelationMethod::Pearson => Some(&self.correlations),
282            CorrelationMethod::Spearman => self.rank_correlations.as_ref(),
283        };
284        matrix
285            .and_then(|m| m.get(row))
286            .and_then(|r| r.get(col))
287            .copied()
288            .unwrap_or(f64::NAN)
289    }
290
291    /// The pair's p-value by `method`, when the matrix has them.
292    pub fn p_value(&self, method: CorrelationMethod, row: usize, col: usize) -> Option<f64> {
293        let matrix = match method {
294            CorrelationMethod::Pearson => self.p_values.as_ref(),
295            CorrelationMethod::Spearman => self.rank_p_values.as_ref(),
296        };
297        matrix.and_then(|m| m.get(row)?.get(col).copied())
298    }
299}
300
301#[derive(Clone)]
302pub struct CorrelationPair {
303    pub column1: String,
304    pub column2: String,
305    pub correlation: f64,
306    pub p_value: Option<f64>,
307    pub sample_size: usize,
308    pub covariance: f64,
309    pub r_squared: f64,
310    pub stats1: ColumnStats,
311    pub stats2: ColumnStats,
312}
313
314#[derive(Clone)]
315pub struct ColumnStats {
316    pub mean: f64,
317    pub std: f64,
318    pub min: f64,
319    pub max: f64,
320}
321
322#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash)]
323pub enum DistributionType {
324    #[default]
325    Normal,
326    LogNormal,
327    Uniform,
328    PowerLaw,
329    Exponential,
330    Beta,
331    Gamma,
332    ChiSquared,
333    StudentsT,
334    Poisson,
335    Bernoulli,
336    Binomial,
337    Geometric,
338    Weibull,
339    /// One value throughout: nothing to fit.
340    Constant,
341    /// Every candidate was rejected: shown as "No clear fit".
342    Unknown,
343}
344
345impl std::fmt::Display for DistributionType {
346    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
347        match self {
348            DistributionType::Normal => write!(f, "Normal"),
349            DistributionType::LogNormal => write!(f, "Log-Normal"),
350            DistributionType::Uniform => write!(f, "Uniform"),
351            DistributionType::PowerLaw => write!(f, "Power Law"),
352            DistributionType::Exponential => write!(f, "Exponential"),
353            DistributionType::Beta => write!(f, "Beta"),
354            DistributionType::Gamma => write!(f, "Gamma"),
355            DistributionType::ChiSquared => write!(f, "Chi-Squared"),
356            DistributionType::StudentsT => write!(f, "Student's t"),
357            DistributionType::Poisson => write!(f, "Poisson"),
358            DistributionType::Bernoulli => write!(f, "Bernoulli"),
359            DistributionType::Binomial => write!(f, "Binomial"),
360            DistributionType::Geometric => write!(f, "Geometric"),
361            DistributionType::Weibull => write!(f, "Weibull"),
362            DistributionType::Constant => write!(f, "Constant"),
363            DistributionType::Unknown => write!(f, "No clear fit"),
364        }
365    }
366}
367
368#[derive(Clone)]
369pub struct AnalysisResults {
370    pub column_statistics: Vec<ColumnStatistics>,
371    pub total_rows: usize,
372    pub sample_size: Option<usize>,
373    /// Rows an equal-per-value sample kept of each value. See [`crate::sampling::PerValue`].
374    pub per_value: Option<usize>,
375    pub sample_seed: u64,
376    pub correlation_matrix: Option<CorrelationMatrix>,
377    pub distribution_analyses: Vec<DistributionAnalysis>,
378}
379
380pub struct AnalysisContext {
381    pub has_query: bool,
382    pub query: String,
383    pub has_filters: bool,
384    pub filter_count: usize,
385    pub is_drilled_down: bool,
386    pub group_key: Option<Vec<String>>,
387    pub group_columns: Option<Vec<String>>,
388}
389
390#[derive(Debug, Clone, Copy)]
391pub struct ComputeOptions {
392    pub include_distribution_info: bool,
393    pub include_distribution_analyses: bool,
394    pub include_correlation_matrix: bool,
395    pub include_skewness_kurtosis_outliers: bool,
396    /// When true, use Polars streaming engine for LazyFrame collect when the streaming feature is enabled.
397    pub polars_streaming: bool,
398}
399
400impl Default for ComputeOptions {
401    fn default() -> Self {
402        Self {
403            include_distribution_info: false,
404            include_distribution_analyses: false,
405            include_correlation_matrix: false,
406            include_skewness_kurtosis_outliers: false,
407            polars_streaming: true,
408        }
409    }
410}
411
412/// Computes statistics for a LazyFrame with default options.
413///
414/// Convenience wrapper around `compute_statistics_with_options`.
415pub fn compute_statistics(
416    lf: &LazyFrame,
417    sample_size: Option<usize>,
418    seed: u64,
419) -> Result<AnalysisResults> {
420    compute_statistics_with_options(lf, sample_size, seed, ComputeOptions::default())
421}
422
423/// Computes comprehensive statistics for a LazyFrame.
424///
425/// Main entry point for statistical analysis. Computes:
426/// - Basic statistics (count, nulls, min, max, mean) for all columns
427/// - Numeric statistics (percentiles, skewness, kurtosis, outliers) for numeric columns
428/// - Categorical statistics (unique count, mode, top values) for categorical columns
429/// - Distribution detection and analysis for numeric columns (if enabled)
430/// - Correlation matrix for numeric columns (if enabled)
431///
432/// A table with more than `sample_size` rows is analyzed from a sample of that many;
433/// see [`analysis_rows`]. `None` reads every row.
434pub fn compute_statistics_with_options(
435    lf: &LazyFrame,
436    sample_size: Option<usize>,
437    seed: u64,
438    options: ComputeOptions,
439) -> Result<AnalysisResults> {
440    let sample = crate::sampling::Sample {
441        method: if sample_size.is_some() {
442            crate::sampling::SampleMethod::Spread
443        } else {
444            crate::sampling::SampleMethod::EveryRow
445        },
446        rows: sample_size.unwrap_or(0),
447        seed,
448        ..crate::sampling::Sample::default()
449    };
450    compute_statistics_for_sample(lf, &sample, None, options)
451}
452
453/// Describe's temporal statistics for one collected column, through the same
454/// aggregation the lazy path runs.
455fn temporal_stats_of(series: &Series) -> Result<Option<TemporalStatistics>> {
456    if !is_temporal_type(series.dtype()) {
457        return Ok(None);
458    }
459    let frame = DataFrame::new_infer_height(vec![series.clone().into()])?;
460    let schema = frame.schema().clone();
461    let agg_df = frame
462        .lazy()
463        .select(build_describe_aggregation_exprs(&schema))
464        .collect()?;
465    Ok(parse_describe_agg_row(&agg_df, &schema)
466        .pop()
467        .and_then(|stats| stats.temporal_stats))
468}
469
470/// A value as the table writes it; `None` for a null.
471fn get_value_str(df: &DataFrame, col_name: &str, row: usize) -> Option<String> {
472    match df.column(col_name).ok()?.get(row).ok()? {
473        AnyValue::Null => None,
474        v => Some(crate::exact::str_value(&v).to_string()),
475    }
476}
477
478/// [`compute_statistics_with_options`] over the rows a [`crate::sampling::Sample`]
479/// picks from `lf`, which is already cut to the sample's scope.
480pub fn compute_statistics_for_sample(
481    lf: &LazyFrame,
482    sample: &crate::sampling::Sample,
483    known_total: Option<usize>,
484    options: ComputeOptions,
485) -> Result<AnalysisResults> {
486    let schema = lf.clone().collect_schema()?;
487    let use_streaming = options.polars_streaming;
488    let seed = sample.seed;
489    let rows = crate::sampling::read(lf, sample, known_total, use_streaming)?;
490    let total_rows = rows.total_rows;
491    let actual_sample_size = rows.sample_size;
492    let should_sample = actual_sample_size.is_some();
493    let per_value = rows.per_value.as_ref().map(|per_value| per_value.kept);
494    let df = rows.df;
495
496    let mut column_statistics = Vec::new();
497
498    for (name, dtype) in schema.iter() {
499        let col = df.column(name)?;
500        let series = col.as_materialized_series();
501        let count = series.len();
502        let null_count = series.null_count();
503
504        let numeric_stats = if is_numeric_type(dtype) {
505            Some(compute_numeric_stats(
506                series,
507                options.include_skewness_kurtosis_outliers,
508            )?)
509        } else {
510            None
511        };
512
513        let categorical_stats = if is_categorical_type(dtype) {
514            Some(compute_categorical_stats(series)?)
515        } else {
516            None
517        };
518
519        let distribution_info =
520            if options.include_distribution_info && is_numeric_type(dtype) && null_count < count {
521                // Get sample for distribution inference
522                Some(infer_distribution(
523                    series,
524                    series,
525                    actual_sample_size.unwrap_or(count),
526                    should_sample,
527                ))
528            } else {
529                None
530            };
531
532        column_statistics.push(ColumnStatistics {
533            name: name.to_string(),
534            dtype: dtype.clone(),
535            count,
536            null_count,
537            numeric_stats,
538            categorical_stats,
539            temporal_stats: temporal_stats_of(series)?,
540            distribution_info,
541        });
542    }
543
544    let distribution_analyses = if options.include_distribution_analyses {
545        column_statistics
546            .iter()
547            .filter_map(|col_stat| {
548                if let (Some(numeric_stats), Some(dist_info)) =
549                    (&col_stat.numeric_stats, &col_stat.distribution_info)
550                {
551                    if let Ok(series_col) = df.column(&col_stat.name) {
552                        let series = series_col.as_materialized_series();
553                        Some(compute_advanced_distribution_analysis(
554                            &col_stat.name,
555                            series,
556                            numeric_stats,
557                            dist_info,
558                            actual_sample_size.unwrap_or(total_rows),
559                            should_sample,
560                        ))
561                    } else {
562                        None
563                    }
564                } else {
565                    None
566                }
567            })
568            .collect()
569    } else {
570        Vec::new()
571    };
572
573    let correlation_matrix = if options.include_correlation_matrix {
574        compute_correlation_matrix(&df).ok()
575    } else {
576        None
577    };
578
579    Ok(AnalysisResults {
580        column_statistics,
581        total_rows,
582        sample_size: actual_sample_size,
583        per_value,
584        sample_seed: seed,
585        correlation_matrix,
586        distribution_analyses,
587    })
588}
589
590/// Builds describe-only AnalysisResults from a list of column statistics.
591///
592/// Used when completing chunked describe; correlation and distribution analyses stay empty/None.
593pub fn analysis_results_from_describe(
594    column_statistics: Vec<ColumnStatistics>,
595    total_rows: usize,
596    sample_size: Option<usize>,
597    sample_seed: u64,
598) -> AnalysisResults {
599    AnalysisResults {
600        column_statistics,
601        total_rows,
602        sample_size,
603        per_value: None,
604        sample_seed,
605        correlation_matrix: None,
606        distribution_analyses: Vec::new(),
607    }
608}
609
610/// Builds aggregation expressions for describe (count, null_count, mean, std, percentiles, min, max).
611/// Used so we can run a single collect on a LazyFrame without materializing all rows.
612fn build_describe_aggregation_exprs(schema: &Schema) -> Vec<Expr> {
613    let mut exprs = Vec::new();
614    for (name, dtype) in schema.iter() {
615        let name = name.as_str();
616        let prefix = format!("{}::", name);
617        exprs.push(col(name).count().alias(format!("{}count", prefix)));
618        exprs.push(
619            col(name)
620                .null_count()
621                .alias(format!("{}null_count", prefix)),
622        );
623        if is_numeric_type(dtype) {
624            let c = col(name).cast(DataType::Float64);
625            exprs.push(c.clone().mean().alias(format!("{}mean", prefix)));
626            exprs.push(c.clone().std(1).alias(format!("{}std", prefix)));
627            exprs.push(c.clone().min().alias(format!("{}min", prefix)));
628            // Without NaN, which sorts above every number and would be the upper
629            // quantiles of a column holding enough of it.
630            let numbers = c.clone().drop_nans();
631            exprs.push(
632                numbers
633                    .clone()
634                    .quantile(lit(0.25), QuantileMethod::Nearest)
635                    .alias(format!("{}q25", prefix)),
636            );
637            exprs.push(
638                numbers
639                    .clone()
640                    .quantile(lit(0.5), QuantileMethod::Nearest)
641                    .alias(format!("{}median", prefix)),
642            );
643            exprs.push(
644                numbers
645                    .clone()
646                    .quantile(lit(0.75), QuantileMethod::Nearest)
647                    .alias(format!("{}q75", prefix)),
648            );
649            exprs.push(c.max().alias(format!("{}max", prefix)));
650        } else if is_categorical_type(dtype) {
651            exprs.push(col(name).min().alias(format!("{}min", prefix)));
652            exprs.push(col(name).max().alias(format!("{}max", prefix)));
653        } else if is_temporal_type(dtype) {
654            // Mean and quantiles run on the physical integers and are cast back, so
655            // each is a value of the column's own type, as in Polars' describe.
656            let physical = dtype.to_physical();
657            let back = |e: Expr| e.cast(physical.clone()).cast(dtype.clone());
658            let p = col(name).cast(physical.clone());
659            exprs.push(
660                back(p.clone().cast(DataType::Float64).mean()).alias(format!("{}mean", prefix)),
661            );
662            exprs.push(col(name).min().alias(format!("{}min", prefix)));
663            for (q, stat) in [(0.25, "q25"), (0.5, "median"), (0.75, "q75")] {
664                exprs.push(
665                    back(p.clone().quantile(lit(q), QuantileMethod::Nearest))
666                        .alias(format!("{}{}", prefix, stat)),
667                );
668            }
669            exprs.push(col(name).max().alias(format!("{}max", prefix)));
670        }
671    }
672    exprs
673}
674
675/// Parses the single-row aggregation result from describe into column statistics.
676fn parse_describe_agg_row(agg_df: &DataFrame, schema: &Schema) -> Vec<ColumnStatistics> {
677    let row = 0usize;
678    let mut column_statistics = Vec::with_capacity(schema.len());
679    for (name, dtype) in schema.iter() {
680        let name_str = name.as_str();
681        let prefix = format!("{}::", name_str);
682        let count: usize = agg_df
683            .column(&format!("{}count", prefix))
684            .ok()
685            .map(|s| match s.get(row) {
686                Ok(AnyValue::UInt32(x)) => x as usize,
687                _ => 0,
688            })
689            .unwrap_or(0);
690        let null_count: usize = agg_df
691            .column(&format!("{}null_count", prefix))
692            .ok()
693            .map(|s| match s.get(row) {
694                Ok(AnyValue::UInt32(x)) => x as usize,
695                _ => 0,
696            })
697            .unwrap_or(0);
698        let numeric_stats = if is_numeric_type(dtype) {
699            let mean = get_f64(agg_df, &format!("{}mean", prefix), row);
700            let std = get_f64(agg_df, &format!("{}std", prefix), row);
701            let min = get_f64(agg_df, &format!("{}min", prefix), row);
702            let q25 = get_f64(agg_df, &format!("{}q25", prefix), row);
703            let median = get_f64(agg_df, &format!("{}median", prefix), row);
704            let q75 = get_f64(agg_df, &format!("{}q75", prefix), row);
705            let max = get_f64(agg_df, &format!("{}max", prefix), row);
706            let mut percentiles = HashMap::new();
707            percentiles.insert(25u8, q25);
708            percentiles.insert(50u8, median);
709            percentiles.insert(75u8, q75);
710            Some(NumericStatistics {
711                mean,
712                std,
713                min,
714                max,
715                median,
716                q25,
717                q75,
718                percentiles,
719                skewness: 0.0,
720                kurtosis: 3.0,
721                outliers_iqr: 0,
722                outliers_zscore: 0,
723            })
724        } else {
725            None
726        };
727        let categorical_stats = if is_categorical_type(dtype) {
728            let min = get_str(agg_df, &format!("{}min", prefix), row);
729            let max = get_str(agg_df, &format!("{}max", prefix), row);
730            Some(CategoricalStatistics {
731                unique_count: 0,
732                mode: None,
733                top_values: Vec::new(),
734                min,
735                max,
736            })
737        } else {
738            None
739        };
740        let temporal_stats = is_temporal_type(dtype).then(|| {
741            let value = |stat: &str| get_value_str(agg_df, &format!("{}{}", prefix, stat), row);
742            TemporalStatistics {
743                mean: value("mean"),
744                min: value("min"),
745                q25: value("q25"),
746                median: value("median"),
747                q75: value("q75"),
748                max: value("max"),
749            }
750        });
751        column_statistics.push(ColumnStatistics {
752            name: name_str.to_string(),
753            dtype: dtype.clone(),
754            count,
755            null_count,
756            numeric_stats,
757            categorical_stats,
758            temporal_stats,
759            distribution_info: None,
760        });
761    }
762    column_statistics
763}
764
765/// Computes describe statistics from a LazyFrame without materializing all rows.
766/// When sampling is disabled, runs a single aggregation collect (like Polars describe) for similar performance.
767/// When sampling is enabled, samples then runs describe on the sample.
768/// Describe statistics for a frame. With `sample_size`, a table with more rows than
769/// that is described from a sample (see [`analysis_rows`]); without it, every row is
770/// aggregated in one streaming pass, never held. `known_total` saves a count.
771pub fn compute_describe_from_lazy(
772    lf: &LazyFrame,
773    known_total: Option<usize>,
774    sample: &crate::sampling::Sample,
775    polars_streaming: bool,
776) -> Result<AnalysisResults> {
777    let schema = lf.clone().collect_schema()?;
778    let seed = sample.seed;
779    if sample.method != crate::sampling::SampleMethod::EveryRow {
780        let rows = crate::sampling::read(lf, sample, known_total, polars_streaming)?;
781        let mut results = compute_describe_single_aggregation(
782            &rows.df,
783            &schema,
784            rows.total_rows,
785            rows.sample_size,
786            seed,
787            polars_streaming,
788        )?;
789        results.per_value = rows.per_value.map(|per_value| per_value.kept);
790        return Ok(results);
791    }
792    let total_rows = match known_total {
793        Some(total) => total,
794        None => count_rows(lf, polars_streaming)?,
795    };
796    let exprs = build_describe_aggregation_exprs(&schema);
797    let agg_df = collect_lazy(lf.clone().select(exprs), polars_streaming).map_err(Report::from)?;
798    let column_statistics = parse_describe_agg_row(&agg_df, &schema);
799    Ok(analysis_results_from_describe(
800        column_statistics,
801        total_rows,
802        None,
803        seed,
804    ))
805}
806
807/// Computes describe statistics in a single aggregation pass over the DataFrame.
808/// Uses one collect() with aggregated expressions for all columns (count, null_count, mean, std, min, percentiles, max).
809pub fn compute_describe_single_aggregation(
810    df: &DataFrame,
811    schema: &Schema,
812    total_rows: usize,
813    sample_size: Option<usize>,
814    sample_seed: u64,
815    polars_streaming: bool,
816) -> Result<AnalysisResults> {
817    let exprs = build_describe_aggregation_exprs(schema);
818    let agg_df =
819        collect_lazy(df.clone().lazy().select(exprs), polars_streaming).map_err(Report::from)?;
820    let column_statistics = parse_describe_agg_row(&agg_df, schema);
821    Ok(analysis_results_from_describe(
822        column_statistics,
823        total_rows,
824        sample_size,
825        sample_seed,
826    ))
827}
828
829fn get_f64(df: &DataFrame, col_name: &str, row: usize) -> f64 {
830    df.column(col_name)
831        .ok()
832        .and_then(|s| {
833            let v = s.get(row).ok()?;
834            match v {
835                AnyValue::Float64(x) => Some(x),
836                AnyValue::Float32(x) => Some(x as f64),
837                AnyValue::Int32(x) => Some(x as f64),
838                AnyValue::Int64(x) => Some(x as f64),
839                AnyValue::UInt32(x) => Some(x as f64),
840                AnyValue::Null => Some(f64::NAN),
841                _ => None,
842            }
843        })
844        .unwrap_or(f64::NAN)
845}
846
847fn get_str(df: &DataFrame, col_name: &str, row: usize) -> Option<String> {
848    df.column(col_name).ok().and_then(|s| {
849        s.get(row)
850            .ok()
851            .map(|v| crate::exact::str_value(&v).to_string())
852    })
853}
854
855/// Uses Polars' definition so Int128, UInt128, Decimal, and future numeric types are included.
856fn is_numeric_type(dtype: &DataType) -> bool {
857    dtype.is_numeric()
858}
859
860fn is_categorical_type(dtype: &DataType) -> bool {
861    matches!(dtype, DataType::String | DataType::Categorical(..))
862}
863
864fn is_temporal_type(dtype: &DataType) -> bool {
865    matches!(
866        dtype,
867        DataType::Date | DataType::Datetime(..) | DataType::Time | DataType::Duration(_)
868    )
869}
870
871/// The rows an analysis reads, and how many the table has.
872pub struct AnalysisRows {
873    pub df: DataFrame,
874    pub total_rows: usize,
875    /// How many rows were sampled, when the table had more than the analysis reads.
876    pub sample_size: Option<usize>,
877    /// What an equal-per-value sample kept and counted.
878    pub per_value: Option<crate::sampling::PerValue>,
879}
880
881/// How many places across the table a block sample reads from. Enough that no one
882/// stretch of it decides the answer, few enough that each is a row group or two.
883const SAMPLE_BLOCKS: usize = 50;
884
885/// How many runs of a block sample are read at once.
886const SAMPLE_READERS: usize = 8;
887
888/// The row index the streaming sampler ranks rows by, dropped before anyone sees it.
889const SAMPLE_POSITION: &str = "__datui_sample_position";
890
891/// Count a frame's rows.
892pub fn count_rows(lf: &LazyFrame, polars_streaming: bool) -> Result<usize> {
893    let count_df = collect_lazy(
894        crate::widgets::datatable::row_count_lf(lf),
895        polars_streaming,
896    )
897    .map_err(Report::from)?;
898    Ok(match count_df.get(0).and_then(|row| row.first().cloned()) {
899        Some(AnyValue::UInt64(n)) => n as usize,
900        Some(AnyValue::UInt32(n)) => n as usize,
901        _ => 0,
902    })
903}
904
905/// Read the rows an analysis works on: all of them when the table has no more than
906/// `sample_rows` (or `sample_rows` is `None`), and otherwise a seeded sample of that
907/// many, spread across the whole table rather than taken from its head.
908///
909/// Two ways to spread it, chosen by what the plan can do cheaply:
910///
911/// - A plan whose slices reach into a single Parquet or IPC scan reads
912///   [`SAMPLE_BLOCKS`] short runs at seeded places across the table. Each run is a
913///   row group or two, so a sample of a 400-million-row hive table reads a few dozen
914///   row groups, not the table. `known_total` saves the count; the footers give it
915///   cheaply otherwise.
916/// - Anything else — a filter, a query, a union of files, a CSV — is read once as a
917///   stream, keeping the rows whose seeded rank is lowest. That is a uniform sample
918///   in bounded memory, and the same pass counts the rows, so a filtered view is
919///   read once rather than counted and then read.
920pub fn analysis_rows(
921    lf: &LazyFrame,
922    sample_rows: Option<usize>,
923    known_total: Option<usize>,
924    seed: u64,
925    polars_streaming: bool,
926) -> Result<AnalysisRows> {
927    analysis_rows_watched(lf, sample_rows, known_total, seed, polars_streaming, None)
928}
929
930/// [`analysis_rows`], stopping when `watch` says to: the streamed pass between
931/// batches, the seeded runs between runs. A whole read is one collect, which runs to
932/// its end.
933pub(crate) fn analysis_rows_watched(
934    lf: &LazyFrame,
935    sample_rows: Option<usize>,
936    known_total: Option<usize>,
937    seed: u64,
938    polars_streaming: bool,
939    watch: Option<&crate::sampling::ReadWatch>,
940) -> Result<AnalysisRows> {
941    sample_rows_counting(
942        lf,
943        sample_rows,
944        known_total,
945        seed,
946        polars_streaming,
947        watch,
948        None,
949    )
950    .map(|read| read.rows)
951}
952
953/// [`analysis_rows_watched`], keeping where each row sat, and counting every row by
954/// `count` when the read sees every row: a streamed pass, or a table read whole
955/// because it is under twice the sample. Seeded runs see too few rows to count, and
956/// a read of the whole scope is not a sample, so neither counts.
957pub(crate) fn sample_rows_counting(
958    lf: &LazyFrame,
959    sample_rows: Option<usize>,
960    known_total: Option<usize>,
961    seed: u64,
962    polars_streaming: bool,
963    watch: Option<&crate::sampling::ReadWatch>,
964    count: Option<&Expr>,
965) -> Result<crate::sampling::SampledRows> {
966    let whole = |df: DataFrame, total_rows: usize| crate::sampling::SampledRows {
967        positions: (0..df.height() as IdxSize).collect(),
968        rows: AnalysisRows {
969            df,
970            total_rows,
971            sample_size: None,
972            per_value: None,
973        },
974        counted: None,
975    };
976    let Some(n) = sample_rows.filter(|n| *n > 0) else {
977        let df = collect_lazy(lf.clone(), polars_streaming).map_err(Report::from)?;
978        let total_rows = df.height();
979        return Ok(whole(df, total_rows));
980    };
981    if !slices_reach_into_the_scan(lf) {
982        let read = stream_sample(lf, n, seed, watch, count)?;
983        let sample_size = (read.seen > n).then_some(read.df.height());
984        return Ok(crate::sampling::SampledRows {
985            rows: AnalysisRows {
986                df: read.df,
987                total_rows: read.seen,
988                sample_size,
989                per_value: None,
990            },
991            positions: read.positions,
992            counted: read.counted,
993        });
994    }
995    let total_rows = match known_total {
996        Some(total) => total,
997        None => count_rows(lf, polars_streaming)?,
998    };
999    if total_rows <= n {
1000        let df = collect_lazy(lf.clone(), polars_streaming).map_err(Report::from)?;
1001        return Ok(whole(df, total_rows));
1002    }
1003    let along = Along {
1004        watch,
1005        count,
1006        on_run: None,
1007    };
1008    let read = block_sample(lf, total_rows, n, seed, polars_streaming, along)?;
1009    Ok(crate::sampling::SampledRows {
1010        rows: AnalysisRows {
1011            sample_size: Some(read.df.height()),
1012            df: read.df,
1013            total_rows,
1014            per_value: None,
1015        },
1016        positions: read.positions,
1017        counted: read.counted,
1018    })
1019}
1020
1021/// Whether a slice of this plan is read by the scan of one file, skipping what comes
1022/// before it: true of a single Parquet or IPC file, which seeks by row group, with or
1023/// without columns stubbed above it. Not of a filter or a CSV, whose slice reads
1024/// everything ahead of it, nor of a scan of many files, where each slice opens the
1025/// footer of every file before it — measured on 135 files in S3, fifty slices took
1026/// longer than streaming all 37 million rows once.
1027///
1028/// Asked of the optimized plan because that is where the answer is, for every route a
1029/// frame can have been built by: pushed into the scan, the slice is a property of the
1030/// `SCAN` (`SLICE: Positive`); left above it, a node of its own (`SLICE[`). Should a
1031/// Polars upgrade change how the plan is described, this says no and the streaming
1032/// sampler takes over: slower, never wrong.
1033pub fn slices_reach_into_the_scan(lf: &LazyFrame) -> bool {
1034    let Ok(plan) = lf.clone().slice(1, 1).describe_optimized_plan() else {
1035        return false;
1036    };
1037    let scans: Vec<&str> = plan
1038        .lines()
1039        .map(str::trim)
1040        .filter(|l| l.starts_with("Parquet SCAN") || l.starts_with("IPC SCAN"))
1041        .collect();
1042    let [scan] = scans.as_slice() else {
1043        return false;
1044    };
1045    let one_source = !scan.contains("other sources") && !scan.contains(", ");
1046    let total_scans = plan.matches(" SCAN").count();
1047    one_source && total_scans == 1 && plan.contains("SLICE: Positive") && !plan.contains("SLICE[")
1048}
1049
1050/// [`block_sample`] for a sample shown as it is drawn: each run goes to `on_run`, with
1051/// where it starts, as it lands. A table under twice the sample is read whole and cut,
1052/// and comes back as one frame instead. A stop ends the read with the runs so far
1053/// delivered.
1054pub(crate) fn block_sample_live(
1055    lf: &LazyFrame,
1056    total_rows: usize,
1057    n: usize,
1058    seed: u64,
1059    polars_streaming: bool,
1060    watch: &crate::sampling::ReadWatch,
1061    on_run: &OnRun<'_>,
1062) -> Result<Option<DataFrame>> {
1063    let along = Along {
1064        watch: Some(watch),
1065        count: None,
1066        on_run: Some(on_run),
1067    };
1068    let read = block_sample(lf, total_rows, n, seed, polars_streaming, along)?;
1069    Ok((total_rows < 2 * n).then_some(read.df))
1070}
1071
1072/// What a block sample does beside reading its runs: stops when `watch` says to,
1073/// counts `count`'s key, and hands each run to `on_run` as it lands.
1074struct Along<'a> {
1075    watch: Option<&'a crate::sampling::ReadWatch>,
1076    count: Option<&'a Expr>,
1077    on_run: Option<&'a OnRun<'a>>,
1078}
1079
1080/// Told of each run of a block sample as it lands, with where it starts.
1081pub(crate) type OnRun<'a> = dyn Fn(usize, &DataFrame) + Sync + 'a;
1082
1083/// `n` rows as [`SAMPLE_BLOCKS`] runs at seeded places across `total_rows`, in table
1084/// order. Each run is collected on its own: as one union the runs share a subplan, and
1085/// Polars caches a shared subplan whole. They are collected [`SAMPLE_READERS`] at a
1086/// time, because on an object store each is a round trip and fifty in a row is the
1087/// wait this exists to avoid.
1088fn block_sample(
1089    lf: &LazyFrame,
1090    total_rows: usize,
1091    n: usize,
1092    seed: u64,
1093    polars_streaming: bool,
1094    along: Along<'_>,
1095) -> Result<StreamRead> {
1096    let Along {
1097        watch,
1098        count,
1099        on_run,
1100    } = along;
1101    // Under twice the sample, reading the table is about as cheap as reading runs of
1102    // it, and runs that must fit side by side would crowd or overlap. Read it and keep
1103    // a seeded uniform `n` of it instead, counting `count`'s key from the rows read.
1104    if total_rows < 2 * n {
1105        let df = collect_lazy(lf.clone(), polars_streaming).map_err(Report::from)?;
1106        let counted = match count {
1107            Some(key) => {
1108                let mut keys = df
1109                    .clone()
1110                    .lazy()
1111                    .select([key.clone().alias(crate::sampling::COUNT_KEY)])
1112                    .collect()?;
1113                let mut counter = crate::sampling::KeyCounter::default();
1114                counter.observe(&mut keys)?;
1115                Some(counter.finish())
1116            }
1117            None => None,
1118        };
1119        let mut ranked: Vec<(u64, IdxSize)> = (0..df.height())
1120            .map(|i| (sample_rank(seed, i as u64), i as IdxSize))
1121            .collect();
1122        ranked.sort_unstable();
1123        let mut keep: Vec<IdxSize> = ranked.into_iter().take(n).map(|(_, i)| i).collect();
1124        keep.sort_unstable();
1125        let df = df.take(&IdxCa::from_vec("sample".into(), keep.clone()))?;
1126        return Ok(StreamRead {
1127            df,
1128            seen: total_rows,
1129            positions: keep,
1130            counted,
1131        });
1132    }
1133    let blocks = SAMPLE_BLOCKS.min(n).max(1);
1134    // Exactly `n` rows between the runs, so none is cut off the end, and each fits in
1135    // its own stretch of the table: a stretch is at least `2n / blocks` rows long.
1136    let stride = total_rows / blocks;
1137    let runs: Vec<(usize, usize)> = (0..blocks)
1138        .map(|block| {
1139            let run = (block + 1) * n / blocks - block * n / blocks;
1140            let room = stride.saturating_sub(run) as u64;
1141            let offset = block * stride + (sample_rank(seed, block as u64) % (room + 1)) as usize;
1142            (offset, run)
1143        })
1144        .collect();
1145    let next = std::sync::atomic::AtomicUsize::new(0);
1146    let read: Vec<Result<(usize, DataFrame)>> = std::thread::scope(|scope| {
1147        let workers: Vec<_> = (0..SAMPLE_READERS.min(blocks))
1148            .map(|_| {
1149                scope.spawn(|| {
1150                    let mut read = Vec::new();
1151                    loop {
1152                        if watch.is_some_and(|watch| watch.stopped()) {
1153                            break;
1154                        }
1155                        let block = next.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1156                        let Some((offset, run)) = runs.get(block) else {
1157                            break;
1158                        };
1159                        let rows = collect_lazy(
1160                            lf.clone().slice(*offset as i64, *run as IdxSize),
1161                            polars_streaming,
1162                        )
1163                        .map(|df| {
1164                            if let Some(watch) = watch {
1165                                watch.saw(df.height());
1166                            }
1167                            if let Some(on_run) = on_run {
1168                                on_run(*offset, &df);
1169                            }
1170                            (block, df)
1171                        })
1172                        .map_err(Report::from);
1173                        read.push(rows);
1174                    }
1175                    read
1176                })
1177            })
1178            .collect();
1179        workers
1180            .into_iter()
1181            .flat_map(|worker| match worker.join() {
1182                Ok(read) => read,
1183                // A reader that died is an error, not a smaller sample.
1184                Err(_) => vec![Err(Report::msg("a sample reader failed"))],
1185            })
1186            .collect()
1187    });
1188    if let Some(watch) = watch {
1189        watch.check()?;
1190    }
1191    let mut read = read.into_iter().collect::<Result<Vec<_>>>()?;
1192    read.sort_by_key(|(block, _)| *block);
1193    let mut out: Option<DataFrame> = None;
1194    let mut positions = Vec::with_capacity(n);
1195    for (block, rows) in read {
1196        let offset = runs[block].0;
1197        positions.extend((0..rows.height()).map(|row| (offset + row) as IdxSize));
1198        out = Some(match out {
1199            Some(frame) => frame.vstack(&rows)?,
1200            None => rows,
1201        });
1202    }
1203    Ok(StreamRead {
1204        df: out.unwrap_or_default(),
1205        seen: total_rows,
1206        positions,
1207        counted: None,
1208    })
1209}
1210
1211/// What [`stream_sample`] or [`block_sample`] read.
1212struct StreamRead {
1213    df: DataFrame,
1214    /// Rows in the scope.
1215    seen: usize,
1216    positions: Vec<IdxSize>,
1217    counted: Option<crate::sampling::Counted>,
1218}
1219
1220/// A uniform sample of `n` rows from one streamed pass, and how many rows there were,
1221/// with every row counted by `count` on the way.
1222fn stream_sample(
1223    lf: &LazyFrame,
1224    n: usize,
1225    seed: u64,
1226    watch: Option<&crate::sampling::ReadWatch>,
1227    count: Option<&Expr>,
1228) -> Result<StreamRead> {
1229    let state = std::sync::Arc::new(std::sync::Mutex::new(Reservoir::new(n, seed)));
1230    let callback_state = std::sync::Arc::clone(&state);
1231    let callback_watch = watch.cloned();
1232    let sink = crate::sampling::with_count_key(lf.clone(), count)
1233        .with_row_index(SAMPLE_POSITION, None)
1234        .sink_batches(
1235            PlanCallback::new(move |batch: DataFrame| {
1236                // True stops the sink: a cancel ends the read at the next batch.
1237                if let Some(watch) = &callback_watch {
1238                    if watch.stopped() {
1239                        return Ok(true);
1240                    }
1241                    watch.saw(batch.height());
1242                }
1243                let mut reservoir = callback_state
1244                    .lock()
1245                    .map_err(|_| PolarsError::ComputeError("sampler lock failed".into()))?;
1246                reservoir.observe(batch)?;
1247                if let Some(watch) = &callback_watch {
1248                    let held = reservoir.kept.as_ref();
1249                    watch.hold(
1250                        held.map_or(0, |kept| kept.estimated_size() as u64),
1251                        held.map_or(0, DataFrame::height),
1252                    );
1253                }
1254                Ok(false)
1255            }),
1256            true,
1257            None,
1258        )?;
1259    // Streaming whatever the setting: holding the table is what this is here to avoid.
1260    collect_lazy(sink, true).map_err(Report::from)?;
1261    if let Some(watch) = watch {
1262        watch.check()?;
1263    }
1264    let mut reservoir = std::mem::take(
1265        &mut *state
1266            .lock()
1267            .map_err(|_| Report::msg("sampler lock failed"))?,
1268    );
1269    let seen = reservoir.seen;
1270    let counted = count
1271        .is_some()
1272        .then(|| std::mem::take(&mut reservoir.counter).finish());
1273    let (df, positions) = match reservoir.finish()? {
1274        Some(kept) => kept,
1275        // Nothing came through: an empty frame of the right shape.
1276        None => (
1277            collect_lazy(lf.clone().limit(0), true).map_err(Report::from)?,
1278            Vec::new(),
1279        ),
1280    };
1281    Ok(StreamRead {
1282        df,
1283        seen,
1284        positions,
1285        counted,
1286    })
1287}
1288
1289/// The `n` rows with the lowest seeded rank seen so far. Held to at most twice `n`
1290/// between prunes, so memory is bounded by the sample and not by the table.
1291#[derive(Default)]
1292struct Reservoir {
1293    n: usize,
1294    seed: u64,
1295    seen: usize,
1296    kept: Option<DataFrame>,
1297    ranks: Vec<u64>,
1298    /// Rows ranked at or above this cannot make the sample: `n` lower ones are held.
1299    bar: u64,
1300    counter: crate::sampling::KeyCounter,
1301}
1302
1303impl Reservoir {
1304    fn new(n: usize, seed: u64) -> Self {
1305        Self {
1306            n,
1307            seed,
1308            bar: u64::MAX,
1309            ..Default::default()
1310        }
1311    }
1312
1313    fn observe(&mut self, mut batch: DataFrame) -> PolarsResult<()> {
1314        self.counter.observe(&mut batch)?;
1315        self.seen += batch.height();
1316        let positions = batch.column(SAMPLE_POSITION)?.idx()?;
1317        let mut picked = Vec::new();
1318        let mut ranks = Vec::new();
1319        for (index, position) in positions.into_no_null_iter().enumerate() {
1320            let rank = sample_rank(self.seed, position as u64);
1321            if rank < self.bar {
1322                picked.push(index as IdxSize);
1323                ranks.push(rank);
1324            }
1325        }
1326        if picked.is_empty() {
1327            return Ok(());
1328        }
1329        let rows = batch.take(&IdxCa::from_vec("picked".into(), picked))?;
1330        self.kept = Some(match self.kept.take() {
1331            Some(kept) => kept.vstack(&rows)?,
1332            None => rows,
1333        });
1334        self.ranks.extend(ranks);
1335        if self.ranks.len() > 2 * self.n {
1336            self.prune()?;
1337        }
1338        Ok(())
1339    }
1340
1341    /// Keep the `n` lowest-ranked rows, and raise the bar to the highest of them.
1342    fn prune(&mut self) -> PolarsResult<()> {
1343        let Some(kept) = self.kept.take() else {
1344            return Ok(());
1345        };
1346        let mut order: Vec<usize> = (0..self.ranks.len()).collect();
1347        order.sort_unstable_by_key(|i| self.ranks[*i]);
1348        order.truncate(self.n);
1349        let take: Vec<IdxSize> = order.iter().map(|i| *i as IdxSize).collect();
1350        self.kept = Some(kept.take(&IdxCa::from_vec("kept".into(), take))?);
1351        self.ranks = order.iter().map(|i| self.ranks[*i]).collect();
1352        if self.ranks.len() == self.n {
1353            self.bar = self.ranks.iter().copied().max().unwrap_or(u64::MAX);
1354        }
1355        Ok(())
1356    }
1357
1358    /// The sample, back in table order without the position column, and where each
1359    /// of its rows sat.
1360    fn finish(mut self) -> PolarsResult<Option<(DataFrame, Vec<IdxSize>)>> {
1361        self.prune()?;
1362        let Some(kept) = self.kept else {
1363            return Ok(None);
1364        };
1365        let sorted = kept.sort([SAMPLE_POSITION], SortMultipleOptions::default())?;
1366        let positions = sorted
1367            .column(SAMPLE_POSITION)?
1368            .idx()?
1369            .into_no_null_iter()
1370            .collect();
1371        Ok(Some((sorted.drop(SAMPLE_POSITION)?, positions)))
1372    }
1373}
1374
1375/// A seeded, well-mixed rank for a row position (SplitMix64's finalizer). The same seed
1376/// and table give the same sample; another seed gives another.
1377pub(crate) fn sample_rank(seed: u64, position: u64) -> u64 {
1378    let mut value = seed ^ position.wrapping_mul(0x9e37_79b9_7f4a_7c15);
1379    value = (value ^ (value >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
1380    value = (value ^ (value >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
1381    value ^ (value >> 31)
1382}
1383
1384/// Up to ten thousand of a column's values, spread across it, as `f64`. NaN and
1385/// infinities are left out: no distribution has them, and one NaN is enough to leave
1386/// a sort by `partial_cmp` out of order.
1387fn get_numeric_values_as_f64(series: &Series) -> Vec<f64> {
1388    let max_len = 10000;
1389    // Every k-th value, not the first ten thousand: a sample is spread across the
1390    // table, and its head is one stretch of it.
1391    let limited_series = if series.len() > max_len {
1392        let step = series.len().div_ceil(max_len);
1393        series
1394            .gather_every(step, 0)
1395            .unwrap_or_else(|_| series.slice(0, max_len))
1396    } else {
1397        series.clone()
1398    };
1399
1400    if let Ok(f64_series) = limited_series.f64() {
1401        f64_series
1402            .iter()
1403            .flatten()
1404            .filter(|v| v.is_finite())
1405            .take(max_len)
1406            .collect()
1407    } else if let Ok(i64_series) = limited_series.i64() {
1408        i64_series
1409            .iter()
1410            .filter_map(|v| v.map(|x| x as f64))
1411            .take(max_len)
1412            .collect()
1413    } else if let Ok(i32_series) = limited_series.i32() {
1414        i32_series
1415            .iter()
1416            .filter_map(|v| v.map(|x| x as f64))
1417            .take(max_len)
1418            .collect()
1419    } else if let Ok(u64_series) = limited_series.u64() {
1420        u64_series
1421            .iter()
1422            .filter_map(|v| v.map(|x| x as f64))
1423            .take(max_len)
1424            .collect()
1425    } else if let Ok(u32_series) = limited_series.u32() {
1426        u32_series
1427            .iter()
1428            .filter_map(|v| v.map(|x| x as f64))
1429            .take(max_len)
1430            .collect()
1431    } else if let Ok(f32_series) = limited_series.f32() {
1432        f32_series
1433            .iter()
1434            .flatten()
1435            .filter(|v| v.is_finite())
1436            .map(f64::from)
1437            .take(max_len)
1438            .collect()
1439    } else {
1440        match limited_series.cast(&DataType::Float64) {
1441            Ok(cast_series) => {
1442                if let Ok(f64_series) = cast_series.f64() {
1443                    f64_series
1444                        .iter()
1445                        .flatten()
1446                        .filter(|v| v.is_finite())
1447                        .take(max_len)
1448                        .collect()
1449                } else {
1450                    Vec::new()
1451                }
1452            }
1453            Err(_) => Vec::new(),
1454        }
1455    }
1456}
1457
1458/// Every finite value of a numeric column as `f64`: nulls, NaN and infinities left out.
1459fn finite_values(series: &Series) -> Vec<f64> {
1460    let Ok(floats) = series.cast(&DataType::Float64) else {
1461        return Vec::new();
1462    };
1463    let Ok(floats) = floats.f64() else {
1464        return Vec::new();
1465    };
1466    floats.iter().flatten().filter(|v| v.is_finite()).collect()
1467}
1468
1469fn compute_numeric_stats(series: &Series, include_advanced: bool) -> Result<NumericStatistics> {
1470    // Cast and aggregate as Describe does (`build_describe_aggregation_exprs`), so a
1471    // sample's median is one number wherever it is shown.
1472    let floats = series.cast(&DataType::Float64)?;
1473    let mean = floats.mean().unwrap_or(f64::NAN);
1474    let std = floats.std(1).unwrap_or(f64::NAN);
1475    let min = floats.min::<f64>()?.unwrap_or(f64::NAN);
1476    let max = floats.max::<f64>()?.unwrap_or(f64::NAN);
1477
1478    // NaN sorts above every number, so a column with some would have them as its
1479    // upper percentiles and fences; Describe leaves them out too. One sort for all.
1480    let floats = floats.f64()?;
1481    let numbers = floats.filter(&floats.is_not_nan())?;
1482    const PERCENTILES: [u8; 7] = [1, 5, 25, 50, 75, 95, 99];
1483    let quantiles = PERCENTILES.map(|p| f64::from(p) / 100.0);
1484    let values = numbers.quantiles(&quantiles, QuantileMethod::Nearest)?;
1485    let percentiles: HashMap<u8, f64> = PERCENTILES
1486        .into_iter()
1487        .zip(values)
1488        .map(|(p, value)| (p, value.unwrap_or(f64::NAN)))
1489        .collect();
1490
1491    let median = percentiles[&50];
1492    let q25 = percentiles[&25];
1493    let q75 = percentiles[&75];
1494
1495    let (skewness, kurtosis, outliers_iqr, outliers_zscore) = if include_advanced {
1496        let values = finite_values(series);
1497        let (skewness, kurtosis) = skewness_and_kurtosis(&values);
1498        let (out_iqr, out_zscore) = detect_outliers(&values, q25, q75);
1499        (skewness, kurtosis, out_iqr, out_zscore)
1500    } else {
1501        (0.0, 3.0, 0, 0) // Default values when not computed
1502    };
1503
1504    Ok(NumericStatistics {
1505        mean,
1506        std,
1507        min,
1508        max,
1509        median,
1510        q25,
1511        q75,
1512        percentiles,
1513        skewness,
1514        kurtosis,
1515        outliers_iqr,
1516        outliers_zscore,
1517    })
1518}
1519
1520/// Mean and sample standard deviation (ddof 1) of a set of values.
1521fn mean_and_std(values: &[f64]) -> (f64, f64) {
1522    let n = values.len() as f64;
1523    if values.len() < 2 {
1524        return (values.first().copied().unwrap_or(f64::NAN), f64::NAN);
1525    }
1526    let mean = values.iter().sum::<f64>() / n;
1527    let sum_squares: f64 = values.iter().map(|v| (v - mean).powi(2)).sum();
1528    (mean, (sum_squares / (n - 1.0)).sqrt())
1529}
1530
1531/// Skewness and kurtosis of one set of values, `n` being their count: the
1532/// bias-corrected forms of Polars' `skew(bias=False)` and `kurtosis(bias=False)`,
1533/// kurtosis on the scale where a normal is 3. One value throughout, or too few to
1534/// say, is 0 and 3.
1535fn skewness_and_kurtosis(values: &[f64]) -> (f64, f64) {
1536    let count = values.len();
1537    if count < 3 || values.iter().all(|v| *v == values[0]) {
1538        return (0.0, 3.0);
1539    }
1540    let n = count as f64;
1541    let (mean, std) = mean_and_std(values);
1542    let (mut cubes, mut fourths) = (0.0, 0.0);
1543    for v in values {
1544        let z = (v - mean) / std;
1545        let z2 = z * z;
1546        cubes += z2 * z;
1547        fourths += z2 * z2;
1548    }
1549    let skewness = n / ((n - 1.0) * (n - 2.0)) * cubes;
1550    if count < 4 {
1551        return (skewness, 3.0);
1552    }
1553    let excess = n * (n + 1.0) / ((n - 1.0) * (n - 2.0) * (n - 3.0)) * fourths
1554        - 3.0 * (n - 1.0) * (n - 1.0) / ((n - 2.0) * (n - 3.0));
1555    (skewness, excess + 3.0)
1556}
1557
1558/// Where a value falls against the IQR fences and three standard deviations.
1559struct OutlierTest {
1560    lower_fence: f64,
1561    upper_fence: f64,
1562    mean: f64,
1563    std: f64,
1564}
1565
1566impl OutlierTest {
1567    /// Fences from the quartiles; the z-score from the values' own mean and std, so
1568    /// a NaN elsewhere in the column cannot void it.
1569    fn new(values: &[f64], q25: f64, q75: f64) -> Option<Self> {
1570        let (mean, std) = mean_and_std(values);
1571        if q25.is_nan() || q75.is_nan() || std.is_nan() || std == 0.0 {
1572            return None;
1573        }
1574        let iqr = q75 - q25;
1575        Some(Self {
1576            lower_fence: q25 - 1.5 * iqr,
1577            upper_fence: q75 + 1.5 * iqr,
1578            mean,
1579            std,
1580        })
1581    }
1582
1583    fn iqr_position(&self, value: f64) -> Option<IqrPosition> {
1584        if value < self.lower_fence {
1585            Some(IqrPosition::BelowLowerFence)
1586        } else if value > self.upper_fence {
1587            Some(IqrPosition::AboveUpperFence)
1588        } else {
1589            None
1590        }
1591    }
1592
1593    fn z_score(&self, value: f64) -> f64 {
1594        (value - self.mean).abs() / self.std
1595    }
1596}
1597
1598const Z_THRESHOLD: f64 = 3.0;
1599
1600/// IQR and z-score outliers among every value given.
1601fn detect_outliers(values: &[f64], q25: f64, q75: f64) -> (usize, usize) {
1602    let Some(test) = OutlierTest::new(values, q25, q75) else {
1603        return (0, 0);
1604    };
1605    values.iter().fold((0, 0), |(iqr, zscore), &v| {
1606        (
1607            iqr + usize::from(test.iqr_position(v).is_some()),
1608            zscore + usize::from(test.z_score(v) > Z_THRESHOLD),
1609        )
1610    })
1611}
1612
1613fn compute_categorical_stats(series: &Series) -> Result<CategoricalStatistics> {
1614    let value_counts = series.value_counts(false, false, "counts".into(), false)?;
1615    let unique_count = value_counts.height();
1616
1617    let mode = if unique_count > 0 {
1618        match value_counts.get(0) {
1619            Some(col) => col.first().map(|v| crate::exact::str_value(v).to_string()),
1620            _ => None,
1621        }
1622    } else {
1623        None
1624    };
1625
1626    let mut top_values = Vec::new();
1627    for i in 0..unique_count.min(10) {
1628        if let (Some(value_col), Some(count_col)) = (value_counts.get(0), value_counts.get(1))
1629            && let (Some(value), Some(count)) = (value_col.get(i), count_col.get(i))
1630        {
1631            let value_str = crate::exact::str_value(value);
1632            if let Ok(count_u32) = count.try_extract::<u32>() {
1633                top_values.push((value_str.to_string(), count_u32 as usize));
1634            }
1635        }
1636    }
1637
1638    let min = if let Ok(str_series) = series.str() {
1639        let mut min_val: Option<String> = None;
1640        for s in str_series.iter().flatten() {
1641            let s_str = s.to_string();
1642            min_val = match min_val {
1643                None => Some(s_str.clone()),
1644                Some(ref current) if s_str < *current => Some(s_str),
1645                Some(current) => Some(current),
1646            };
1647        }
1648        min_val
1649    } else {
1650        None
1651    };
1652
1653    let max = if let Ok(str_series) = series.str() {
1654        let mut max_val: Option<String> = None;
1655        for s in str_series.iter().flatten() {
1656            let s_str = s.to_string();
1657            max_val = match max_val {
1658                None => Some(s_str.clone()),
1659                Some(ref current) if s_str > *current => Some(s_str),
1660                Some(current) => Some(current),
1661            };
1662        }
1663        max_val
1664    } else {
1665        None
1666    };
1667
1668    Ok(CategoricalStatistics {
1669        unique_count,
1670        mode,
1671        top_values,
1672        min,
1673        max,
1674    })
1675}
1676
1677/// Seeds the fit tests' simulations, so the same values get the same p-values.
1678const FIT_SEED: u64 = 0x5eed_d157;
1679
1680fn infer_distribution(
1681    _series: &Series,
1682    sample: &Series,
1683    sample_size: usize,
1684    is_sampled: bool,
1685) -> DistributionInfo {
1686    if sample_size < 3 {
1687        return DistributionInfo {
1688            distribution_type: DistributionType::Unknown,
1689            confidence: 0.0,
1690            sample_size,
1691            is_sampled,
1692            fit_quality: None,
1693            fits: Vec::new(),
1694        };
1695    }
1696
1697    let max_convert = 10000.min(sample.len());
1698    let values: Vec<f64> = if sample.len() > max_convert {
1699        let all_values = get_numeric_values_as_f64(sample);
1700        all_values.into_iter().take(max_convert).collect()
1701    } else {
1702        get_numeric_values_as_f64(sample)
1703    };
1704
1705    if values.is_empty() {
1706        return DistributionInfo {
1707            distribution_type: DistributionType::Unknown,
1708            confidence: 0.0,
1709            sample_size,
1710            is_sampled,
1711            fit_quality: None,
1712            fits: Vec::new(),
1713        };
1714    }
1715
1716    let mean: f64 = values.iter().sum::<f64>() / values.len() as f64;
1717    let variance: f64 =
1718        values.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / (values.len() - 1) as f64;
1719    let std = variance.sqrt();
1720
1721    // One value throughout fits every distribution's degenerate case and none of
1722    // them usefully; a year column in a partitioned table is the usual one.
1723    if std == 0.0 {
1724        return DistributionInfo {
1725            distribution_type: DistributionType::Constant,
1726            confidence: 1.0,
1727            sample_size,
1728            is_sampled,
1729            fit_quality: None,
1730            fits: Vec::new(),
1731        };
1732    }
1733
1734    // Counts are described by a count distribution when one holds.
1735    let counts = values
1736        .iter()
1737        .all(|v| *v >= 0.0 && *v == v.floor() && v.is_finite());
1738    let fits = crate::distribution_fit::test_all(&values, FIT_SEED);
1739    let distribution_type = crate::distribution_fit::select(&fits, counts);
1740    // The figure beside the name is that family's p-value; with no clear fit, the best
1741    // any family managed, so the table can say how far from fitting it was.
1742    let confidence = fits
1743        .iter()
1744        .find(|(family, _)| *family == distribution_type)
1745        .and_then(|(_, outcome)| outcome.p_value())
1746        .or_else(|| {
1747            fits.iter()
1748                .filter_map(|(_, outcome)| outcome.p_value())
1749                .max_by(f64::total_cmp)
1750        })
1751        .unwrap_or(0.0);
1752    DistributionInfo {
1753        distribution_type,
1754        confidence,
1755        sample_size,
1756        is_sampled,
1757        fit_quality: Some(confidence),
1758        fits,
1759    }
1760}
1761
1762fn approximate_shapiro_wilk(values: &[f64]) -> (Option<f64>, Option<f64>) {
1763    let n = values.len();
1764    if n < 3 {
1765        return (None, None);
1766    }
1767
1768    let mean: f64 = values.iter().sum::<f64>() / n as f64;
1769    let variance: f64 = values.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / (n - 1) as f64;
1770    let std = variance.sqrt();
1771
1772    if std == 0.0 {
1773        return (None, None);
1774    }
1775
1776    let mut sorted = values.to_vec();
1777    sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
1778
1779    let mut sum_expected_sq = 0.0;
1780    let mut sum_data_sq = 0.0;
1781    let mut sum_product = 0.0;
1782
1783    for (i, &value) in sorted.iter().enumerate() {
1784        let p = (i as f64 + 1.0 - 0.375) / (n as f64 + 0.25);
1785        let expected_quantile = normal_quantile(p);
1786        let standardized_value = (value - mean) / std;
1787
1788        sum_expected_sq += expected_quantile * expected_quantile;
1789        sum_data_sq += standardized_value * standardized_value;
1790        sum_product += expected_quantile * standardized_value;
1791    }
1792
1793    let sw_stat = if sum_expected_sq > 0.0 && sum_data_sq > 0.0 {
1794        (sum_product * sum_product) / (sum_expected_sq * sum_data_sq)
1795    } else {
1796        0.0
1797    };
1798
1799    let sw_stat = sw_stat.clamp(0.0, 1.0);
1800    (Some(sw_stat), shapiro_francia_pvalue(sw_stat, n))
1801}
1802
1803/// The p-value of a normality statistic computed as above: the squared correlation of
1804/// the sorted values with normal scores, which is the Shapiro-Francia form of the
1805/// Shapiro-Wilk test. Royston's (1993) approximation, `ln(1 - W')` being close to
1806/// normal, for 5 to 5,000 values; `None` outside them.
1807///
1808/// It replaces a blend of W with skew and kurtosis penalties that was not a p-value:
1809/// 2,590 prices with W' = 0.929 read p = 0.855, "normal", where the test says p < 1e-20.
1810fn shapiro_francia_pvalue(w: f64, n: usize) -> Option<f64> {
1811    if !(5..=5_000).contains(&n) {
1812        return None;
1813    }
1814    if w >= 1.0 {
1815        return Some(1.0);
1816    }
1817    let u = (n as f64).ln();
1818    let v = u.ln();
1819    let mu = -1.2725 + 1.0521 * (v - u);
1820    let sigma = 1.0308 - 0.26758 * (v + 2.0 / u);
1821    let z = ((1.0 - w).ln() - mu) / sigma;
1822    Some((1.0 - normal_cdf(z, 0.0, 1.0)).clamp(0.0, 1.0))
1823}
1824
1825// Advanced distribution analysis computation
1826fn compute_advanced_distribution_analysis(
1827    column_name: &str,
1828    series: &Series,
1829    numeric_stats: &NumericStatistics,
1830    dist_info: &DistributionInfo,
1831    _sample_size: usize,
1832    is_sampled: bool,
1833) -> DistributionAnalysis {
1834    // At most five thousand, spread across the rows: the head of a table sorted by
1835    // date is its first few years.
1836    const MAX_VALUES: usize = 5_000;
1837    let mut values = get_numeric_values_as_f64(series);
1838    if values.len() > MAX_VALUES {
1839        let step = values.len().div_ceil(MAX_VALUES);
1840        values = values.into_iter().step_by(step).collect();
1841    }
1842
1843    // Sort values for Q-Q plot (all data if not sampled, or sampled data if >= threshold)
1844    values.sort_by(f64::total_cmp);
1845    let sorted_sample_values = values.clone();
1846    let actual_sample_size = sorted_sample_values.len();
1847
1848    // Compute distribution characteristics
1849    let (sw_stat, sw_pvalue) = if values.len() >= 3 {
1850        approximate_shapiro_wilk(&values)
1851    } else {
1852        (None, None)
1853    };
1854    let coefficient_of_variation = if numeric_stats.mean != 0.0 {
1855        numeric_stats.std / numeric_stats.mean.abs()
1856    } else {
1857        0.0
1858    };
1859
1860    let mode = compute_mode(&values);
1861
1862    let characteristics = DistributionCharacteristics {
1863        shapiro_wilk_stat: sw_stat,
1864        shapiro_wilk_pvalue: sw_pvalue,
1865        skewness: numeric_stats.skewness,
1866        kurtosis: numeric_stats.kurtosis,
1867        mean: numeric_stats.mean,
1868        median: numeric_stats.median,
1869        std_dev: numeric_stats.std,
1870        variance: numeric_stats.std * numeric_stats.std,
1871        coefficient_of_variation,
1872        mode,
1873    };
1874
1875    let fit_quality = dist_info.fit_quality.unwrap_or(dist_info.confidence);
1876    let qq = dist_info
1877        .fits
1878        .iter()
1879        .filter_map(|(family, outcome)| {
1880            let test = outcome.test()?;
1881            Some((
1882                *family,
1883                crate::distribution_fit::qq_quantiles(&test.fitted, sorted_sample_values.len()),
1884            ))
1885        })
1886        .collect();
1887
1888    let outliers = compute_outlier_analysis(&finite_values(series), numeric_stats);
1889
1890    let percentiles = PercentileBreakdown {
1891        p1: numeric_stats
1892            .percentiles
1893            .get(&1)
1894            .copied()
1895            .unwrap_or(f64::NAN),
1896        p5: numeric_stats
1897            .percentiles
1898            .get(&5)
1899            .copied()
1900            .unwrap_or(f64::NAN),
1901        p25: numeric_stats.q25,
1902        p50: numeric_stats.median,
1903        p75: numeric_stats.q75,
1904        p95: numeric_stats
1905            .percentiles
1906            .get(&95)
1907            .copied()
1908            .unwrap_or(f64::NAN),
1909        p99: numeric_stats
1910            .percentiles
1911            .get(&99)
1912            .copied()
1913            .unwrap_or(f64::NAN),
1914    };
1915
1916    DistributionAnalysis {
1917        column_name: column_name.to_string(),
1918        distribution_type: dist_info.distribution_type,
1919        confidence: dist_info.confidence,
1920        fit_quality,
1921        characteristics,
1922        outliers,
1923        percentiles,
1924        sorted_sample_values,
1925        is_sampled,
1926        sample_size: actual_sample_size,
1927        fits: dist_info.fits.clone(),
1928        qq,
1929    }
1930}
1931
1932fn compute_mode(values: &[f64]) -> Option<f64> {
1933    if values.is_empty() {
1934        return None;
1935    }
1936
1937    // Bin values and find most frequent bin
1938    let min = values.iter().fold(f64::INFINITY, |a, &b| a.min(b));
1939    let max = values.iter().fold(f64::NEG_INFINITY, |a, &b| a.max(b));
1940    let range = max - min;
1941
1942    if range == 0.0 {
1943        return Some(min);
1944    }
1945
1946    let bins = 50.min(values.len());
1947    let mut bin_counts = vec![0; bins];
1948    let mut bin_sums = vec![0.0; bins];
1949
1950    for &v in values {
1951        let bin = (((v - min) / range) * (bins - 1) as f64) as usize;
1952        let bin = bin.min(bins - 1);
1953        bin_counts[bin] += 1;
1954        bin_sums[bin] += v;
1955    }
1956
1957    // Find bin with maximum count
1958    let max_bin = bin_counts
1959        .iter()
1960        .enumerate()
1961        .max_by_key(|&(_, &count)| count)
1962        .map(|(idx, _)| idx);
1963
1964    max_bin.map(|idx| bin_sums[idx] / bin_counts[idx] as f64)
1965}
1966
1967/// The standard normal quantile. See [`crate::distribution_fit::normal_quantile`].
1968pub(crate) fn normal_quantile(p: f64) -> f64 {
1969    crate::distribution_fit::normal_quantile(p)
1970}
1971
1972// CDF (Cumulative Distribution Function) implementations for histogram theoretical probabilities
1973fn normal_cdf(x: f64, mean: f64, std: f64) -> f64 {
1974    if std <= 0.0 {
1975        return if x < mean { 0.0 } else { 1.0 };
1976    }
1977    crate::distribution_fit::normal_cdf((x - mean) / std)
1978}
1979
1980fn lognormal_cdf(x: f64, mu: f64, sigma: f64) -> f64 {
1981    if x <= 0.0 {
1982        return 0.0;
1983    }
1984    if sigma <= 0.0 {
1985        return if x < mu.exp() { 0.0 } else { 1.0 };
1986    }
1987    // Lognormal: CDF(x) = Normal CDF of ln(x) with parameters mu, sigma
1988    normal_cdf(x.ln(), mu, sigma)
1989}
1990
1991fn exponential_cdf(x: f64, lambda: f64) -> f64 {
1992    if x < 0.0 {
1993        return 0.0;
1994    }
1995    if lambda <= 0.0 {
1996        return if x < 0.0 { 0.0 } else { 1.0 };
1997    }
1998    // Exponential CDF: 1 - exp(-lambda * x)
1999    1.0 - (-lambda * x).exp()
2000}
2001
2002fn powerlaw_cdf(x: f64, xmin: f64, alpha: f64) -> f64 {
2003    if x < xmin {
2004        return 0.0;
2005    }
2006    if alpha <= 1.0 {
2007        return if x >= xmin { 1.0 } else { 0.0 };
2008    }
2009    // Power law CDF: 1 - (x/xmin)^(-alpha + 1) for x >= xmin
2010    // Valid for alpha > 1
2011    if alpha <= 1.0 || xmin <= 0.0 {
2012        return if x >= xmin { 1.0 } else { 0.0 };
2013    }
2014    1.0 - (x / xmin).powf(-alpha + 1.0)
2015}
2016
2017/// The regularized incomplete beta function `I_x(a, b)`, by its continued fraction
2018/// (Lentz's method), flipped to the side where the fraction converges fast.
2019fn regularized_incomplete_beta(x: f64, a: f64, b: f64) -> f64 {
2020    if x <= 0.0 {
2021        return 0.0;
2022    }
2023    if x >= 1.0 {
2024        return 1.0;
2025    }
2026    let ln_gamma = crate::distribution_fit::ln_gamma;
2027    let ln_front = ln_gamma(a + b) - ln_gamma(a) - ln_gamma(b) + a * x.ln() + b * (1.0 - x).ln();
2028    let front = ln_front.exp();
2029    if x < (a + 1.0) / (a + b + 2.0) {
2030        front * beta_continued_fraction(x, a, b) / a
2031    } else {
2032        1.0 - front * beta_continued_fraction(1.0 - x, b, a) / b
2033    }
2034}
2035
2036fn beta_continued_fraction(x: f64, a: f64, b: f64) -> f64 {
2037    const TINY: f64 = 1e-300;
2038    let mut c = 1.0;
2039    let mut d = 1.0 - (a + b) * x / (a + 1.0);
2040    if d.abs() < TINY {
2041        d = TINY;
2042    }
2043    d = 1.0 / d;
2044    let mut h = d;
2045    for m in 1..=300 {
2046        let m = m as f64;
2047        let m2 = 2.0 * m;
2048        let even = m * (b - m) * x / ((a + m2 - 1.0) * (a + m2));
2049        d = 1.0 + even * d;
2050        if d.abs() < TINY {
2051            d = TINY;
2052        }
2053        c = 1.0 + even / c;
2054        if c.abs() < TINY {
2055            c = TINY;
2056        }
2057        d = 1.0 / d;
2058        h *= d * c;
2059        let odd = -(a + m) * (a + b + m) * x / ((a + m2) * (a + m2 + 1.0));
2060        d = 1.0 + odd * d;
2061        if d.abs() < TINY {
2062            d = TINY;
2063        }
2064        c = 1.0 + odd / c;
2065        if c.abs() < TINY {
2066            c = TINY;
2067        }
2068        d = 1.0 / d;
2069        let step = d * c;
2070        h *= step;
2071        if (step - 1.0).abs() < 1e-12 {
2072            break;
2073        }
2074    }
2075    h
2076}
2077
2078// Beta distribution CDF: the regularized incomplete beta function.
2079fn beta_cdf(x: f64, alpha: f64, beta: f64) -> f64 {
2080    if alpha <= 0.0 || beta <= 0.0 {
2081        return 0.0;
2082    }
2083    regularized_incomplete_beta(x, alpha, beta)
2084}
2085
2086// Gamma distribution CDF (requires incomplete gamma function approximation)
2087pub(crate) fn gamma_cdf(x: f64, shape: f64, scale: f64) -> f64 {
2088    if x <= 0.0 {
2089        return 0.0;
2090    }
2091    if shape <= 0.0 || scale <= 0.0 {
2092        return 0.0;
2093    }
2094    // Gamma CDF uses incomplete gamma function
2095    // For large shape, use normal approximation
2096    if shape > 30.0 {
2097        let mean = shape * scale;
2098        let variance = shape * scale * scale;
2099        if variance > 0.0 {
2100            normal_cdf(x, mean, variance.sqrt())
2101        } else if x < mean {
2102            0.0
2103        } else {
2104            1.0
2105        }
2106    } else {
2107        // Series approximation for incomplete gamma: P(x, k) = gamma(k, x) / Gamma(k)
2108        // Simplified approximation for small shape
2109        let z = x / scale;
2110        let sum: f64 = (0..(shape as usize * 10).min(100))
2111            .map(|n| {
2112                if (n as f64) < shape {
2113                    (-z).exp() * z.powi(n as i32) / (1..=n).map(|i| i as f64).product::<f64>()
2114                } else {
2115                    0.0
2116                }
2117            })
2118            .sum();
2119        (1.0 - sum).clamp(0.0, 1.0)
2120    }
2121}
2122
2123// Chi-squared distribution CDF (special case of Gamma with shape = df/2, scale = 2)
2124fn chi_squared_cdf(x: f64, df: f64) -> f64 {
2125    if x <= 0.0 {
2126        return 0.0;
2127    }
2128    if df <= 0.0 {
2129        return 0.0;
2130    }
2131    gamma_cdf(x, df / 2.0, 2.0)
2132}
2133
2134// Student's t distribution CDF (approximation)
2135fn students_t_cdf(x: f64, df: f64) -> f64 {
2136    if df <= 0.0 {
2137        return 0.5; // Invalid, return median
2138    }
2139    // Exact, through the incomplete beta: the tails are what tell a t from a normal,
2140    // and a scaled normal standing in for it had none.
2141    let tail = 0.5 * regularized_incomplete_beta(df / (df + x * x), df / 2.0, 0.5);
2142    if x >= 0.0 { 1.0 - tail } else { tail }
2143}
2144
2145// Poisson CDF (discrete, but return as continuous approximation)
2146fn poisson_cdf(x: f64, lambda: f64) -> f64 {
2147    if x < 0.0 {
2148        return 0.0;
2149    }
2150    if lambda <= 0.0 {
2151        return if x >= 0.0 { 1.0 } else { 0.0 };
2152    }
2153    // For large lambda, use normal approximation
2154    if lambda > 20.0 {
2155        normal_cdf(x, lambda, lambda.sqrt())
2156    } else {
2157        // Sum Poisson PMF from 0 to floor(x)
2158        let k_max = x.floor() as usize;
2159        let mut cdf = 0.0;
2160        let mut factorial = 1.0;
2161        for k in 0..=k_max.min(100) {
2162            if k > 0 {
2163                factorial *= k as f64;
2164            }
2165            let ln_pmf = (k as f64) * lambda.ln() - lambda - factorial.ln();
2166            let pmf = ln_pmf.exp();
2167            cdf += pmf;
2168            if cdf > 1.0 {
2169                break;
2170            }
2171        }
2172        cdf.min(1.0)
2173    }
2174}
2175
2176// Bernoulli CDF (discrete, p = probability of success)
2177fn bernoulli_cdf(x: f64, p: f64) -> f64 {
2178    if x < 0.0 {
2179        return 0.0;
2180    }
2181    if x >= 1.0 {
2182        return 1.0;
2183    }
2184    if p < 0.0 {
2185        return 0.0;
2186    }
2187    if p > 1.0 {
2188        return 1.0;
2189    }
2190    1.0 - p // CDF(x) = 0 for x < 0, 1-p for 0 <= x < 1, 1 for x >= 1
2191}
2192
2193// Binomial coefficient helper
2194fn binomial_coeff(n: usize, k: usize) -> f64 {
2195    if k > n {
2196        0.0
2197    } else if k == 0 || k == n {
2198        1.0
2199    } else {
2200        let k = k.min(n - k); // Use symmetry
2201        (1..=k).map(|i| (n - k + i) as f64 / i as f64).product()
2202    }
2203}
2204
2205// Binomial CDF (discrete)
2206fn binomial_cdf(x: f64, n: usize, p: f64) -> f64 {
2207    if x < 0.0 {
2208        return 0.0;
2209    }
2210    if p <= 0.0 {
2211        return if x >= n as f64 { 1.0 } else { 0.0 };
2212    }
2213    if p >= 1.0 {
2214        return if x >= 0.0 { 1.0 } else { 0.0 };
2215    }
2216    // For large n, use normal approximation
2217    if n > 50 {
2218        let mean = n as f64 * p;
2219        let variance = n as f64 * p * (1.0 - p);
2220        if variance > 0.0 {
2221            normal_cdf(x + 0.5, mean, variance.sqrt()) // Continuity correction
2222        } else if x < mean {
2223            0.0
2224        } else {
2225            1.0
2226        }
2227    } else {
2228        // Sum binomial PMF
2229        let k_max = x.floor() as usize;
2230        let mut cdf = 0.0;
2231        for k in 0..=k_max.min(n) {
2232            let coeff = binomial_coeff(n, k);
2233            let pmf = coeff * p.powi(k as i32) * (1.0 - p).powi((n - k) as i32);
2234            cdf += pmf;
2235        }
2236        cdf.min(1.0)
2237    }
2238}
2239
2240// Geometric CDF (discrete, number of failures before first success)
2241fn geometric_cdf(x: f64, p: f64) -> f64 {
2242    if x < 0.0 {
2243        return 0.0;
2244    }
2245    if p <= 0.0 || p >= 1.0 {
2246        return if x >= 0.0 && p >= 1.0 { 1.0 } else { 0.0 };
2247    }
2248
2249    // Geometric CDF: 1 - (1-p)^(k+1) for k failures
2250    // Use log-space to avoid numerical underflow: (1-p)^(k+1) = exp((k+1) * ln(1-p))
2251    // But cap k aggressively: beyond k=50, CDF is essentially 1.0 for most p values
2252    let k = x.floor().min(50.0); // Aggressive cap at 50 (was 1000)
2253
2254    // For very small (1-p)^(k+1), we can approximate as 0
2255    let log_one_minus_p = (1.0 - p).ln();
2256    if log_one_minus_p.is_nan() || log_one_minus_p.is_infinite() {
2257        return if x >= 0.0 { 1.0 } else { 0.0 };
2258    }
2259
2260    // Calculate (k+1) * ln(1-p)
2261    let exponent = (k + 1.0) * log_one_minus_p;
2262
2263    // If exponent is very negative, (1-p)^(k+1) is essentially 0, so CDF ≈ 1.0
2264    if exponent < -50.0 {
2265        return 1.0;
2266    }
2267
2268    // Otherwise calculate normally using exp
2269    let one_minus_p_power = exponent.exp();
2270    let result = 1.0 - one_minus_p_power;
2271    result.clamp(0.0, 1.0)
2272}
2273
2274// Weibull distribution CDF
2275fn weibull_cdf(x: f64, shape: f64, scale: f64) -> f64 {
2276    if x <= 0.0 {
2277        return 0.0;
2278    }
2279    if shape <= 0.0 || scale <= 0.0 {
2280        return 0.0;
2281    }
2282    // Weibull CDF: 1 - exp(-(x/scale)^shape)
2283    1.0 - (-(x / scale).powf(shape)).exp()
2284}
2285
2286// Calculate theoretical probability in an interval [lower, upper] for a distribution
2287// Helper function for dense sampling of theoretical distribution
2288/// Calculates the probability that a value falls in [lower, upper] for the given distribution.
2289///
2290/// Uses the distribution's CDF to compute P(lower ≤ X < upper).
2291pub fn calculate_theoretical_probability_in_interval(
2292    dist: &DistributionAnalysis,
2293    dist_type: DistributionType,
2294    lower: f64,
2295    upper: f64,
2296) -> f64 {
2297    let mean = dist.characteristics.mean;
2298    let std = dist.characteristics.std_dev;
2299    let sorted_data = &dist.sorted_sample_values;
2300
2301    match dist_type {
2302        DistributionType::Normal => {
2303            let cdf_upper = normal_cdf(upper, mean, std);
2304            let cdf_lower = normal_cdf(lower, mean, std);
2305            cdf_upper - cdf_lower
2306        }
2307        DistributionType::LogNormal => {
2308            if sorted_data.is_empty() || !sorted_data.iter().all(|&v| v > 0.0) {
2309                0.0
2310            } else {
2311                let e_x = mean;
2312                let var_x = std * std;
2313                let sigma_sq = (1.0 + var_x / (e_x * e_x)).ln();
2314                let mu = e_x.ln() - sigma_sq / 2.0;
2315                let sigma = sigma_sq.sqrt();
2316
2317                if lower > 0.0 && upper > 0.0 {
2318                    let cdf_upper = lognormal_cdf(upper, mu, sigma);
2319                    let cdf_lower = lognormal_cdf(lower, mu, sigma);
2320                    cdf_upper - cdf_lower
2321                } else {
2322                    0.0
2323                }
2324            }
2325        }
2326        DistributionType::Uniform => {
2327            if sorted_data.is_empty() {
2328                0.0
2329            } else {
2330                let data_min = sorted_data[0];
2331                let data_max = sorted_data[sorted_data.len() - 1];
2332                let data_range = data_max - data_min;
2333                if data_range > 0.0 {
2334                    (upper - lower) / data_range
2335                } else {
2336                    0.0
2337                }
2338            }
2339        }
2340        DistributionType::Exponential if mean > 0.0 => {
2341            let lambda = 1.0 / mean;
2342            let cdf_upper = exponential_cdf(upper, lambda);
2343            let cdf_lower = exponential_cdf(lower, lambda);
2344            cdf_upper - cdf_lower
2345        }
2346        DistributionType::PowerLaw => {
2347            if sorted_data.is_empty() || !sorted_data.iter().any(|&v| v > 0.0) {
2348                0.0
2349            } else {
2350                let positive_values: Vec<f64> =
2351                    sorted_data.iter().filter(|&&v| v > 0.0).copied().collect();
2352                if positive_values.is_empty() {
2353                    0.0
2354                } else {
2355                    let xmin = positive_values[0];
2356                    let n_pos = positive_values.len();
2357                    if n_pos < 2 || xmin <= 0.0 {
2358                        0.0
2359                    } else {
2360                        let sum_log = positive_values
2361                            .iter()
2362                            .map(|&x| (x / xmin).ln())
2363                            .sum::<f64>();
2364                        if sum_log > 0.0 {
2365                            let alpha = 1.0 + (n_pos as f64) / sum_log;
2366                            let cdf_upper = powerlaw_cdf(upper, xmin, alpha);
2367                            let cdf_lower = powerlaw_cdf(lower, xmin, alpha);
2368                            cdf_upper - cdf_lower
2369                        } else {
2370                            0.0
2371                        }
2372                    }
2373                }
2374            }
2375        }
2376        DistributionType::Beta => {
2377            // Estimate parameters from mean and variance
2378            let mean_val = mean;
2379            let variance = std * std;
2380            if mean_val > 0.0 && mean_val < 1.0 && variance > 0.0 {
2381                let max_var = mean_val * (1.0 - mean_val);
2382                if variance < max_var {
2383                    let sum = mean_val * (1.0 - mean_val) / variance - 1.0;
2384                    let alpha = mean_val * sum;
2385                    let beta = (1.0 - mean_val) * sum;
2386                    if alpha > 0.0 && beta > 0.0 {
2387                        let cdf_upper = beta_cdf(upper, alpha, beta);
2388                        let cdf_lower = beta_cdf(lower, alpha, beta);
2389                        cdf_upper - cdf_lower
2390                    } else {
2391                        0.0
2392                    }
2393                } else {
2394                    0.0
2395                }
2396            } else {
2397                0.0
2398            }
2399        }
2400        DistributionType::Gamma if mean > 0.0 && std > 0.0 => {
2401            let variance = std * std;
2402            let shape = (mean * mean) / variance;
2403            let scale = variance / mean;
2404            if shape > 0.0 && scale > 0.0 {
2405                let cdf_upper = gamma_cdf(upper, shape, scale);
2406                let cdf_lower = gamma_cdf(lower, shape, scale);
2407                cdf_upper - cdf_lower
2408            } else {
2409                0.0
2410            }
2411        }
2412        DistributionType::ChiSquared => {
2413            // Chi-squared is gamma(df/2, 2)
2414            let df = mean; // For chi-squared, mean = df
2415            if df > 0.0 {
2416                let cdf_upper = chi_squared_cdf(upper, df);
2417                let cdf_lower = chi_squared_cdf(lower, df);
2418                cdf_upper - cdf_lower
2419            } else {
2420                0.0
2421            }
2422        }
2423        DistributionType::StudentsT => {
2424            // Estimate df from variance
2425            let variance = std * std;
2426            let df = if variance > 1.0 {
2427                2.0 * variance / (variance - 1.0)
2428            } else {
2429                30.0
2430            };
2431            let cdf_upper = students_t_cdf(upper, df);
2432            let cdf_lower = students_t_cdf(lower, df);
2433            cdf_upper - cdf_lower
2434        }
2435        DistributionType::Poisson => {
2436            let lambda = mean;
2437            if lambda > 0.0 {
2438                let cdf_upper = poisson_cdf(upper, lambda);
2439                let cdf_lower = poisson_cdf(lower, lambda);
2440                cdf_upper - cdf_lower
2441            } else {
2442                0.0
2443            }
2444        }
2445        DistributionType::Bernoulli => {
2446            let p = mean; // For Bernoulli, mean = p
2447            let cdf_upper = bernoulli_cdf(upper, p);
2448            let cdf_lower = bernoulli_cdf(lower, p);
2449            cdf_upper - cdf_lower
2450        }
2451        DistributionType::Binomial => {
2452            // Estimate n from data range
2453            let sorted_data = &dist.sorted_sample_values;
2454            if !sorted_data.is_empty() {
2455                let max_val = sorted_data[sorted_data.len() - 1];
2456                let n = max_val.floor() as usize;
2457                let p = if n > 0 { mean / n as f64 } else { 0.5 };
2458                if n > 0 && p > 0.0 && p < 1.0 {
2459                    let cdf_upper = binomial_cdf(upper, n, p);
2460                    let cdf_lower = binomial_cdf(lower, n, p);
2461                    cdf_upper - cdf_lower
2462                } else {
2463                    0.0
2464                }
2465            } else {
2466                0.0
2467            }
2468        }
2469        DistributionType::Geometric => {
2470            let mean_val = mean; // mean = (1-p)/p for geometric
2471            if mean_val > 0.0 {
2472                let p = 1.0 / (mean_val + 1.0);
2473                let cdf_upper = geometric_cdf(upper, p);
2474                let cdf_lower = geometric_cdf(lower, p);
2475                cdf_upper - cdf_lower
2476            } else {
2477                0.0
2478            }
2479        }
2480        DistributionType::Weibull if mean > 0.0 && std > 0.0 => {
2481            // Approximate shape from CV
2482            let cv = std / mean;
2483            let shape = if cv < 1.0 { 1.0 / cv } else { 1.0 };
2484            // Scale from mean
2485            let gamma_1_over_shape = 1.0 + 1.0 / shape; // Approximation
2486            let scale = mean / gamma_1_over_shape;
2487            if shape > 0.0 && scale > 0.0 {
2488                let cdf_upper = weibull_cdf(upper, shape, scale);
2489                let cdf_lower = weibull_cdf(lower, shape, scale);
2490                cdf_upper - cdf_lower
2491            } else {
2492                0.0
2493            }
2494        }
2495        _ => 0.0,
2496    }
2497}
2498
2499/// How many outlier examples an analysis keeps, most extreme first.
2500const OUTLIER_EXAMPLES: usize = 100;
2501
2502/// Outliers among every finite value of the column. The counts and the percentage
2503/// cover all of them; only the list of examples is cut to [`OUTLIER_EXAMPLES`].
2504fn compute_outlier_analysis(values: &[f64], numeric_stats: &NumericStatistics) -> OutlierAnalysis {
2505    let mut analysis = OutlierAnalysis {
2506        total_count: 0,
2507        percentage: 0.0,
2508        iqr_count: 0,
2509        zscore_count: 0,
2510        outlier_rows: Vec::new(),
2511    };
2512    let Some(test) = OutlierTest::new(values, numeric_stats.q25, numeric_stats.q75) else {
2513        return analysis;
2514    };
2515
2516    for (idx, &value) in values.iter().enumerate() {
2517        let iqr_position = test.iqr_position(value);
2518        let z_score = test.z_score(value);
2519        let detection_method = match (iqr_position.is_some(), z_score > Z_THRESHOLD) {
2520            (true, true) => OutlierMethod::Both,
2521            (true, false) => OutlierMethod::IQR,
2522            (false, true) => OutlierMethod::ZScore,
2523            (false, false) => continue,
2524        };
2525        analysis.total_count += 1;
2526        analysis.iqr_count += usize::from(iqr_position.is_some());
2527        analysis.zscore_count += usize::from(z_score > Z_THRESHOLD);
2528        analysis.outlier_rows.push(OutlierRow {
2529            row_index: idx,
2530            column_value: value,
2531            context_data: HashMap::new(),
2532            detection_method,
2533            z_score: Some(z_score),
2534            iqr_position,
2535        });
2536    }
2537
2538    analysis.percentage = analysis.total_count as f64 / values.len() as f64 * 100.0;
2539    // Most extreme first; a z-score is the distance from the mean in one scale. The
2540    // hundred are picked before sorting: a long tail has tens of thousands.
2541    let most_extreme = |a: &OutlierRow, b: &OutlierRow| {
2542        b.z_score
2543            .unwrap_or(0.0)
2544            .total_cmp(&a.z_score.unwrap_or(0.0))
2545    };
2546    if analysis.outlier_rows.len() > OUTLIER_EXAMPLES {
2547        analysis
2548            .outlier_rows
2549            .select_nth_unstable_by(OUTLIER_EXAMPLES, most_extreme);
2550        analysis.outlier_rows.truncate(OUTLIER_EXAMPLES);
2551    }
2552    analysis.outlier_rows.sort_by(most_extreme);
2553    analysis
2554}
2555
2556// Correlation matrix computation
2557/// Computes pairwise Pearson correlation matrix for all numeric columns.
2558///
2559/// Returns correlations, p-values, and sample sizes for each pair.
2560/// Requires at least 2 numeric columns.
2561pub fn compute_correlation_matrix(df: &DataFrame) -> Result<CorrelationMatrix> {
2562    let columns = df
2563        .schema()
2564        .iter()
2565        .filter(|(_, dtype)| is_numeric_type(dtype))
2566        .count()
2567        .max(1);
2568    let band = CORRELATION_SCRATCH_BYTES / (columns * std::mem::size_of::<f64>());
2569    correlation_matrix_in_bands(df, band.max(MIN_BAND_ROWS))
2570}
2571
2572/// Bytes of converted values a correlation matrix holds at once, beyond the sample
2573/// itself and the matrix: one band of the 100,000-row sample is the whole of it up
2574/// to 83 columns.
2575const CORRELATION_SCRATCH_BYTES: usize = 64 * 1024 * 1024;
2576
2577/// The fewest rows in a band, so a schema of thousands of columns is not read a few
2578/// rows at a time.
2579const MIN_BAND_ROWS: usize = 1024;
2580
2581/// Rows cast to floats at a time, so no cast is the size of a column.
2582const CAST_ROWS: usize = 16 * 1024;
2583
2584/// [`compute_correlation_matrix`] converting `band` rows of every column at a time.
2585fn correlation_matrix_in_bands(df: &DataFrame, band: usize) -> Result<CorrelationMatrix> {
2586    // Get all numeric columns
2587    let schema = df.schema();
2588    let numeric_cols: Vec<String> = schema
2589        .iter()
2590        .filter(|(_, dtype)| is_numeric_type(dtype))
2591        .map(|(name, _)| name.to_string())
2592        .collect();
2593
2594    if numeric_cols.len() < 2 {
2595        return Err(color_eyre::eyre::eyre!(
2596            "Need at least 2 numeric columns for correlation matrix"
2597        ));
2598    }
2599
2600    let series = numeric_cols
2601        .iter()
2602        .map(|name| Ok(df.column(name)?.as_materialized_series()))
2603        .collect::<Result<Vec<_>>>()?;
2604
2605    let n = numeric_cols.len();
2606    let rows = df.height();
2607    let band = band.clamp(1, rows.max(1));
2608
2609    // Each column's mean first, then a band of rows of every column at a time, less
2610    // those means. Each pair's sums run on from one band to the next in row order:
2611    // the same sums as one pass over whole columns, with every column converted once
2612    // and no more than a band of them held.
2613    let mut shifts = vec![Shift::default(); n];
2614    across_threads(
2615        series.iter().zip(shifts.iter_mut()).collect(),
2616        |(series, shift)| *shift = Shift::new(series),
2617    );
2618    let mut sums: Vec<Vec<PairSums>> = (0..n)
2619        .map(|i| vec![PairSums::default(); n - i - 1])
2620        .collect();
2621    let mut bands: Vec<Vec<f64>> = (0..n).map(|_| Vec::with_capacity(band)).collect();
2622    for start in (0..rows).step_by(band) {
2623        let within = start..(start + band).min(rows);
2624        across_threads(
2625            series.iter().zip(&shifts).zip(bands.iter_mut()).collect(),
2626            |((series, shift), values)| shift.fill(series, within.clone(), values),
2627        );
2628        // Every pair is one pass over two columns; fifty columns are 1,225 pairs, so
2629        // the rows of the matrix are shared out across threads, interleaved to even
2630        // the load.
2631        let (bands, shifts) = (&bands, &shifts);
2632        across_threads(sums.iter_mut().enumerate().collect(), |(i, row)| {
2633            for (k, sums) in row.iter_mut().enumerate() {
2634                let j = i + 1 + k;
2635                let both = shifts[i].complete && shifts[j].complete;
2636                sums.add_pairs(&bands[i], &bands[j], both);
2637            }
2638        });
2639    }
2640
2641    let mut correlations = vec![vec![1.0; n]; n];
2642    let mut p_values = vec![vec![0.0; n]; n];
2643    let mut sample_sizes = vec![vec![0; n]; n];
2644    for (i, row) in sums.iter().enumerate() {
2645        for (k, sums) in row.iter().enumerate() {
2646            let j = i + 1 + k;
2647            let sample_size = sums.count;
2648            sample_sizes[i][j] = sample_size;
2649            sample_sizes[j][i] = sample_size;
2650            // Fewer than three pairs say nothing.
2651            let correlation = if sample_size < 3 {
2652                f64::NAN
2653            } else {
2654                sums.correlation()
2655            };
2656            correlations[i][j] = correlation;
2657            correlations[j][i] = correlation;
2658            if !correlation.is_nan() {
2659                let p_value = compute_correlation_p_value(correlation, sample_size);
2660                p_values[i][j] = p_value;
2661                p_values[j][i] = p_value;
2662            }
2663        }
2664    }
2665
2666    let ranked = (rows.saturating_mul(n) <= RANK_VALUES)
2667        .then(|| rank_correlation_matrix(&series, &sample_sizes));
2668    let (rank_correlations, rank_p_values) = ranked.unzip();
2669    Ok(CorrelationMatrix {
2670        columns: numeric_cols,
2671        correlations,
2672        p_values: Some(p_values),
2673        sample_sizes,
2674        rank_correlations,
2675        rank_p_values,
2676    })
2677}
2678
2679/// Marks a row with no finite value in [`Ranked::ranks`].
2680const NO_RANK: u32 = u32::MAX;
2681
2682/// One column's ranks, for Spearman's ρ.
2683struct Ranked {
2684    /// Twice each finite value's average rank (from 1) among the column's finite
2685    /// values, so a tie's half rank stays whole; [`NO_RANK`] where there is none.
2686    /// Equal values share a rank, so the ranks also tell ties apart.
2687    ranks: Vec<u32>,
2688    /// The rows with a finite value, in order of value.
2689    order: Vec<u32>,
2690    /// Every row has a finite value.
2691    complete: bool,
2692}
2693
2694impl Ranked {
2695    fn new(series: &Series) -> Option<Self> {
2696        let rows = series.len();
2697        // Doubled ranks reach twice the rows.
2698        if rows >= (NO_RANK / 2) as usize {
2699            return None;
2700        }
2701        let mut values = Vec::with_capacity(rows);
2702        let mut row = 0u32;
2703        for_each_float(series, 0..rows, |v| {
2704            if let Some(v) = v.filter(|v| v.is_finite()) {
2705                values.push((v, row));
2706            }
2707            row += 1;
2708        });
2709        values.sort_unstable_by(|a, b| a.0.total_cmp(&b.0));
2710        let mut ranks = vec![NO_RANK; rows];
2711        let mut start = 0;
2712        while start < values.len() {
2713            // `==` so that -0 and 0 tie, which total_cmp sorts side by side.
2714            let end = start
2715                + values[start..]
2716                    .iter()
2717                    .take_while(|(v, _)| *v == values[start].0)
2718                    .count();
2719            let doubled = (start + 1 + end) as u32;
2720            for &(_, row) in &values[start..end] {
2721                ranks[row as usize] = doubled;
2722            }
2723            start = end;
2724        }
2725        Some(Self {
2726            complete: values.len() == rows,
2727            order: values.into_iter().map(|(_, row)| row).collect(),
2728            ranks,
2729        })
2730    }
2731
2732    /// `out` becomes this column's doubled ranks among the rows where `other` also
2733    /// has a value, [`NO_RANK`] elsewhere; the number of those rows is returned.
2734    /// Walking the rows in order of value, a run of equal global ranks is a run of
2735    /// equal values. `kept` is scratch, reused from pair to pair.
2736    fn ranks_beside(&self, other: &Ranked, out: &mut Vec<u32>, kept: &mut Vec<u32>) -> usize {
2737        out.clear();
2738        out.resize(self.ranks.len(), NO_RANK);
2739        kept.clear();
2740        kept.extend(
2741            self.order
2742                .iter()
2743                .copied()
2744                .filter(|&row| other.ranks[row as usize] != NO_RANK),
2745        );
2746        let mut start = 0;
2747        while start < kept.len() {
2748            let tie = self.ranks[kept[start] as usize];
2749            let end = start
2750                + kept[start..]
2751                    .iter()
2752                    .take_while(|&&row| self.ranks[row as usize] == tie)
2753                    .count();
2754            let doubled = (start + 1 + end) as u32;
2755            for &row in &kept[start..end] {
2756                out[row as usize] = doubled;
2757            }
2758            start = end;
2759        }
2760        kept.len()
2761    }
2762}
2763
2764/// Spearman's ρ for every pair of `series`, with its p-values: Pearson's r of the
2765/// ranks, each pair ranked over the rows where both hold a finite value, as Pearson
2766/// pairs them. Where neither column misses a value the column's own ranks are the
2767/// pair's, and no pair needs ranking again. Holds four bytes of rank and four of
2768/// order per value, about what the sample's own values take.
2769fn rank_correlation_matrix(
2770    series: &[&Series],
2771    sample_sizes: &[Vec<usize>],
2772) -> (Vec<Vec<f64>>, Vec<Vec<f64>>) {
2773    let n = series.len();
2774    let mut ranked: Vec<Option<Ranked>> = (0..n).map(|_| None).collect();
2775    across_threads(
2776        series.iter().zip(ranked.iter_mut()).collect(),
2777        |(series, ranked)| *ranked = Ranked::new(series),
2778    );
2779    let mut rows: Vec<Vec<f64>> = (0..n).map(|i| vec![f64::NAN; n - i - 1]).collect();
2780    let ranked = &ranked;
2781    across_threads(rows.iter_mut().enumerate().collect(), |(i, row)| {
2782        // Ranks as doubled u32s, half the size of floats, held once per row of the
2783        // matrix rather than once per pair.
2784        let (mut a, mut b, mut kept) = (Vec::new(), Vec::new(), Vec::new());
2785        for (k, rho) in row.iter_mut().enumerate() {
2786            let (Some(x), Some(y)) = (&ranked[i], &ranked[i + 1 + k]) else {
2787                continue;
2788            };
2789            let mut sums = PairSums::default();
2790            if x.complete && y.complete {
2791                // Ranks centered on their mean, which is the row count plus one.
2792                let mean = (x.ranks.len() + 1) as f64;
2793                for (&rx, &ry) in x.ranks.iter().zip(&y.ranks) {
2794                    sums.add(rx as f64 - mean, ry as f64 - mean);
2795                }
2796            } else {
2797                let pairs = x.ranks_beside(y, &mut a, &mut kept);
2798                y.ranks_beside(x, &mut b, &mut kept);
2799                let mean = (pairs + 1) as f64;
2800                for (&rx, &ry) in a.iter().zip(&b) {
2801                    if rx != NO_RANK {
2802                        sums.add(rx as f64 - mean, ry as f64 - mean);
2803                    }
2804                }
2805            }
2806            if sums.count >= 3 {
2807                *rho = sums.correlation();
2808            }
2809        }
2810    });
2811    let mut rho = vec![vec![1.0; n]; n];
2812    let mut p_values = vec![vec![0.0; n]; n];
2813    for (i, row) in rows.iter().enumerate() {
2814        for (k, &r) in row.iter().enumerate() {
2815            let j = i + 1 + k;
2816            rho[i][j] = r;
2817            rho[j][i] = r;
2818            if !r.is_nan() {
2819                // The same t approximation as Pearson's, over the same pairs.
2820                let p = compute_correlation_p_value(r, sample_sizes[i][j]);
2821                p_values[i][j] = p;
2822                p_values[j][i] = p;
2823            }
2824        }
2825    }
2826    (rho, p_values)
2827}
2828
2829/// Runs `work` on every item, the items dealt out across threads in turn. A worker's
2830/// panic is raised again here rather than leaving its items undone.
2831fn across_threads<T: Send>(items: Vec<T>, work: impl Fn(T) + Sync) {
2832    let threads = std::thread::available_parallelism()
2833        .map_or(1, usize::from)
2834        .min(items.len())
2835        .max(1);
2836    let mut shares: Vec<Vec<T>> = (0..threads).map(|_| Vec::new()).collect();
2837    for (k, item) in items.into_iter().enumerate() {
2838        shares[k % threads].push(item);
2839    }
2840    let work = &work;
2841    std::thread::scope(|scope| {
2842        let handles: Vec<_> = shares
2843            .into_iter()
2844            .map(|share| scope.spawn(move || share.into_iter().for_each(work)))
2845            .collect();
2846        for handle in handles {
2847            if let Err(panic) = handle.join() {
2848                std::panic::resume_unwind(panic);
2849            }
2850        }
2851    });
2852}
2853
2854/// `len` values of `series` from `start`, as floats.
2855fn float_piece(series: &Series, start: usize, len: usize) -> Option<Float64Chunked> {
2856    let piece = series
2857        .slice(start as i64, len)
2858        .cast(&DataType::Float64)
2859        .ok()?;
2860    piece.f64().ok().cloned()
2861}
2862
2863/// The values of `series` in `rows` as floats, None where null, cast [`CAST_ROWS`]
2864/// at a time.
2865fn for_each_float(series: &Series, rows: Range<usize>, mut f: impl FnMut(Option<f64>)) {
2866    for start in rows.clone().step_by(CAST_ROWS) {
2867        let len = CAST_ROWS.min(rows.end - start);
2868        match float_piece(series, start, len) {
2869            Some(floats) => floats.iter().for_each(&mut f),
2870            None => (0..len).for_each(|_| f(None)),
2871        }
2872    }
2873}
2874
2875/// A numeric column's mean over its finite values, taken off each value before it
2876/// is correlated: centered, the one-pass sums below stay exact enough.
2877#[derive(Clone, Copy, Default)]
2878struct Shift {
2879    mean: f64,
2880    /// No value is missing, so every row pairs.
2881    complete: bool,
2882}
2883
2884impl Shift {
2885    fn new(series: &Series) -> Self {
2886        let (mut sum, mut count) = (0.0, 0usize);
2887        for_each_float(series, 0..series.len(), |v| {
2888            if let Some(v) = v.filter(|v| v.is_finite()) {
2889                sum += v;
2890                count += 1;
2891            }
2892        });
2893        Self {
2894            mean: if count > 0 { sum / count as f64 } else { 0.0 },
2895            complete: count == series.len(),
2896        }
2897    }
2898
2899    /// `values` becomes the column's `rows` less the mean, NaN where a value is null
2900    /// or not finite.
2901    fn fill(&self, series: &Series, rows: Range<usize>, values: &mut Vec<f64>) {
2902        values.clear();
2903        for_each_float(series, rows, |v| {
2904            values.push(
2905                v.filter(|v| v.is_finite())
2906                    .map_or(f64::NAN, |v| v - self.mean),
2907            );
2908        });
2909    }
2910}
2911
2912/// One pass of sums over paired values, each less a shift near its mean. The
2913/// spreads about the pairs' own means follow exactly whatever the shift; a shift
2914/// near the mean keeps them clear of rounding.
2915#[derive(Clone, Copy, Default)]
2916struct PairSums {
2917    count: usize,
2918    x: f64,
2919    y: f64,
2920    xx: f64,
2921    yy: f64,
2922    xy: f64,
2923}
2924
2925impl PairSums {
2926    fn add(&mut self, v1: f64, v2: f64) {
2927        self.count += 1;
2928        self.x += v1;
2929        self.y += v2;
2930        self.xx += v1 * v1;
2931        self.yy += v2 * v2;
2932        self.xy += v1 * v2;
2933    }
2934
2935    /// Adds the rows where both centered columns have a value: with `both` complete,
2936    /// every row.
2937    fn add_pairs(&mut self, a: &[f64], b: &[f64], both: bool) {
2938        // Summed in a copy, which the loop can keep in registers.
2939        let mut sums = *self;
2940        for (&v1, &v2) in a.iter().zip(b) {
2941            if !both && (v1.is_nan() || v2.is_nan()) {
2942                continue;
2943            }
2944            sums.add(v1, v2);
2945        }
2946        *self = sums;
2947    }
2948
2949    /// The sums of squares and of products about the pairs' means.
2950    fn spreads(&self) -> (f64, f64, f64) {
2951        let n = self.count as f64;
2952        (
2953            self.xx - self.x * self.x / n,
2954            self.yy - self.y * self.y / n,
2955            self.xy - self.x * self.y / n,
2956        )
2957    }
2958
2959    /// NaN for fewer than two pairs or a column of one value.
2960    fn correlation(&self) -> f64 {
2961        let (sxx, syy, sxy) = self.spreads();
2962        // Measured against the sums of squares: what a column with one value leaves
2963        // behind is rounding, not spread.
2964        if self.count < 2 || sxx <= self.xx * 1e-12 || syy <= self.yy * 1e-12 {
2965            return f64::NAN;
2966        }
2967        (sxy / (sxx * syy).sqrt()).clamp(-1.0, 1.0)
2968    }
2969}
2970
2971/// The two-sided p-value of Pearson's r over `n` pairs: Student's t with `n - 2`
2972/// degrees of freedom. With `t² = r²·df / (1 - r²)`, both tails together are
2973/// `I_x(df/2, 1/2)` at `x = df / (df + t²) = 1 - r²`, taken directly so that a small p
2974/// is not lost to `1 - cdf`.
2975fn compute_correlation_p_value(correlation: f64, n: usize) -> f64 {
2976    if n < 3 || correlation.is_nan() {
2977        return 1.0;
2978    }
2979    if correlation.abs() >= 1.0 {
2980        return 0.0;
2981    }
2982    let df = (n - 2) as f64;
2983    regularized_incomplete_beta(1.0 - correlation * correlation, df / 2.0, 0.5).clamp(0.0, 1.0)
2984}
2985
2986/// Computes correlation statistics for a pair of columns.
2987///
2988/// Returns Pearson correlation coefficient, p-value, covariance, and sample size,
2989/// over the rows where both columns hold a finite value. Requires at least 3 such
2990/// rows. Two passes over the columns as they are, a rough mean first and then the
2991/// sums about it, as the matrix's: no list of the pairs is built.
2992pub fn compute_correlation_pair(
2993    df: &DataFrame,
2994    col1_name: &str,
2995    col2_name: &str,
2996) -> Result<CorrelationPair> {
2997    let series1 = df.column(col1_name)?.as_materialized_series();
2998    let series2 = df.column(col2_name)?.as_materialized_series();
2999
3000    let mut sample_size = 0usize;
3001    let (mut sum1, mut sum2) = (0.0, 0.0);
3002    let (mut min1, mut max1, mut min2, mut max2) = (f64::NAN, f64::NAN, f64::NAN, f64::NAN);
3003    for_each_finite_pair(series1, series2, |v1, v2| {
3004        sample_size += 1;
3005        sum1 += v1;
3006        sum2 += v2;
3007        (min1, max1) = (min1.min(v1), max1.max(v1));
3008        (min2, max2) = (min2.min(v2), max2.max(v2));
3009    });
3010    if sample_size < 3 {
3011        return Err(color_eyre::eyre::eyre!("Not enough data for correlation"));
3012    }
3013    let n = sample_size as f64;
3014    let (shift1, shift2) = (sum1 / n, sum2 / n);
3015    let mut sums = PairSums::default();
3016    for_each_finite_pair(series1, series2, |v1, v2| {
3017        sums.add(v1 - shift1, v2 - shift2)
3018    });
3019    let (sxx, syy, sxy) = sums.spreads();
3020
3021    // A column with one value has no correlation with anything: undefined, not 0,
3022    // which reads as a finding.
3023    let correlation = sums.correlation();
3024    let p_value = Some(compute_correlation_p_value(correlation, sample_size));
3025    // Rounding can leave a constant's sum of squares a hair below zero.
3026    let stats = |mean: f64, squares: f64, min: f64, max: f64| ColumnStats {
3027        mean,
3028        std: (squares.max(0.0) / (n - 1.0)).sqrt(),
3029        min,
3030        max,
3031    };
3032
3033    Ok(CorrelationPair {
3034        column1: col1_name.to_string(),
3035        column2: col2_name.to_string(),
3036        correlation,
3037        p_value,
3038        sample_size,
3039        covariance: sxy / (n - 1.0),
3040        r_squared: correlation * correlation,
3041        stats1: stats(shift1 + sums.x / n, sxx, min1, max1),
3042        stats2: stats(shift2 + sums.y / n, syy, min2, max2),
3043    })
3044}
3045
3046/// The rows where both columns hold a finite value, in order, cast [`CAST_ROWS`] at
3047/// a time.
3048fn for_each_finite_pair(a: &Series, b: &Series, mut f: impl FnMut(f64, f64)) {
3049    let rows = a.len().min(b.len());
3050    for start in (0..rows).step_by(CAST_ROWS) {
3051        let len = CAST_ROWS.min(rows - start);
3052        let (Some(a), Some(b)) = (float_piece(a, start, len), float_piece(b, start, len)) else {
3053            continue;
3054        };
3055        for pair in a.iter().zip(b.iter()) {
3056            if let (Some(v1), Some(v2)) = pair
3057                && v1.is_finite()
3058                && v2.is_finite()
3059            {
3060                f(v1, v2);
3061            }
3062        }
3063    }
3064}
3065
3066#[cfg(test)]
3067mod sampling_tests {
3068    use super::*;
3069
3070    /// `n` rows of `id` and a `score` that climbs with it, as a Parquet file: a head
3071    /// sample of it is biased, and a spread one is not.
3072    fn climbing(dir: &std::path::Path, n: i64) -> LazyFrame {
3073        let ids: Vec<i64> = (0..n).collect();
3074        let scores: Vec<f64> = ids.iter().map(|i| *i as f64).collect();
3075        let mut df = df!("id" => ids, "score" => scores).unwrap();
3076        let path = dir.join("climbing.parquet");
3077        ParquetWriter::new(std::fs::File::create(&path).unwrap())
3078            .with_row_group_size(Some(1_000))
3079            .finish(&mut df)
3080            .unwrap();
3081        LazyFrame::scan_parquet(PlRefPath::try_from_path(&path).unwrap(), Default::default())
3082            .unwrap()
3083    }
3084
3085    fn mean(df: &DataFrame) -> f64 {
3086        df.column("score")
3087            .unwrap()
3088            .as_materialized_series()
3089            .mean()
3090            .unwrap()
3091    }
3092
3093    #[test]
3094    fn a_small_table_is_read_whole() {
3095        let dir = tempfile::tempdir().unwrap();
3096        let rows = analysis_rows(&climbing(dir.path(), 500), Some(1_000), None, 1, false).unwrap();
3097        assert_eq!(rows.df.height(), 500);
3098        assert_eq!(rows.total_rows, 500);
3099        assert_eq!(rows.sample_size, None);
3100    }
3101
3102    #[test]
3103    fn a_parquet_scan_is_sampled_in_blocks_across_all_of_it() {
3104        let dir = tempfile::tempdir().unwrap();
3105        let lf = climbing(dir.path(), 100_000);
3106        assert!(slices_reach_into_the_scan(&lf), "a Parquet scan seeks");
3107        let rows = analysis_rows(&lf, Some(5_000), None, 7, false).unwrap();
3108        assert_eq!(rows.df.height(), 5_000);
3109        assert_eq!(rows.total_rows, 100_000);
3110        assert_eq!(rows.sample_size, Some(5_000));
3111        // The table's mean is 49,999.5; a head sample's would be 2,499.5.
3112        assert!(
3113            (mean(&rows.df) - 49_999.5).abs() < 2_500.0,
3114            "{}",
3115            mean(&rows.df)
3116        );
3117        // In table order, and reaching its last stretch.
3118        let ids = rows.df.column("id").unwrap().i64().unwrap();
3119        assert!(ids.into_no_null_iter().is_sorted());
3120        assert!(ids.max().unwrap() > 95_000);
3121        // The same seed, the same sample; another seed, another.
3122        let again = analysis_rows(&lf, Some(5_000), None, 7, false).unwrap();
3123        assert!(rows.df.equals(&again.df));
3124        let other = analysis_rows(&lf, Some(5_000), None, 8, false).unwrap();
3125        assert!(!rows.df.equals(&other.df));
3126    }
3127
3128    /// Seeded runs, and a table read whole because it is under twice the sample, say
3129    /// where each kept row sat: the row's id, in a table whose id is its position.
3130    /// The runs see too few rows to count; the whole read counts what it read.
3131    #[test]
3132    fn a_block_sample_says_where_its_rows_sat() {
3133        let tens = (col("id") / lit(10_000i64)).alias("tens");
3134        for (rows, n) in [(100_000, 5_000), (10_000, 6_000)] {
3135            let dir = tempfile::tempdir().unwrap();
3136            let lf = climbing(dir.path(), rows);
3137            let read =
3138                sample_rows_counting(&lf, Some(n), None, 7, false, None, Some(&tens)).unwrap();
3139            let ids: Vec<IdxSize> = read
3140                .rows
3141                .df
3142                .column("id")
3143                .unwrap()
3144                .i64()
3145                .unwrap()
3146                .into_no_null_iter()
3147                .map(|id| id as IdxSize)
3148                .collect();
3149            assert_eq!(ids.len(), n);
3150            assert_eq!(ids, read.positions, "{rows} rows, n={n}");
3151            if rows < 2 * n as i64 {
3152                assert_eq!(
3153                    read.counted,
3154                    Some(crate::sampling::Counted::Totals(
3155                        [(Some("0".to_string()), 10_000)].into()
3156                    ))
3157                );
3158            } else {
3159                assert_eq!(read.counted, None, "runs see too few rows to count");
3160            }
3161        }
3162    }
3163
3164    /// Odd and small sizes: exactly `n` distinct rows, reaching the end of the table.
3165    /// Runs rounded up and then cut to `n` once dropped the last blocks; runs wider
3166    /// than their stretch overlapped.
3167    #[test]
3168    fn a_block_sample_of_any_size_is_n_distinct_rows_across_the_table() {
3169        let dir = tempfile::tempdir().unwrap();
3170        let lf = climbing(dir.path(), 10_000);
3171        for (n, seed) in [
3172            (60, 1),
3173            (101, 2),
3174            (1_234, 3),
3175            (4_999, 4),
3176            (5_001, 5),
3177            (9_999, 6),
3178        ] {
3179            let rows = analysis_rows(&lf, Some(n), None, seed, false).unwrap();
3180            let ids: Vec<i64> = rows
3181                .df
3182                .column("id")
3183                .unwrap()
3184                .i64()
3185                .unwrap()
3186                .into_no_null_iter()
3187                .collect();
3188            assert_eq!(ids.len(), n, "n={n}");
3189            let unique: std::collections::HashSet<_> = ids.iter().collect();
3190            assert_eq!(unique.len(), n, "no duplicates at n={n}");
3191            assert!(ids.is_sorted(), "table order at n={n}");
3192            assert!(
3193                *ids.last().unwrap() > 9_000,
3194                "reaches the end at n={n}: {ids:?}"
3195            );
3196        }
3197    }
3198
3199    #[test]
3200    fn only_a_single_file_scan_is_sampled_in_blocks() {
3201        let dir = tempfile::tempdir().unwrap();
3202        let one = climbing(dir.path(), 1_000);
3203        // A column stubbed above the scan, as binary columns are: still seekable.
3204        let stubbed = one
3205            .clone()
3206            .select([col("id"), lit(NULL).cast(DataType::Binary).alias("blob")]);
3207        assert!(slices_reach_into_the_scan(&stubbed));
3208        // Two files: each slice would open the footers of those before it.
3209        let two = concat([one.clone(), one.clone()], UnionArgs::default()).unwrap();
3210        assert!(!slices_reach_into_the_scan(&two));
3211        let glob = LazyFrame::scan_parquet(
3212            PlRefPath::try_from_path(&dir.path().join("*.parquet")).unwrap(),
3213            Default::default(),
3214        )
3215        .unwrap();
3216        std::fs::copy(
3217            dir.path().join("climbing.parquet"),
3218            dir.path().join("again.parquet"),
3219        )
3220        .unwrap();
3221        assert!(!slices_reach_into_the_scan(&glob), "two files in one scan");
3222        assert!(!slices_reach_into_the_scan(
3223            &one.clone().sort(["id"], Default::default())
3224        ));
3225    }
3226
3227    #[test]
3228    fn a_filtered_view_is_sampled_in_one_uniform_pass() {
3229        let dir = tempfile::tempdir().unwrap();
3230        let lf = climbing(dir.path(), 100_000).filter(col("id").gt_eq(lit(50_000)));
3231        assert!(
3232            !slices_reach_into_the_scan(&lf),
3233            "a slice of a filter reads what is ahead of it"
3234        );
3235        let rows = analysis_rows(&lf, Some(5_000), None, 7, false).unwrap();
3236        assert_eq!(rows.df.height(), 5_000);
3237        assert_eq!(rows.total_rows, 50_000, "the pass counts as it goes");
3238        assert!(
3239            (mean(&rows.df) - 74_999.5).abs() < 1_500.0,
3240            "{}",
3241            mean(&rows.df)
3242        );
3243        assert!(
3244            rows.df.column(SAMPLE_POSITION).is_err(),
3245            "the position column does not leak"
3246        );
3247        let again = analysis_rows(&lf, Some(5_000), None, 7, false).unwrap();
3248        assert!(rows.df.equals(&again.df), "seeded");
3249    }
3250
3251    #[test]
3252    fn no_sample_size_reads_every_row() {
3253        let dir = tempfile::tempdir().unwrap();
3254        let rows = analysis_rows(&climbing(dir.path(), 20_000), None, None, 1, false).unwrap();
3255        assert_eq!(rows.df.height(), 20_000);
3256        assert_eq!(rows.sample_size, None);
3257    }
3258
3259    #[test]
3260    fn a_matrix_in_bands_is_the_matrix_in_one() {
3261        // Columns with nulls in different rows, one with none and one of integers,
3262        // converted a band of rows at a time as well as all at once: the same matrix,
3263        // bit for bit, whether or not the bands divide the rows.
3264        let rows = 1_000;
3265        let mut columns: Vec<Column> = (0..6)
3266            .map(|c| {
3267                let v: Vec<Option<f64>> = (0..rows)
3268                    .map(|r| {
3269                        (c == 0 || (r + c) % (5 + c) != 0)
3270                            .then(|| ((r * (c + 1)) as f64 * 0.37).sin() + (r % 13) as f64)
3271                    })
3272                    .collect();
3273                Series::new(format!("c{c}").into(), v).into()
3274            })
3275            .collect();
3276        let integers: Vec<Option<i64>> = (0..rows)
3277            .map(|r| (r % 9 != 4).then_some((r * r % 101) as i64))
3278            .collect();
3279        columns.push(Series::new("i".into(), integers).into());
3280        let df = DataFrame::new(rows, columns).unwrap();
3281        let whole = compute_correlation_matrix(&df).unwrap();
3282        let bits = |m: &CorrelationMatrix| {
3283            m.correlations
3284                .iter()
3285                .flatten()
3286                .chain(m.p_values.iter().flatten().flatten())
3287                .map(|v| v.to_bits())
3288                .collect::<Vec<_>>()
3289        };
3290        assert!(whole.correlations[0][6].is_finite());
3291        for band in [1, 2, 7, 333, 999, 1_000, 5_000] {
3292            let bands = correlation_matrix_in_bands(&df, band).unwrap();
3293            assert_eq!(bits(&bands), bits(&whole), "{band} rows at a time");
3294            assert_eq!(bands.sample_sizes, whole.sample_sizes);
3295        }
3296    }
3297
3298    #[test]
3299    fn a_constant_column_correlates_with_nothing() {
3300        let df = df!(
3301            "year" => vec![2020.0f64; 50],
3302            "value" => (0..50).map(|i| i as f64).collect::<Vec<_>>(),
3303            "double" => (0..50).map(|i| 2.0 * i as f64).collect::<Vec<_>>(),
3304            // Its mean is not exactly 0.1, so centering leaves rounding behind.
3305            "tenth" => vec![0.1f64; 50]
3306        )
3307        .unwrap();
3308        let matrix = compute_correlation_matrix(&df).unwrap();
3309        assert!(matrix.correlations[0][1].is_nan(), "undefined, not 0");
3310        assert!(
3311            matrix.correlations[3][1].is_nan(),
3312            "{}",
3313            matrix.correlations[3][1]
3314        );
3315        assert!((matrix.correlations[1][2] - 1.0).abs() < 1e-9);
3316    }
3317
3318    /// Spearman's ρ by the book: rank both columns over the pair's rows, average
3319    /// ranks for ties, then Pearson's r of the ranks.
3320    fn spearman_by_hand(x: &[Option<f64>], y: &[Option<f64>]) -> f64 {
3321        let pairs: Vec<(f64, f64)> = x
3322            .iter()
3323            .zip(y)
3324            .filter_map(|(a, b)| Some((a.filter(|v| v.is_finite())?, b.filter(|v| v.is_finite())?)))
3325            .collect();
3326        let rank = |values: Vec<f64>| -> Vec<f64> {
3327            values
3328                .iter()
3329                .map(|v| {
3330                    let below = values.iter().filter(|w| *w < v).count() as f64;
3331                    let equal = values.iter().filter(|w| *w == v).count() as f64;
3332                    below + (equal + 1.0) / 2.0
3333                })
3334                .collect()
3335        };
3336        let a = rank(pairs.iter().map(|p| p.0).collect());
3337        let b = rank(pairs.iter().map(|p| p.1).collect());
3338        let n = a.len() as f64;
3339        let (ma, mb) = (a.iter().sum::<f64>() / n, b.iter().sum::<f64>() / n);
3340        let sab: f64 = a.iter().zip(&b).map(|(x, y)| (x - ma) * (y - mb)).sum();
3341        let saa: f64 = a.iter().map(|x| (x - ma).powi(2)).sum();
3342        let sbb: f64 = b.iter().map(|y| (y - mb).powi(2)).sum();
3343        sab / (saa * sbb).sqrt()
3344    }
3345
3346    #[test]
3347    fn spearman_ranks_each_pair_over_the_rows_both_hold() {
3348        let rows = 400;
3349        let x: Vec<Option<f64>> = (0..rows)
3350            .map(|r| (r % 7 != 3).then(|| ((r * 37 % 101) as f64 * 0.5).floor()))
3351            .collect();
3352        // Monotone in x but far from a line, with its own gaps and a NaN.
3353        let y: Vec<Option<f64>> = (0..rows)
3354            .map(|r| match r {
3355                _ if r % 11 == 5 => None,
3356                17 => Some(f64::NAN),
3357                _ => x[r].map(|v| v.powi(5)),
3358            })
3359            .collect();
3360        let z: Vec<Option<f64>> = (0..rows).map(|r| Some(((r * 13) % 29) as f64)).collect();
3361        let w: Vec<Option<f64>> = (0..rows)
3362            .map(|r| Some(((r * 7) % 31) as f64 - (r % 3) as f64))
3363            .collect();
3364        let df = DataFrame::new(
3365            rows,
3366            vec![
3367                Series::new("x".into(), &x).into(),
3368                Series::new("y".into(), &y).into(),
3369                Series::new("z".into(), &z).into(),
3370                Series::new("w".into(), &w).into(),
3371            ],
3372        )
3373        .unwrap();
3374        let m = compute_correlation_matrix(&df).unwrap();
3375        let columns = [&x, &y, &z, &w];
3376        for i in 0..4 {
3377            assert_eq!(m.coefficient(CorrelationMethod::Spearman, i, i), 1.0);
3378            for j in (0..4).filter(|&j| j != i) {
3379                let expected = spearman_by_hand(columns[i], columns[j]);
3380                let got = m.coefficient(CorrelationMethod::Spearman, i, j);
3381                assert!(
3382                    (got - expected).abs() < 1e-12,
3383                    "{i},{j}: {got} vs {expected}"
3384                );
3385            }
3386        }
3387        assert!((m.coefficient(CorrelationMethod::Spearman, 0, 1) - 1.0).abs() < 1e-12);
3388        assert!(m.coefficient(CorrelationMethod::Pearson, 0, 1) < 0.99);
3389        assert!(m.p_value(CorrelationMethod::Spearman, 2, 3).is_some());
3390    }
3391
3392    #[test]
3393    fn spearman_of_a_constant_column_is_undefined() {
3394        let df = df!(
3395            "year" => vec![2020.0f64; 50],
3396            "value" => (0..50).map(|i| i as f64).collect::<Vec<_>>()
3397        )
3398        .unwrap();
3399        let matrix = compute_correlation_matrix(&df).unwrap();
3400        assert!(
3401            matrix
3402                .coefficient(CorrelationMethod::Spearman, 0, 1)
3403                .is_nan()
3404        );
3405    }
3406
3407    #[test]
3408    fn the_incomplete_beta_and_the_t_cdf_are_exact() {
3409        // I_0.5(2, 3) = 11/16.
3410        let i = regularized_incomplete_beta(0.5, 2.0, 3.0);
3411        assert!((i - 0.6875).abs() < 1e-6, "{i}");
3412        // The 97.5th percentile of t with 5 degrees of freedom is 2.5706.
3413        assert!((students_t_cdf(2.5706, 5.0) - 0.975).abs() < 1e-4);
3414        assert!((students_t_cdf(-2.5706, 5.0) - 0.025).abs() < 1e-4);
3415        assert!((students_t_cdf(0.0, 5.0) - 0.5).abs() < 1e-12);
3416        // Beta(1, 1) is the uniform.
3417        assert!((beta_cdf(0.3, 1.0, 1.0) - 0.3).abs() < 1e-6);
3418    }
3419
3420    #[test]
3421    fn a_constant_column_is_constant_and_a_hopeless_one_has_no_clear_fit() {
3422        let constant = Series::new("c".into(), vec![2020.0f64; 500]);
3423        let info = infer_distribution(&constant, &constant, 500, false);
3424        assert_eq!(info.distribution_type, DistributionType::Constant);
3425
3426        // Two far-apart clusters of non-integers: every candidate is rejected.
3427        let values: Vec<f64> = (0..2_000)
3428            .map(|i| {
3429                let jitter = (i % 97) as f64 * 0.013;
3430                if i % 2 == 0 {
3431                    -1_000.3 + jitter
3432                } else {
3433                    1_000.7 + jitter
3434                }
3435            })
3436            .collect();
3437        let bimodal = Series::new("b".into(), values);
3438        let info = infer_distribution(&bimodal, &bimodal, 2_000, false);
3439        assert_eq!(info.distribution_type, DistributionType::Unknown);
3440        assert_eq!(info.distribution_type.to_string(), "No clear fit");
3441
3442        // Integers, half of them negative: no count distribution, whatever the
3443        // non-negative half looks like on its own.
3444        let values: Vec<f64> = (0..2_000)
3445            .map(|i| {
3446                if i % 2 == 0 {
3447                    -1_000.0 + (i % 7) as f64
3448                } else {
3449                    1_000.0 + (i % 5) as f64
3450                }
3451            })
3452            .collect();
3453        let integers = Series::new("i".into(), values);
3454        let info = infer_distribution(&integers, &integers, 2_000, false);
3455        assert!(
3456            !matches!(
3457                info.distribution_type,
3458                DistributionType::Binomial
3459                    | DistributionType::Poisson
3460                    | DistributionType::Geometric
3461            ),
3462            "{:?}",
3463            info.distribution_type
3464        );
3465    }
3466}
3467
3468#[cfg(test)]
3469mod normality_tests {
3470    use super::*;
3471
3472    /// Against SciPy's `2 * t.sf(|t|, n - 2)`. A normal CDF standing in for the t,
3473    /// and a tanh standing in for the normal, gave r = 0.01 over 100,000 pairs p = 0.023
3474    /// where it is 0.0016.
3475    #[test]
3476    fn correlation_p_values_are_students_t() {
3477        let cases = [
3478            (0.5, 10, 0.14111328125000006),
3479            (-0.2, 50, 0.16375308124541754),
3480            (0.3, 30, 0.10724594805795436),
3481            (0.9, 5, 0.03738607346849862),
3482            (0.1, 100, 0.32221736303061954),
3483            (0.02, 20_000, 0.004676184609440329),
3484            (0.01, 100_000, 0.0015651897452783157),
3485            (0.005, 1_000_000, 5.732288112893878e-7),
3486            (0.1, 10_000, 1.1970504236520445e-23),
3487        ];
3488        for (r, n, expected) in cases {
3489            let p = compute_correlation_p_value(r, n);
3490            assert!(
3491                (p - expected).abs() <= 1e-9 * expected,
3492                "r {r}, n {n}: {p} against {expected}"
3493            );
3494        }
3495        assert_eq!(compute_correlation_p_value(1.0, 10), 0.0);
3496        assert_eq!(compute_correlation_p_value(0.0, 10), 1.0);
3497    }
3498
3499    /// Shapiro-Francia's p-value is a p-value: a normal sample passes, and a price
3500    /// series of two regimes, W' = 0.929 over 2,590 values, does not.
3501    #[test]
3502    fn the_normality_p_value_is_a_p_value() {
3503        assert!(shapiro_francia_pvalue(0.929, 2_590).unwrap() < 1e-10);
3504        assert!(shapiro_francia_pvalue(0.9995, 2_590).unwrap() > 0.05);
3505        assert_eq!(shapiro_francia_pvalue(0.99, 4), None);
3506        let normal: Vec<f64> = (1..=500)
3507            .map(|i| normal_quantile(i as f64 / 501.0))
3508            .collect();
3509        let (_, p) = approximate_shapiro_wilk(&normal);
3510        assert!(p.unwrap() > 0.5, "{p:?}");
3511    }
3512
3513    /// Royston's approximation is calibrated: normal samples fall below 0.05 about one
3514    /// time in twenty, and skewed ones nearly always.
3515    #[test]
3516    fn the_normality_p_value_is_calibrated() {
3517        let mut rng = crate::distribution_fit::Rng::new(2_026);
3518        let mut sample = |skewed: bool| -> Vec<f64> {
3519            (0..100)
3520                .map(|_| {
3521                    let z = rng.normal();
3522                    if skewed { z.exp() } else { z }
3523                })
3524                .collect()
3525        };
3526        let below = |values: Vec<f64>| approximate_shapiro_wilk(&values).1.unwrap() < 0.05;
3527        let false_alarms = (0..400).filter(|_| below(sample(false))).count();
3528        assert!(
3529            (8..=36).contains(&false_alarms),
3530            "{false_alarms} of 400 normal samples below 0.05"
3531        );
3532        let caught = (0..100).filter(|_| below(sample(true))).count();
3533        assert!(caught >= 95, "{caught} of 100 log-normal samples caught");
3534    }
3535
3536    /// NaN and infinities never reach a fit: one NaN left a `partial_cmp` sort out of
3537    /// order and folded the Q-Q plot.
3538    #[test]
3539    fn non_finite_values_are_left_out() {
3540        let series = Series::new("x".into(), &[3.0, f64::NAN, 1.0, f64::INFINITY, 2.0]);
3541        assert_eq!(get_numeric_values_as_f64(&series), vec![3.0, 1.0, 2.0]);
3542        assert_eq!(finite_values(&series), vec![3.0, 1.0, 2.0]);
3543        let integers = Series::new("i".into(), &[Some(4i16), None, Some(-2)]);
3544        assert_eq!(finite_values(&integers), vec![4.0, -2.0]);
3545    }
3546
3547    /// One value throughout, even one a float cannot hold exactly, has no skew; and a
3548    /// symmetric set has none either.
3549    #[test]
3550    fn a_constant_has_no_shape() {
3551        assert_eq!(skewness_and_kurtosis(&[0.1; 50]), (0.0, 3.0));
3552        assert_eq!(skewness_and_kurtosis(&[1.0, 2.0]), (0.0, 3.0));
3553        let (skewness, _) = skewness_and_kurtosis(&[1.0, 2.0, 3.0, 4.0, 5.0]);
3554        assert!(skewness.abs() < 1e-12);
3555    }
3556}
3557
3558#[cfg(test)]
3559pub(crate) mod describe_tests {
3560    use super::*;
3561
3562    /// A frame with one column of each temporal type, a zoned datetime too, five
3563    /// values and a null each.
3564    pub(crate) fn temporal_frame() -> DataFrame {
3565        let day = 20_089i32; // 2025-01-01
3566        let dates = Series::new(
3567            "day".into(),
3568            &[
3569                Some(day + 4),
3570                Some(day),
3571                None,
3572                Some(day + 2),
3573                Some(day + 1),
3574                Some(day + 3),
3575            ],
3576        )
3577        .cast(&DataType::Date)
3578        .unwrap();
3579        let hour = 3_600_000_000i64; // microseconds
3580        let start = 1_735_678_075_000_000i64; // 2024-12-31 20:47:55
3581        let pickups = Series::new(
3582            "pickup".into(),
3583            &[
3584                Some(start + 4 * hour),
3585                Some(start),
3586                None,
3587                Some(start + 2 * hour),
3588                Some(start + hour),
3589                Some(start + 3 * hour),
3590            ],
3591        )
3592        .cast(&DataType::Datetime(TimeUnit::Microseconds, None))
3593        .unwrap();
3594        let second = 1_000_000_000i64; // nanoseconds
3595        let times = Series::new(
3596            "at".into(),
3597            &[
3598                Some(9 * 3600 * second + 40 * second),
3599                Some(9 * 3600 * second),
3600                None,
3601                Some(9 * 3600 * second + 20 * second),
3602                Some(9 * 3600 * second + 10 * second),
3603                Some(9 * 3600 * second + 30 * second),
3604            ],
3605        )
3606        .cast(&DataType::Time)
3607        .unwrap();
3608        // The same instants in New York, in milliseconds: the zone survives the cast back.
3609        let local = Series::new(
3610            "local".into(),
3611            &[
3612                Some(start / 1000 + 4 * hour / 1000),
3613                Some(start / 1000),
3614                None,
3615                Some(start / 1000 + 2 * hour / 1000),
3616                Some(start / 1000 + hour / 1000),
3617                Some(start / 1000 + 3 * hour / 1000),
3618            ],
3619        )
3620        .cast(&DataType::Datetime(
3621            TimeUnit::Milliseconds,
3622            TimeZone::opt_try_new(Some("America/New_York")).unwrap(),
3623        ))
3624        .unwrap();
3625        let minute = 60_000i64; // milliseconds
3626        let waits = Series::new(
3627            "wait".into(),
3628            &[
3629                Some(5 * minute),
3630                Some(minute),
3631                None,
3632                Some(3 * minute),
3633                Some(2 * minute),
3634                Some(4 * minute),
3635            ],
3636        )
3637        .cast(&DataType::Duration(TimeUnit::Milliseconds))
3638        .unwrap();
3639        DataFrame::new_infer_height(vec![
3640            dates.into(),
3641            pickups.into(),
3642            times.into(),
3643            local.into(),
3644            waits.into(),
3645        ])
3646        .unwrap()
3647    }
3648
3649    #[test]
3650    fn describe_gives_dates_and_times_their_range_in_their_own_format() {
3651        let df = temporal_frame();
3652        let every_row = crate::sampling::Sample {
3653            method: crate::sampling::SampleMethod::EveryRow,
3654            ..crate::sampling::Sample::default()
3655        };
3656        let lazy = compute_describe_from_lazy(&df.clone().lazy(), Some(6), &every_row, false)
3657            .unwrap()
3658            .column_statistics;
3659        let schema = df.schema().clone();
3660        let sampled = compute_describe_single_aggregation(&df, &schema, 6, None, 0, false)
3661            .unwrap()
3662            .column_statistics;
3663        let expected = [
3664            [
3665                "2025-01-03",
3666                "2025-01-01",
3667                "2025-01-02",
3668                "2025-01-03",
3669                "2025-01-04",
3670                "2025-01-05",
3671            ],
3672            [
3673                "2024-12-31 22:47:55",
3674                "2024-12-31 20:47:55",
3675                "2024-12-31 21:47:55",
3676                "2024-12-31 22:47:55",
3677                "2024-12-31 23:47:55",
3678                "2025-01-01 00:47:55",
3679            ],
3680            [
3681                "09:00:20", "09:00:00", "09:00:10", "09:00:20", "09:00:30", "09:00:40",
3682            ],
3683            [
3684                "2024-12-31 17:47:55 EST",
3685                "2024-12-31 15:47:55 EST",
3686                "2024-12-31 16:47:55 EST",
3687                "2024-12-31 17:47:55 EST",
3688                "2024-12-31 18:47:55 EST",
3689                "2024-12-31 19:47:55 EST",
3690            ],
3691            ["3m", "1m", "2m", "3m", "4m", "5m"],
3692        ];
3693        for stats in [&lazy, &sampled] {
3694            assert_eq!(stats.len(), expected.len());
3695            for (column, want) in stats.iter().zip(expected) {
3696                assert!(column.numeric_stats.is_none(), "{}", column.name);
3697                assert_eq!(column.null_count, 1);
3698                let t = column.temporal_stats.as_ref().expect("temporal stats");
3699                let got = [&t.mean, &t.min, &t.q25, &t.median, &t.q75, &t.max]
3700                    .map(|v| v.clone().unwrap_or_default());
3701                assert_eq!(got, want.map(String::from), "{}", column.name);
3702            }
3703        }
3704    }
3705
3706    #[test]
3707    fn describe_of_an_all_null_datetime_is_empty() {
3708        let empty = Series::new("never".into(), &[None::<i64>, None])
3709            .cast(&DataType::Datetime(TimeUnit::Microseconds, None))
3710            .unwrap();
3711        let df = DataFrame::new_infer_height(vec![empty.into()]).unwrap();
3712        let schema = df.schema().clone();
3713        let stats = compute_describe_single_aggregation(&df, &schema, 2, None, 0, false)
3714            .unwrap()
3715            .column_statistics;
3716        let t = stats[0].temporal_stats.as_ref().expect("temporal stats");
3717        assert!(
3718            [&t.mean, &t.min, &t.q25, &t.median, &t.q75, &t.max]
3719                .iter()
3720                .all(|v| v.is_none())
3721        );
3722    }
3723}
3724
3725#[cfg(all(test, feature = "streaming"))]
3726mod streaming_guard_tests {
3727    use super::*;
3728
3729    fn frame() -> LazyFrame {
3730        df!("i" => (0..1_000i64).collect::<Vec<_>>(), "j" => (0..1_000i64).map(|v| v % 7).collect::<Vec<_>>())
3731            .unwrap()
3732            .lazy()
3733            .with_columns([
3734                col("i").cast(DataType::Decimal(38, 2)).alias("d"),
3735                col("i").cast(DataType::Int128).alias("w"),
3736            ])
3737    }
3738
3739    struct Anonymous;
3740
3741    impl AnonymousScan for Anonymous {
3742        fn as_any(&self) -> &dyn std::any::Any {
3743            self
3744        }
3745
3746        fn schema(&self, _: Option<usize>) -> PolarsResult<SchemaRef> {
3747            Ok(Arc::new(Schema::from_iter([Field::new(
3748                "a".into(),
3749                DataType::Int64,
3750            )])))
3751        }
3752
3753        fn scan(&self, _: AnonymousScanArgs) -> PolarsResult<DataFrame> {
3754            df!("a" => [1i64, 2, 3])
3755        }
3756    }
3757
3758    /// An anonymous scan (a SQLite table) has no streaming implementation in Polars
3759    /// 0.55, so a query over one runs on the in-memory engine whatever is asked.
3760    #[test]
3761    fn an_anonymous_scan_stays_off_the_streaming_engine() {
3762        let lf = LazyFrame::anonymous_scan(Arc::new(Anonymous), Default::default())
3763            .unwrap()
3764            .filter(col("a").gt(lit(1i64)));
3765        assert!(!may_stream(&lf, true));
3766        assert!(may_stream(&frame(), true));
3767        assert_eq!(collect_lazy(lf, true).unwrap().height(), 2);
3768    }
3769
3770    /// The shapes Polars 0.55's streaming top-k panics on go to the in-memory engine
3771    /// and read; the ones it runs stay on the streaming engine.
3772    #[test]
3773    fn only_a_top_k_by_one_wide_key_leaves_the_streaming_engine() {
3774        let lf = frame();
3775        let desc = SortMultipleOptions::default().with_order_descending(true);
3776        for key in ["d", "w"] {
3777            let top = [
3778                lf.clone().sort([key], desc.clone()).slice(0, 3),
3779                lf.clone()
3780                    .filter(col("j").eq(lit(1)))
3781                    .sort([key], Default::default())
3782                    .limit(3),
3783                lf.clone()
3784                    .group_by([col(key)])
3785                    .agg([len()])
3786                    .sort([key], desc.clone())
3787                    .slice(0, 3),
3788            ];
3789            for query in top {
3790                assert!(sorts_by_one_wide_key(&query), "{key}");
3791                assert_eq!(collect_lazy(query, true).unwrap().height(), 3);
3792            }
3793            let streams = [
3794                lf.clone().sort([key], desc.clone()),
3795                lf.clone().sort([key], desc.clone()).slice(10, 3),
3796                lf.clone().sort([key, "j"], desc.clone()).slice(0, 3),
3797                lf.clone()
3798                    .sort([key], desc.clone().with_maintain_order(true))
3799                    .slice(0, 3),
3800                lf.clone().sort(["i"], desc.clone()).slice(0, 3),
3801            ];
3802            for query in streams {
3803                assert!(!sorts_by_one_wide_key(&query), "{key}");
3804                query
3805                    .collect_with_engine(Engine::Streaming)
3806                    .unwrap()
3807                    .unwrap_single();
3808            }
3809        }
3810    }
3811}