1use polars::prelude::*;
17use std::sync::Arc;
18
19pub const INDEX: &str = "__datui_row";
21
22pub const MAX_ROWS: usize = IdxSize::MAX as usize;
24
25pub trait RowSource: Send + Sync + 'static {
27 fn height(&self) -> usize;
29
30 fn schema(&self) -> SchemaRef;
31
32 fn decode(&self, column: usize, index: &IdxCa) -> PolarsResult<Column>;
37}
38
39pub fn lazy<S: RowSource>(source: &Arc<S>) -> LazyFrame {
41 lazy_with_height(source, source.height())
42}
43
44pub fn lazy_with_height<S: RowSource>(source: &Arc<S>, height: usize) -> LazyFrame {
47 frame(source, height, false)
48}
49
50pub fn lazy_numbered<S: RowSource>(source: &Arc<S>, height: usize) -> LazyFrame {
54 frame(source, height, true)
55}
56
57fn frame<S: RowSource>(source: &Arc<S>, height: usize, numbered: bool) -> LazyFrame {
58 let height = DataFrame::empty_with_height(height.min(MAX_ROWS)).lazy();
59 frame_over(source, height, numbered)
60}
61
62pub fn lazy_numbered_over<S: RowSource>(source: &Arc<S>, height: LazyFrame) -> LazyFrame {
65 frame_over(source, height, true)
66}
67
68fn frame_over<S: RowSource>(source: &Arc<S>, height: LazyFrame, numbered: bool) -> LazyFrame {
69 let base = height.with_row_index(INDEX, None);
70 let mut exprs: Vec<Expr> = source
71 .schema()
72 .iter()
73 .enumerate()
74 .map(|(column, (name, dtype))| {
75 let source = Arc::clone(source);
76 let field = Field::new(name.clone(), dtype.clone());
77 col(INDEX)
78 .map(
79 move |c: Column| source.decode(column, c.as_materialized_series().idx()?),
80 move |_, _| Ok(field.clone()),
81 )
82 .alias(name.clone())
83 })
84 .collect();
85 if numbered {
86 exprs.push(col(INDEX));
87 }
88 base.select(exprs)
89}
90
91pub fn checked(index: &IdxCa, rows: usize) -> PolarsResult<std::borrow::Cow<'_, [IdxSize]>> {
94 polars_ensure!(
95 index.null_count() == 0,
96 ComputeError: "a row index has a missing row"
97 );
98 if let Some(max) = index.max() {
99 polars_ensure!(
100 (max as usize) < rows,
101 OutOfBounds: "row {max} is past the {rows} rows on hand"
102 );
103 }
104 Ok(match index.cont_slice() {
105 Ok(rows) => std::borrow::Cow::Borrowed(rows),
106 Err(_) => std::borrow::Cow::Owned(index.into_no_null_iter().collect()),
107 })
108}
109
110#[cfg(test)]
111mod tests {
112 use super::*;
113
114 struct Counting {
116 rows: usize,
117 decoded: std::sync::atomic::AtomicUsize,
118 }
119
120 impl RowSource for Counting {
121 fn height(&self) -> usize {
122 self.rows
123 }
124
125 fn schema(&self) -> SchemaRef {
126 Arc::new(Schema::from_iter([
127 Field::new("a".into(), DataType::Int64),
128 Field::new("b".into(), DataType::Int64),
129 ]))
130 }
131
132 fn decode(&self, column: usize, index: &IdxCa) -> PolarsResult<Column> {
133 let rows = checked(index, self.rows)?;
134 self.decoded
135 .fetch_add(rows.len(), std::sync::atomic::Ordering::Relaxed);
136 Ok(Int64Chunked::from_iter_values(
137 PlSmallStr::EMPTY,
138 rows.iter().map(|&i| i as i64 * 10 + column as i64),
139 )
140 .into_column())
141 }
142 }
143
144 fn counting(rows: usize) -> Arc<Counting> {
145 Arc::new(Counting {
146 rows,
147 decoded: Default::default(),
148 })
149 }
150
151 #[test]
152 fn a_slice_decodes_only_its_rows_and_columns() {
153 let source = counting(1_000_000);
154 let df = lazy(&source)
155 .select([col("b")])
156 .slice(999_998, 10)
157 .collect()
158 .unwrap();
159 assert_eq!(df.get_column_names(), ["b"]);
160 assert_eq!(
161 df.column("b").unwrap().i64().unwrap().to_vec(),
162 [Some(9_999_981), Some(9_999_991)]
163 );
164 assert_eq!(source.decoded.load(std::sync::atomic::Ordering::Relaxed), 2);
165 }
166
167 #[test]
168 fn filters_sorts_and_group_bys_run_on_the_streaming_engine() {
169 let source = counting(100_000);
170 let lf = lazy(&source);
171 let streaming = |lf: LazyFrame| crate::statistics::collect_lazy(lf, true).unwrap();
172 let n = streaming(
173 lf.clone()
174 .filter(col("a").gt(lit(500_000i64)))
175 .select([len()]),
176 );
177 assert_eq!(
178 n.column("len").unwrap().get(0).unwrap(),
179 AnyValue::UInt32(49_999)
180 );
181 let top = streaming(
182 lf.clone()
183 .sort(
184 ["a"],
185 SortMultipleOptions::default().with_order_descending(true),
186 )
187 .limit(2),
188 );
189 assert_eq!(
190 top.column("b").unwrap().i64().unwrap().to_vec(),
191 [Some(999_991), Some(999_981)]
192 );
193 let groups = streaming(
194 lf.group_by([(col("a") % lit(20i64)).alias("k")])
195 .agg([len()])
196 .sort(["k"], Default::default()),
197 );
198 assert_eq!(groups.height(), 2);
199 assert_eq!(
200 groups.column("len").unwrap().u32().unwrap().to_vec(),
201 [Some(50_000), Some(50_000)]
202 );
203 }
204
205 #[test]
206 fn a_row_past_the_end_or_missing_is_refused() {
207 let index = IdxCa::from_slice("i".into(), &[0, 5]);
208 assert!(checked(&index, 6).is_ok());
209 assert!(checked(&index, 5).is_err());
210 let missing = IdxCa::from_slice_options("i".into(), &[Some(0), None]);
211 assert!(checked(&missing, 6).is_err());
212 }
213}