1use std::sync::{Arc, Mutex};
15
16use color_eyre::Result;
17use color_eyre::eyre::Report;
18use polars::prelude::*;
19
20use crate::analysis::sampling::{CANCELLED, ReadWatch, Sample, SampleMethod};
21use crate::analysis::sampling::{sample_rank, stream_batches};
22
23pub const MEMORY_SETTING: &str = "analysis.sample_memory_limit";
25
26const POSITION: &str = "__datui_table_sample_position";
28
29#[derive(Default)]
32pub struct SampleRows {
33 inner: Mutex<Inner>,
34}
35
36#[derive(Default)]
37struct Inner {
38 chunks: Vec<(u64, DataFrame)>,
40 taken: usize,
42 rows: usize,
43 bytes: usize,
44 stopped: Option<String>,
46}
47
48impl SampleRows {
49 fn lock(&self) -> std::sync::MutexGuard<'_, Inner> {
50 self.inner.lock().unwrap_or_else(|e| e.into_inner())
51 }
52
53 pub fn push(&self, key: u64, df: DataFrame) {
55 let mut inner = self.lock();
56 inner.rows += df.height();
57 inner.bytes += df.estimated_size();
58 inner.chunks.push((key, df));
59 }
60
61 pub fn take_new(&self) -> Vec<DataFrame> {
63 let mut inner = self.lock();
64 let from = inner.taken;
65 inner.taken = inner.chunks.len();
66 inner.chunks[from..]
67 .iter()
68 .map(|(_, df)| df.clone())
69 .collect()
70 }
71
72 pub fn rows(&self) -> usize {
74 self.lock().rows
75 }
76
77 pub fn bytes(&self) -> usize {
79 self.lock().bytes
80 }
81
82 pub fn stopped(&self) -> Option<String> {
84 self.lock().stopped.clone()
85 }
86
87 fn stop(&self, reason: String) {
88 self.lock().stopped = Some(reason);
89 }
90
91 pub fn take_in_source_order(&self) -> Result<Option<DataFrame>> {
94 let ordered = self.in_source_order()?;
95 let mut inner = self.lock();
96 inner.chunks.clear();
97 inner.taken = 0;
98 drop(inner);
99 Ok(ordered.map(|mut frame| {
100 frame.rechunk_mut_par();
101 frame
102 }))
103 }
104
105 pub fn in_source_order(&self) -> Result<Option<DataFrame>> {
108 let inner = self.lock();
109 let mut order: Vec<&(u64, DataFrame)> = inner.chunks.iter().collect();
110 order.sort_by_key(|(key, _)| *key);
112 let mut out: Option<DataFrame> = None;
113 for (_, df) in order {
114 match out.as_mut() {
115 Some(frame) => {
116 frame.vstack_mut(df)?;
117 }
118 None => out = Some(df.clone()),
119 }
120 }
121 Ok(out)
122 }
123}
124
125#[derive(Debug, Clone, Copy, PartialEq, Eq)]
127pub enum Limit {
128 Available,
130 Fixed(u64),
132 Off,
134}
135
136impl Limit {
137 pub fn of_setting(setting: Option<crate::config::ByteSize>) -> Self {
139 match setting.map(|size| size.bytes()) {
140 None => Self::Available,
141 Some(0) => Self::Off,
142 Some(bytes) => Self::Fixed(bytes),
143 }
144 }
145}
146
147pub type MemoryProbe = Arc<dyn Fn() -> Option<u64> + Send + Sync>;
149
150pub fn available_memory() -> Option<u64> {
152 static SYSTEM: std::sync::LazyLock<Mutex<sysinfo::System>> =
153 std::sync::LazyLock::new(|| Mutex::new(sysinfo::System::new()));
154 let mut system = SYSTEM.lock().unwrap_or_else(|e| e.into_inner());
155 system.refresh_memory_specifics(sysinfo::MemoryRefreshKind::nothing().with_ram());
156 let available = match system.cgroup_limits() {
157 Some(limits) => limits.free_memory,
158 None => system.available_memory(),
159 };
160 (available > 0).then_some(available)
161}
162
163#[derive(Clone)]
165pub struct MemoryCheck {
166 pub limit: Limit,
167 pub probe: MemoryProbe,
168}
169
170impl MemoryCheck {
171 pub fn off() -> Self {
173 Self {
174 limit: Limit::Off,
175 probe: Arc::new(|| None),
176 }
177 }
178
179 fn room(&self, held: u64) -> Option<u64> {
182 match self.limit {
183 Limit::Off => None,
184 Limit::Fixed(bytes) => Some(bytes.saturating_sub(held)),
185 Limit::Available => (self.probe)(),
186 }
187 }
188
189 pub fn refuses(&self, estimate: u64) -> Option<String> {
192 let room = self.room(0)?;
193 if estimate <= room {
194 return None;
195 }
196 let bytes = |n: u64| crate::numfmt::bytes(n);
197 let against = match self.limit {
198 Limit::Fixed(limit) => format!("more than {MEMORY_SETTING} ({})", bytes(limit)),
199 _ => format!("more than the {} available now", bytes(room)),
200 };
201 Some(format!(
202 "~{}, {against}\nEnter again to draw anyway {} set a limit: -c {MEMORY_SETTING}=8GiB",
203 bytes(estimate),
204 crate::glyphs::get().middot
205 ))
206 }
207
208 fn stops(&self, rows: &SampleRows, still: u64) -> Option<String> {
210 self.past(rows.bytes() as u64, rows.rows(), still)
211 }
212
213 pub fn holds_too_much(&self, held: u64, rows: usize) -> Option<String> {
217 self.past(held, rows, held)
218 }
219
220 fn past(&self, held: u64, rows: usize, still: u64) -> Option<String> {
221 let room = self.room(held)?;
222 (still > room).then(|| {
223 let why = match self.limit {
224 Limit::Fixed(_) => format!("{MEMORY_SETTING} reached"),
225 _ => "memory ran low".to_string(),
226 };
227 format!(
228 "Sample stopped at {} ({} rows): {why}; -c {MEMORY_SETTING}=0 draws on",
229 crate::numfmt::bytes(held),
230 crate::numfmt::group_chrome(rows)
231 )
232 })
233 }
234}
235
236#[derive(Clone)]
238pub struct Live {
239 pub rows: Arc<SampleRows>,
240 pub notify: Arc<dyn Fn() + Send + Sync>,
242 pub memory: MemoryCheck,
243 pub watch: ReadWatch,
244 pub bytes_per_row: Option<usize>,
247}
248
249impl Live {
250 fn keep(&self, key: u64, df: DataFrame, expected: Option<usize>) -> bool {
254 let last = df.estimated_size() as u64;
255 if df.height() > 0 {
256 self.rows.push(key, df);
257 (self.notify)();
258 }
259 let held = self.rows.rows();
260 let per_row = match self.rows.bytes().checked_div(held) {
261 Some(measured) if measured > 0 => measured,
262 _ => self.bytes_per_row.unwrap_or(0),
263 } as u64;
264 let still = match expected {
266 Some(rows) => rows.saturating_sub(held) as u64 * per_row,
267 None => last.saturating_mul(2),
268 };
269 if let Some(reason) = self.memory.stops(&self.rows, still) {
270 self.rows.stop(reason);
271 self.watch.stop();
272 return false;
273 }
274 !self.watch.stopped()
275 }
276}
277
278pub(crate) fn rebind(
281 plan: &mut polars::lazy::dsl::DslPlan,
282 old: &Arc<DataFrame>,
283 new: &Arc<DataFrame>,
284) {
285 use polars::lazy::dsl::DslPlan;
286 match plan {
287 DslPlan::IR { dsl, .. } => {
290 let mut inner = Arc::unwrap_or_clone(dsl.clone());
291 rebind(&mut inner, old, new);
292 *plan = inner;
293 return;
294 }
295 DslPlan::DataFrameScan { df, .. } if Arc::ptr_eq(df, old) => {
296 *df = Arc::clone(new);
297 return;
298 }
299 _ => {}
300 }
301 crate::table::for_each_input(plan, &mut |input| rebind(input, old, new));
302}
303
304#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
307#[serde(rename_all = "snake_case", tag = "kind")]
308pub enum DrawPath {
309 Reservoir,
311 Bernoulli { of: usize },
313}
314
315#[derive(Debug, Clone, Default, PartialEq)]
317pub struct Drawn {
318 pub total: Option<usize>,
320 pub about: bool,
322 pub per_value: Option<usize>,
324 pub cut: bool,
326 pub path: Option<DrawPath>,
328}
329
330pub fn draw(
336 lf: &LazyFrame,
337 sample: &Sample,
338 known_total: Option<usize>,
339 path: Option<DrawPath>,
340 polars_streaming: bool,
341 live: &Live,
342) -> Result<Drawn> {
343 let n = sample.rows.max(1);
344 let drawn = match &sample.method {
345 SampleMethod::EveryRow => {
346 let seen = stream(lf, live, known_total)?;
347 Drawn {
348 total: Some(seen),
349 ..Drawn::default()
350 }
351 }
352 SampleMethod::FirstRows => {
353 let seen = stream(&lf.clone().limit(n as IdxSize), live, Some(n))?;
354 Drawn {
355 total: known_total.or((seen < n).then_some(seen)),
356 ..Drawn::default()
357 }
358 }
359 SampleMethod::Spread if crate::analysis::sampling::slices_reach_into_the_scan(lf) => {
360 let total = match known_total {
361 Some(total) => total,
362 None => crate::analysis::sampling::count_rows(lf, polars_streaming)?,
363 };
364 if total <= n {
365 stream(lf, live, Some(total))?;
366 } else {
367 let on_run = |offset: usize, run: &DataFrame| {
368 live.keep(offset as u64, run.clone(), Some(n));
369 };
370 let read = crate::analysis::sampling::block_sample_live(
371 lf,
372 total,
373 n,
374 sample.seed,
375 polars_streaming,
376 &live.watch,
377 &on_run,
378 );
379 match read {
380 Ok(Some(df)) => {
381 live.keep(0, df, Some(n));
384 }
385 Ok(None) => {}
386 Err(error) if error.to_string() == CANCELLED => {}
387 Err(error) => return Err(error),
388 }
389 }
390 Drawn {
391 total: Some(total),
392 ..Drawn::default()
393 }
394 }
395 SampleMethod::Spread => match path.unwrap_or(match known_total {
396 Some(of) => DrawPath::Bernoulli { of },
397 None => DrawPath::Reservoir,
398 }) {
399 DrawPath::Bernoulli { of } => {
400 bernoulli(lf, n, of, sample.seed, live)?;
402 Drawn {
403 total: Some(live.watch.rows_seen().unwrap_or(of)),
404 about: of > n,
405 path: Some(DrawPath::Bernoulli { of }),
406 ..Drawn::default()
407 }
408 }
409 DrawPath::Reservoir => {
410 let read = crate::analysis::sampling::acquire(
411 lf,
412 sample,
413 None,
414 polars_streaming,
415 Some(&live.watch),
416 None,
417 )?;
418 let total = read.rows.total_rows;
419 live.keep(0, read.rows.df, Some(n));
420 Drawn {
421 total: Some(total),
422 path: Some(DrawPath::Reservoir),
423 ..Drawn::default()
424 }
425 }
426 },
427 SampleMethod::PerPartition { .. } => {
428 let read = crate::analysis::sampling::acquire(
429 lf,
430 sample,
431 known_total,
432 polars_streaming,
433 Some(&live.watch),
434 None,
435 )?;
436 let total = read.rows.total_rows;
437 let per_value = read.rows.per_value.as_ref().map(|per_value| per_value.kept);
438 live.keep(0, read.rows.df, None);
439 Drawn {
440 total: Some(total),
441 per_value,
442 ..Drawn::default()
443 }
444 }
445 };
446 if let Some(reason) = live.watch.memory_stopped()
448 && live.rows.stopped().is_none()
449 {
450 live.rows.stop(reason);
451 }
452 let cut = live.watch.stopped() || live.rows.stopped().is_some();
453 if cut && live.rows.rows() == 0 {
454 return Err(Report::msg(CANCELLED));
455 }
456 Ok(Drawn { cut, ..drawn })
457}
458
459fn stream(lf: &LazyFrame, live: &Live, expected: Option<usize>) -> Result<usize> {
462 let seen = Arc::new(std::sync::atomic::AtomicUsize::new(0));
463 let counted = Arc::clone(&seen);
464 let kept = live.clone();
465 stream_batches(lf.clone(), Some(&live.watch), true, move |batch| {
466 let key = counted.fetch_add(batch.height(), std::sync::atomic::Ordering::Relaxed);
467 Ok(!kept.keep(key as u64, batch, expected))
468 })?;
469 Ok(seen.load(std::sync::atomic::Ordering::Relaxed))
470}
471
472fn bernoulli(lf: &LazyFrame, n: usize, total: usize, seed: u64, live: &Live) -> Result<()> {
476 let bar = bernoulli_bar(n, total);
477 let kept = live.clone();
478 stream_batches(
479 lf.clone().with_row_index(POSITION, None),
480 Some(&live.watch),
481 true,
482 move |batch| {
483 let (first, rows) = bernoulli_keep(&batch, seed, bar)?;
484 Ok(!kept.keep(first, rows, Some(n)))
485 },
486 )
487}
488
489fn bernoulli_bar(n: usize, total: usize) -> u128 {
491 let share = (n as f64 / total.max(1) as f64).min(1.0);
492 (share * (u64::MAX as f64 + 1.0)) as u128
493}
494
495fn bernoulli_keep(batch: &DataFrame, seed: u64, bar: u128) -> PolarsResult<(u64, DataFrame)> {
498 let positions = batch.column(POSITION)?.idx()?;
499 let first = positions.get(0).unwrap_or(0) as u64;
500 let picked: Vec<IdxSize> = positions
501 .into_no_null_iter()
502 .enumerate()
503 .filter(|(_, position)| (sample_rank(seed, *position as u64) as u128) < bar)
504 .map(|(index, _)| index as IdxSize)
505 .collect();
506 let kept = batch
507 .take(&IdxCa::from_vec("kept".into(), picked))?
508 .drop(POSITION)?;
509 Ok((first, kept))
510}
511
512#[cfg(test)]
513mod tests;