1use std::sync::{Arc, Mutex};
15
16use color_eyre::Result;
17use color_eyre::eyre::Report;
18use polars::prelude::*;
19
20use crate::sampling::{CANCELLED, ReadWatch, Sample, SampleMethod};
21use crate::statistics::{collect_lazy, sample_rank};
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::widgets::info::format_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::widgets::info::format_bytes(held),
230 crate::numfmt::group_chrome(rows)
231 )
232 })
233 }
234}
235
236pub struct Live {
238 pub rows: Arc<SampleRows>,
239 pub notify: Arc<dyn Fn() + Send + Sync>,
241 pub memory: MemoryCheck,
242 pub watch: ReadWatch,
243 pub bytes_per_row: Option<usize>,
246}
247
248impl Live {
249 fn keep(&self, key: u64, df: DataFrame, expected: Option<usize>) -> bool {
253 let last = df.estimated_size() as u64;
254 if df.height() > 0 {
255 self.rows.push(key, df);
256 (self.notify)();
257 }
258 let held = self.rows.rows();
259 let per_row = match self.rows.bytes().checked_div(held) {
260 Some(measured) if measured > 0 => measured,
261 _ => self.bytes_per_row.unwrap_or(0),
262 } as u64;
263 let still = match expected {
265 Some(rows) => rows.saturating_sub(held) as u64 * per_row,
266 None => last.saturating_mul(2),
267 };
268 if let Some(reason) = self.memory.stops(&self.rows, still) {
269 self.rows.stop(reason);
270 self.watch.stop();
271 return false;
272 }
273 !self.watch.stopped()
274 }
275}
276
277pub(crate) fn rebind(
280 plan: &mut polars::lazy::dsl::DslPlan,
281 old: &Arc<DataFrame>,
282 new: &Arc<DataFrame>,
283) {
284 use polars::lazy::dsl::DslPlan;
285 match plan {
286 DslPlan::IR { dsl, .. } => {
289 let mut inner = Arc::unwrap_or_clone(dsl.clone());
290 rebind(&mut inner, old, new);
291 *plan = inner;
292 return;
293 }
294 DslPlan::DataFrameScan { df, .. } if Arc::ptr_eq(df, old) => {
295 *df = Arc::clone(new);
296 return;
297 }
298 _ => {}
299 }
300 crate::widgets::datatable::for_each_input(plan, &mut |input| rebind(input, old, new));
301}
302
303#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
306#[serde(rename_all = "snake_case", tag = "kind")]
307pub enum DrawPath {
308 Reservoir,
310 Bernoulli { of: usize },
312}
313
314#[derive(Debug, Clone, Default, PartialEq)]
316pub struct Drawn {
317 pub total: Option<usize>,
319 pub about: bool,
321 pub per_value: Option<usize>,
323 pub cut: bool,
325 pub path: Option<DrawPath>,
327}
328
329pub fn draw(
335 lf: &LazyFrame,
336 sample: &Sample,
337 known_total: Option<usize>,
338 path: Option<DrawPath>,
339 polars_streaming: bool,
340 live: &Live,
341) -> Result<Drawn> {
342 let n = sample.rows.max(1);
343 let drawn = match &sample.method {
344 SampleMethod::EveryRow => {
345 let seen = stream(lf, live, known_total)?;
346 Drawn {
347 total: Some(seen),
348 ..Drawn::default()
349 }
350 }
351 SampleMethod::FirstRows => {
352 let seen = stream(&lf.clone().limit(n as IdxSize), live, Some(n))?;
353 Drawn {
354 total: known_total.or((seen < n).then_some(seen)),
355 ..Drawn::default()
356 }
357 }
358 SampleMethod::Spread if crate::statistics::slices_reach_into_the_scan(lf) => {
359 let total = match known_total {
360 Some(total) => total,
361 None => crate::statistics::count_rows(lf, polars_streaming)?,
362 };
363 if total <= n {
364 stream(lf, live, Some(total))?;
365 } else {
366 let on_run = |offset: usize, run: &DataFrame| {
367 live.keep(offset as u64, run.clone(), Some(n));
368 };
369 let read = crate::statistics::block_sample_live(
370 lf,
371 total,
372 n,
373 sample.seed,
374 polars_streaming,
375 &live.watch,
376 &on_run,
377 );
378 match read {
379 Ok(Some(df)) => {
380 live.keep(0, df, Some(n));
383 }
384 Ok(None) => {}
385 Err(error) if error.to_string() == CANCELLED => {}
386 Err(error) => return Err(error),
387 }
388 }
389 Drawn {
390 total: Some(total),
391 ..Drawn::default()
392 }
393 }
394 SampleMethod::Spread => match path.unwrap_or(match known_total {
395 Some(of) => DrawPath::Bernoulli { of },
396 None => DrawPath::Reservoir,
397 }) {
398 DrawPath::Bernoulli { of } => {
399 bernoulli(lf, n, of, sample.seed, live)?;
401 Drawn {
402 total: Some(live.watch.rows_seen().unwrap_or(of)),
403 about: of > n,
404 path: Some(DrawPath::Bernoulli { of }),
405 ..Drawn::default()
406 }
407 }
408 DrawPath::Reservoir => {
409 let read = crate::sampling::acquire(
410 lf,
411 sample,
412 None,
413 polars_streaming,
414 Some(&live.watch),
415 None,
416 )?;
417 let total = read.rows.total_rows;
418 live.keep(0, read.rows.df, Some(n));
419 Drawn {
420 total: Some(total),
421 path: Some(DrawPath::Reservoir),
422 ..Drawn::default()
423 }
424 }
425 },
426 SampleMethod::PerPartition { .. } => {
427 let read = crate::sampling::acquire(
428 lf,
429 sample,
430 known_total,
431 polars_streaming,
432 Some(&live.watch),
433 None,
434 )?;
435 let total = read.rows.total_rows;
436 let per_value = read.rows.per_value.as_ref().map(|per_value| per_value.kept);
437 live.keep(0, read.rows.df, None);
438 Drawn {
439 total: Some(total),
440 per_value,
441 ..Drawn::default()
442 }
443 }
444 };
445 if let Some(reason) = live.watch.memory_stopped()
447 && live.rows.stopped().is_none()
448 {
449 live.rows.stop(reason);
450 }
451 let cut = live.watch.stopped() || live.rows.stopped().is_some();
452 if cut && live.rows.rows() == 0 {
453 return Err(Report::msg(CANCELLED));
454 }
455 Ok(Drawn { cut, ..drawn })
456}
457
458fn stream(lf: &LazyFrame, live: &Live, expected: Option<usize>) -> Result<usize> {
461 let seen = Arc::new(std::sync::atomic::AtomicUsize::new(0));
462 let rows = Arc::clone(&live.rows);
463 let notify = Arc::clone(&live.notify);
464 let memory = live.memory.clone();
465 let watch = live.watch.clone();
466 let bytes_per_row = live.bytes_per_row;
467 let counted = Arc::clone(&seen);
468 let sink = lf.clone().sink_batches(
469 PlanCallback::new(move |batch: DataFrame| {
470 if watch.stopped() {
471 return Ok(true);
472 }
473 watch.saw(batch.height());
474 let key = counted.fetch_add(batch.height(), std::sync::atomic::Ordering::Relaxed);
475 let live = Live {
476 rows: Arc::clone(&rows),
477 notify: Arc::clone(¬ify),
478 memory: memory.clone(),
479 watch: watch.clone(),
480 bytes_per_row,
481 };
482 Ok(!live.keep(key as u64, batch, expected))
483 }),
484 true,
485 None,
486 )?;
487 collect_lazy(sink, true).map_err(Report::from)?;
489 Ok(seen.load(std::sync::atomic::Ordering::Relaxed))
490}
491
492fn bernoulli(lf: &LazyFrame, n: usize, total: usize, seed: u64, live: &Live) -> Result<()> {
496 let bar = bernoulli_bar(n, total);
497 let rows = Arc::clone(&live.rows);
498 let notify = Arc::clone(&live.notify);
499 let memory = live.memory.clone();
500 let watch = live.watch.clone();
501 let bytes_per_row = live.bytes_per_row;
502 let sink = lf.clone().with_row_index(POSITION, None).sink_batches(
503 PlanCallback::new(move |batch: DataFrame| {
504 if watch.stopped() {
505 return Ok(true);
506 }
507 watch.saw(batch.height());
508 let (first, kept) = bernoulli_keep(&batch, seed, bar)?;
509 let live = Live {
510 rows: Arc::clone(&rows),
511 notify: Arc::clone(¬ify),
512 memory: memory.clone(),
513 watch: watch.clone(),
514 bytes_per_row,
515 };
516 Ok(!live.keep(first, kept, Some(n)))
517 }),
518 true,
519 None,
520 )?;
521 collect_lazy(sink, true).map_err(Report::from)?;
522 Ok(())
523}
524
525fn bernoulli_bar(n: usize, total: usize) -> u128 {
527 let share = (n as f64 / total.max(1) as f64).min(1.0);
528 (share * (u64::MAX as f64 + 1.0)) as u128
529}
530
531fn bernoulli_keep(batch: &DataFrame, seed: u64, bar: u128) -> PolarsResult<(u64, DataFrame)> {
534 let positions = batch.column(POSITION)?.idx()?;
535 let first = positions.get(0).unwrap_or(0) as u64;
536 let picked: Vec<IdxSize> = positions
537 .into_no_null_iter()
538 .enumerate()
539 .filter(|(_, position)| (sample_rank(seed, *position as u64) as u128) < bar)
540 .map(|(index, _)| index as IdxSize)
541 .collect();
542 let kept = batch
543 .take(&IdxCa::from_vec("kept".into(), picked))?
544 .drop(POSITION)?;
545 Ok((first, kept))
546}
547
548#[cfg(test)]
549mod tests {
550 use super::*;
551
552 fn table(rows: i64) -> LazyFrame {
553 df!("value" => (0..rows).collect::<Vec<_>>())
554 .unwrap()
555 .lazy()
556 }
557
558 fn live() -> Live {
559 Live {
560 rows: Arc::new(SampleRows::default()),
561 notify: Arc::new(|| {}),
562 memory: MemoryCheck::off(),
563 watch: ReadWatch::default(),
564 bytes_per_row: None,
565 }
566 }
567
568 fn values(df: &DataFrame) -> Vec<i64> {
569 df.column("value")
570 .unwrap()
571 .i64()
572 .unwrap()
573 .into_no_null_iter()
574 .collect()
575 }
576
577 #[test]
580 fn bernoulli_keeps_about_n_rows_in_order_and_repeats() {
581 let (n, total) = (2_000usize, 100_000usize);
582 let one = live();
583 bernoulli(&table(total as i64), n, total, 42_891, &one).unwrap();
584 let kept = one.rows.in_source_order().unwrap().unwrap();
585 let spread = 4.0 * (n as f64).sqrt();
586 assert!(
587 (kept.height() as f64 - n as f64).abs() < spread,
588 "{} rows",
589 kept.height()
590 );
591 let rows = values(&kept);
592 assert!(
593 rows.windows(2).all(|pair| pair[0] < pair[1]),
594 "source order"
595 );
596 assert!(kept.column(POSITION).is_err(), "no helper column leaks");
597 let two = live();
598 bernoulli(&table(total as i64), n, total, 42_891, &two).unwrap();
599 assert_eq!(values(&two.rows.in_source_order().unwrap().unwrap()), rows);
600 let other = live();
601 bernoulli(&table(total as i64), n, total, 7, &other).unwrap();
602 assert_ne!(
603 values(&other.rows.in_source_order().unwrap().unwrap()),
604 rows
605 );
606 }
607
608 #[test]
611 fn the_bernoulli_bar_spans_none_to_all() {
612 let batch = df!(POSITION => (0..1_000 as IdxSize).collect::<Vec<_>>(), "value" => (0..1_000i64).collect::<Vec<_>>()).unwrap();
613 let (_, all) = bernoulli_keep(&batch, 1, bernoulli_bar(1_000, 1_000)).unwrap();
614 assert_eq!(all.height(), 1_000);
615 let (_, none) = bernoulli_keep(&batch, 1, bernoulli_bar(0, 1_000)).unwrap();
616 assert_eq!(none.height(), 0);
617 }
618
619 #[test]
622 fn chunks_arrive_in_any_order_and_end_in_source_order() {
623 let rows = SampleRows::default();
624 let chunk = |from: i64| df!("value" => (from..from + 3).collect::<Vec<_>>()).unwrap();
625 rows.push(30, chunk(30));
626 rows.push(10, chunk(10));
627 let first = rows.take_new();
628 assert_eq!(
629 first.iter().flat_map(values).collect::<Vec<_>>(),
630 [30, 31, 32, 10, 11, 12]
631 );
632 rows.push(20, chunk(20));
633 let next = rows.take_new();
634 assert_eq!(next.len(), 1, "only what arrived since");
635 assert_eq!(rows.rows(), 9);
636 assert_eq!(
637 values(&rows.in_source_order().unwrap().unwrap()),
638 [10, 11, 12, 20, 21, 22, 30, 31, 32]
639 );
640 }
641
642 #[test]
644 fn every_row_and_the_head_stream_in() {
645 let all = live();
646 let sample = Sample {
647 method: SampleMethod::EveryRow,
648 ..Sample::default()
649 };
650 let drawn = draw(&table(5_000), &sample, None, None, false, &all).unwrap();
651 assert_eq!((drawn.total, all.rows.rows()), (Some(5_000), 5_000));
652 let head = live();
653 let sample = Sample {
654 method: SampleMethod::FirstRows,
655 rows: 120,
656 ..Sample::default()
657 };
658 draw(&table(5_000), &sample, Some(5_000), None, false, &head).unwrap();
659 assert_eq!(
660 values(&head.rows.in_source_order().unwrap().unwrap()),
661 (0..120).collect::<Vec<_>>()
662 );
663 }
664
665 #[test]
668 fn a_known_total_draws_about_n_and_an_unknown_one_exactly_n() {
669 let sample = Sample {
670 rows: 500,
671 ..Sample::default()
672 };
673 let known = live();
674 let drawn = draw(&table(20_000), &sample, Some(20_000), None, false, &known).unwrap();
675 assert!(drawn.about);
676 let unknown = live();
677 let drawn = draw(&table(20_000), &sample, None, None, false, &unknown).unwrap();
678 assert!(!drawn.about);
679 assert_eq!((drawn.total, unknown.rows.rows()), (Some(20_000), 500));
680 }
681
682 #[test]
685 fn low_memory_stops_the_draw_and_keeps_what_it_has() {
686 let mut low = live();
687 low.memory = MemoryCheck {
688 limit: Limit::Available,
689 probe: Arc::new(|| Some(1)),
690 };
691 let sample = Sample {
692 method: SampleMethod::EveryRow,
693 ..Sample::default()
694 };
695 let parts: Vec<LazyFrame> = (0..5).map(|_| table(100_000)).collect();
697 let lf = concat(parts, UnionArgs::default()).unwrap();
698 let drawn = draw(&lf, &sample, Some(500_000), None, false, &low).unwrap();
699 assert!(drawn.cut);
700 let held = low.rows.rows();
701 assert!(held > 0 && held < 500_000, "{held}");
702 let reason = low.rows.stopped().unwrap();
703 assert!(reason.contains("memory ran low"), "{reason}");
704 assert!(reason.contains(MEMORY_SETTING), "{reason}");
705 }
706
707 #[test]
711 fn a_recorded_path_draws_the_same_rows_whatever_is_known_now() {
712 let sample = Sample {
713 rows: 500,
714 ..Sample::default()
715 };
716 let rows = |live: &Live| values(&live.rows.in_source_order().unwrap().unwrap());
717 let reservoir = live();
718 draw(&table(20_000), &sample, None, None, false, &reservoir).unwrap();
719 let counted = live();
720 let drawn = draw(
721 &table(20_000),
722 &sample,
723 Some(20_000),
724 Some(DrawPath::Reservoir),
725 false,
726 &counted,
727 )
728 .unwrap();
729 assert_eq!(drawn.path, Some(DrawPath::Reservoir));
730 assert_eq!(rows(&counted), rows(&reservoir));
731
732 let known = live();
733 let drawn = draw(&table(20_000), &sample, Some(20_000), None, false, &known).unwrap();
734 assert_eq!(drawn.path, Some(DrawPath::Bernoulli { of: 20_000 }));
735 let uncounted = live();
736 let path = drawn.path;
737 draw(&table(20_000), &sample, None, path, false, &uncounted).unwrap();
738 assert_eq!(rows(&uncounted), rows(&known));
739 }
740
741 #[test]
744 fn a_reservoir_past_the_memory_stops_and_keeps_what_it_holds() {
745 let check = MemoryCheck {
746 limit: Limit::Fixed(1),
747 probe: Arc::new(|| None),
748 };
749 let mut low = live();
750 low.watch = ReadWatch::judging_held(Arc::new(move |bytes, rows| {
751 check.holds_too_much(bytes, rows)
752 }));
753 let parts: Vec<LazyFrame> = (0..5).map(|_| table(100_000)).collect();
754 let lf = concat(parts, UnionArgs::default()).unwrap();
755 let sample = Sample {
756 rows: 400_000,
757 ..Sample::default()
758 };
759 let drawn = draw(&lf, &sample, None, None, false, &low).unwrap();
760 assert!(drawn.cut);
761 assert!(low.rows.rows() > 0);
762 assert!(drawn.total.unwrap() < 500_000, "it stopped early");
763 let reason = low.rows.stopped().unwrap();
764 assert!(reason.contains(MEMORY_SETTING), "{reason}");
765 }
766
767 #[test]
769 fn an_estimate_past_the_room_is_refused_with_the_way_through() {
770 let check = MemoryCheck {
771 limit: Limit::Available,
772 probe: Arc::new(|| Some(4 << 30)),
773 };
774 let refused = check.refuses(6 << 30).unwrap();
775 assert!(refused.contains("available now"), "{refused}");
776 assert!(refused.contains("Enter again"), "{refused}");
777 assert!(
778 refused.contains("-c analysis.sample_memory_limit"),
779 "{refused}"
780 );
781 assert!(check.refuses(1 << 30).is_none());
782 assert!(MemoryCheck::off().refuses(u64::MAX).is_none());
783 let fixed = MemoryCheck {
784 limit: Limit::Fixed(1 << 20),
785 probe: Arc::new(|| None),
786 };
787 assert!(fixed.refuses(2 << 20).unwrap().contains(MEMORY_SETTING));
788 }
789}