use std::sync::{Arc, Mutex};
use color_eyre::Result;
use color_eyre::eyre::Report;
use polars::prelude::*;
use crate::sampling::{CANCELLED, ReadWatch, Sample, SampleMethod};
use crate::statistics::{collect_lazy, sample_rank};
pub const MEMORY_SETTING: &str = "analysis.sample_memory_limit";
const POSITION: &str = "__datui_table_sample_position";
#[derive(Default)]
pub struct SampleRows {
inner: Mutex<Inner>,
}
#[derive(Default)]
struct Inner {
chunks: Vec<(u64, DataFrame)>,
taken: usize,
rows: usize,
bytes: usize,
stopped: Option<String>,
}
impl SampleRows {
fn lock(&self) -> std::sync::MutexGuard<'_, Inner> {
self.inner.lock().unwrap_or_else(|e| e.into_inner())
}
pub fn push(&self, key: u64, df: DataFrame) {
let mut inner = self.lock();
inner.rows += df.height();
inner.bytes += df.estimated_size();
inner.chunks.push((key, df));
}
pub fn take_new(&self) -> Vec<DataFrame> {
let mut inner = self.lock();
let from = inner.taken;
inner.taken = inner.chunks.len();
inner.chunks[from..]
.iter()
.map(|(_, df)| df.clone())
.collect()
}
pub fn rows(&self) -> usize {
self.lock().rows
}
pub fn bytes(&self) -> usize {
self.lock().bytes
}
pub fn stopped(&self) -> Option<String> {
self.lock().stopped.clone()
}
fn stop(&self, reason: String) {
self.lock().stopped = Some(reason);
}
pub fn take_in_source_order(&self) -> Result<Option<DataFrame>> {
let ordered = self.in_source_order()?;
let mut inner = self.lock();
inner.chunks.clear();
inner.taken = 0;
drop(inner);
Ok(ordered.map(|mut frame| {
frame.rechunk_mut_par();
frame
}))
}
pub fn in_source_order(&self) -> Result<Option<DataFrame>> {
let inner = self.lock();
let mut order: Vec<&(u64, DataFrame)> = inner.chunks.iter().collect();
order.sort_by_key(|(key, _)| *key);
let mut out: Option<DataFrame> = None;
for (_, df) in order {
match out.as_mut() {
Some(frame) => {
frame.vstack_mut(df)?;
}
None => out = Some(df.clone()),
}
}
Ok(out)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Limit {
Available,
Fixed(u64),
Off,
}
impl Limit {
pub fn of_setting(setting: Option<crate::config::ByteSize>) -> Self {
match setting.map(|size| size.bytes()) {
None => Self::Available,
Some(0) => Self::Off,
Some(bytes) => Self::Fixed(bytes),
}
}
}
pub type MemoryProbe = Arc<dyn Fn() -> Option<u64> + Send + Sync>;
pub fn available_memory() -> Option<u64> {
static SYSTEM: std::sync::LazyLock<Mutex<sysinfo::System>> =
std::sync::LazyLock::new(|| Mutex::new(sysinfo::System::new()));
let mut system = SYSTEM.lock().unwrap_or_else(|e| e.into_inner());
system.refresh_memory_specifics(sysinfo::MemoryRefreshKind::nothing().with_ram());
let available = match system.cgroup_limits() {
Some(limits) => limits.free_memory,
None => system.available_memory(),
};
(available > 0).then_some(available)
}
#[derive(Clone)]
pub struct MemoryCheck {
pub limit: Limit,
pub probe: MemoryProbe,
}
impl MemoryCheck {
pub fn off() -> Self {
Self {
limit: Limit::Off,
probe: Arc::new(|| None),
}
}
fn room(&self, held: u64) -> Option<u64> {
match self.limit {
Limit::Off => None,
Limit::Fixed(bytes) => Some(bytes.saturating_sub(held)),
Limit::Available => (self.probe)(),
}
}
pub fn refuses(&self, estimate: u64) -> Option<String> {
let room = self.room(0)?;
if estimate <= room {
return None;
}
let bytes = |n: u64| crate::widgets::info::format_bytes(n);
let against = match self.limit {
Limit::Fixed(limit) => format!("more than {MEMORY_SETTING} ({})", bytes(limit)),
_ => format!("more than the {} available now", bytes(room)),
};
Some(format!(
"~{}, {against}\nEnter again to draw anyway {} set a limit: -c {MEMORY_SETTING}=8GiB",
bytes(estimate),
crate::glyphs::get().middot
))
}
fn stops(&self, rows: &SampleRows, still: u64) -> Option<String> {
self.past(rows.bytes() as u64, rows.rows(), still)
}
pub fn holds_too_much(&self, held: u64, rows: usize) -> Option<String> {
self.past(held, rows, held)
}
fn past(&self, held: u64, rows: usize, still: u64) -> Option<String> {
let room = self.room(held)?;
(still > room).then(|| {
let why = match self.limit {
Limit::Fixed(_) => format!("{MEMORY_SETTING} reached"),
_ => "memory ran low".to_string(),
};
format!(
"Sample stopped at {} ({} rows): {why}; -c {MEMORY_SETTING}=0 draws on",
crate::widgets::info::format_bytes(held),
crate::numfmt::group_chrome(rows)
)
})
}
}
pub struct Live {
pub rows: Arc<SampleRows>,
pub notify: Arc<dyn Fn() + Send + Sync>,
pub memory: MemoryCheck,
pub watch: ReadWatch,
pub bytes_per_row: Option<usize>,
}
impl Live {
fn keep(&self, key: u64, df: DataFrame, expected: Option<usize>) -> bool {
let last = df.estimated_size() as u64;
if df.height() > 0 {
self.rows.push(key, df);
(self.notify)();
}
let held = self.rows.rows();
let per_row = match self.rows.bytes().checked_div(held) {
Some(measured) if measured > 0 => measured,
_ => self.bytes_per_row.unwrap_or(0),
} as u64;
let still = match expected {
Some(rows) => rows.saturating_sub(held) as u64 * per_row,
None => last.saturating_mul(2),
};
if let Some(reason) = self.memory.stops(&self.rows, still) {
self.rows.stop(reason);
self.watch.stop();
return false;
}
!self.watch.stopped()
}
}
pub(crate) fn rebind(
plan: &mut polars::lazy::dsl::DslPlan,
old: &Arc<DataFrame>,
new: &Arc<DataFrame>,
) {
use polars::lazy::dsl::DslPlan;
match plan {
DslPlan::IR { dsl, .. } => {
let mut inner = Arc::unwrap_or_clone(dsl.clone());
rebind(&mut inner, old, new);
*plan = inner;
return;
}
DslPlan::DataFrameScan { df, .. } if Arc::ptr_eq(df, old) => {
*df = Arc::clone(new);
return;
}
_ => {}
}
crate::widgets::datatable::for_each_input(plan, &mut |input| rebind(input, old, new));
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case", tag = "kind")]
pub enum DrawPath {
Reservoir,
Bernoulli { of: usize },
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct Drawn {
pub total: Option<usize>,
pub about: bool,
pub per_value: Option<usize>,
pub cut: bool,
pub path: Option<DrawPath>,
}
pub fn draw(
lf: &LazyFrame,
sample: &Sample,
known_total: Option<usize>,
path: Option<DrawPath>,
polars_streaming: bool,
live: &Live,
) -> Result<Drawn> {
let n = sample.rows.max(1);
let drawn = match &sample.method {
SampleMethod::EveryRow => {
let seen = stream(lf, live, known_total)?;
Drawn {
total: Some(seen),
..Drawn::default()
}
}
SampleMethod::FirstRows => {
let seen = stream(&lf.clone().limit(n as IdxSize), live, Some(n))?;
Drawn {
total: known_total.or((seen < n).then_some(seen)),
..Drawn::default()
}
}
SampleMethod::Spread if crate::statistics::slices_reach_into_the_scan(lf) => {
let total = match known_total {
Some(total) => total,
None => crate::statistics::count_rows(lf, polars_streaming)?,
};
if total <= n {
stream(lf, live, Some(total))?;
} else {
let on_run = |offset: usize, run: &DataFrame| {
live.keep(offset as u64, run.clone(), Some(n));
};
let read = crate::statistics::block_sample_live(
lf,
total,
n,
sample.seed,
polars_streaming,
&live.watch,
&on_run,
);
match read {
Ok(Some(df)) => {
live.keep(0, df, Some(n));
}
Ok(None) => {}
Err(error) if error.to_string() == CANCELLED => {}
Err(error) => return Err(error),
}
}
Drawn {
total: Some(total),
..Drawn::default()
}
}
SampleMethod::Spread => match path.unwrap_or(match known_total {
Some(of) => DrawPath::Bernoulli { of },
None => DrawPath::Reservoir,
}) {
DrawPath::Bernoulli { of } => {
bernoulli(lf, n, of, sample.seed, live)?;
Drawn {
total: Some(live.watch.rows_seen().unwrap_or(of)),
about: of > n,
path: Some(DrawPath::Bernoulli { of }),
..Drawn::default()
}
}
DrawPath::Reservoir => {
let read = crate::sampling::acquire(
lf,
sample,
None,
polars_streaming,
Some(&live.watch),
None,
)?;
let total = read.rows.total_rows;
live.keep(0, read.rows.df, Some(n));
Drawn {
total: Some(total),
path: Some(DrawPath::Reservoir),
..Drawn::default()
}
}
},
SampleMethod::PerPartition { .. } => {
let read = crate::sampling::acquire(
lf,
sample,
known_total,
polars_streaming,
Some(&live.watch),
None,
)?;
let total = read.rows.total_rows;
let per_value = read.rows.per_value.as_ref().map(|per_value| per_value.kept);
live.keep(0, read.rows.df, None);
Drawn {
total: Some(total),
per_value,
..Drawn::default()
}
}
};
if let Some(reason) = live.watch.memory_stopped()
&& live.rows.stopped().is_none()
{
live.rows.stop(reason);
}
let cut = live.watch.stopped() || live.rows.stopped().is_some();
if cut && live.rows.rows() == 0 {
return Err(Report::msg(CANCELLED));
}
Ok(Drawn { cut, ..drawn })
}
fn stream(lf: &LazyFrame, live: &Live, expected: Option<usize>) -> Result<usize> {
let seen = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let rows = Arc::clone(&live.rows);
let notify = Arc::clone(&live.notify);
let memory = live.memory.clone();
let watch = live.watch.clone();
let bytes_per_row = live.bytes_per_row;
let counted = Arc::clone(&seen);
let sink = lf.clone().sink_batches(
PlanCallback::new(move |batch: DataFrame| {
if watch.stopped() {
return Ok(true);
}
watch.saw(batch.height());
let key = counted.fetch_add(batch.height(), std::sync::atomic::Ordering::Relaxed);
let live = Live {
rows: Arc::clone(&rows),
notify: Arc::clone(¬ify),
memory: memory.clone(),
watch: watch.clone(),
bytes_per_row,
};
Ok(!live.keep(key as u64, batch, expected))
}),
true,
None,
)?;
collect_lazy(sink, true).map_err(Report::from)?;
Ok(seen.load(std::sync::atomic::Ordering::Relaxed))
}
fn bernoulli(lf: &LazyFrame, n: usize, total: usize, seed: u64, live: &Live) -> Result<()> {
let bar = bernoulli_bar(n, total);
let rows = Arc::clone(&live.rows);
let notify = Arc::clone(&live.notify);
let memory = live.memory.clone();
let watch = live.watch.clone();
let bytes_per_row = live.bytes_per_row;
let sink = lf.clone().with_row_index(POSITION, None).sink_batches(
PlanCallback::new(move |batch: DataFrame| {
if watch.stopped() {
return Ok(true);
}
watch.saw(batch.height());
let (first, kept) = bernoulli_keep(&batch, seed, bar)?;
let live = Live {
rows: Arc::clone(&rows),
notify: Arc::clone(¬ify),
memory: memory.clone(),
watch: watch.clone(),
bytes_per_row,
};
Ok(!live.keep(first, kept, Some(n)))
}),
true,
None,
)?;
collect_lazy(sink, true).map_err(Report::from)?;
Ok(())
}
fn bernoulli_bar(n: usize, total: usize) -> u128 {
let share = (n as f64 / total.max(1) as f64).min(1.0);
(share * (u64::MAX as f64 + 1.0)) as u128
}
fn bernoulli_keep(batch: &DataFrame, seed: u64, bar: u128) -> PolarsResult<(u64, DataFrame)> {
let positions = batch.column(POSITION)?.idx()?;
let first = positions.get(0).unwrap_or(0) as u64;
let picked: Vec<IdxSize> = positions
.into_no_null_iter()
.enumerate()
.filter(|(_, position)| (sample_rank(seed, *position as u64) as u128) < bar)
.map(|(index, _)| index as IdxSize)
.collect();
let kept = batch
.take(&IdxCa::from_vec("kept".into(), picked))?
.drop(POSITION)?;
Ok((first, kept))
}
#[cfg(test)]
mod tests {
use super::*;
fn table(rows: i64) -> LazyFrame {
df!("value" => (0..rows).collect::<Vec<_>>())
.unwrap()
.lazy()
}
fn live() -> Live {
Live {
rows: Arc::new(SampleRows::default()),
notify: Arc::new(|| {}),
memory: MemoryCheck::off(),
watch: ReadWatch::default(),
bytes_per_row: None,
}
}
fn values(df: &DataFrame) -> Vec<i64> {
df.column("value")
.unwrap()
.i64()
.unwrap()
.into_no_null_iter()
.collect()
}
#[test]
fn bernoulli_keeps_about_n_rows_in_order_and_repeats() {
let (n, total) = (2_000usize, 100_000usize);
let one = live();
bernoulli(&table(total as i64), n, total, 42_891, &one).unwrap();
let kept = one.rows.in_source_order().unwrap().unwrap();
let spread = 4.0 * (n as f64).sqrt();
assert!(
(kept.height() as f64 - n as f64).abs() < spread,
"{} rows",
kept.height()
);
let rows = values(&kept);
assert!(
rows.windows(2).all(|pair| pair[0] < pair[1]),
"source order"
);
assert!(kept.column(POSITION).is_err(), "no helper column leaks");
let two = live();
bernoulli(&table(total as i64), n, total, 42_891, &two).unwrap();
assert_eq!(values(&two.rows.in_source_order().unwrap().unwrap()), rows);
let other = live();
bernoulli(&table(total as i64), n, total, 7, &other).unwrap();
assert_ne!(
values(&other.rows.in_source_order().unwrap().unwrap()),
rows
);
}
#[test]
fn the_bernoulli_bar_spans_none_to_all() {
let batch = df!(POSITION => (0..1_000 as IdxSize).collect::<Vec<_>>(), "value" => (0..1_000i64).collect::<Vec<_>>()).unwrap();
let (_, all) = bernoulli_keep(&batch, 1, bernoulli_bar(1_000, 1_000)).unwrap();
assert_eq!(all.height(), 1_000);
let (_, none) = bernoulli_keep(&batch, 1, bernoulli_bar(0, 1_000)).unwrap();
assert_eq!(none.height(), 0);
}
#[test]
fn chunks_arrive_in_any_order_and_end_in_source_order() {
let rows = SampleRows::default();
let chunk = |from: i64| df!("value" => (from..from + 3).collect::<Vec<_>>()).unwrap();
rows.push(30, chunk(30));
rows.push(10, chunk(10));
let first = rows.take_new();
assert_eq!(
first.iter().flat_map(values).collect::<Vec<_>>(),
[30, 31, 32, 10, 11, 12]
);
rows.push(20, chunk(20));
let next = rows.take_new();
assert_eq!(next.len(), 1, "only what arrived since");
assert_eq!(rows.rows(), 9);
assert_eq!(
values(&rows.in_source_order().unwrap().unwrap()),
[10, 11, 12, 20, 21, 22, 30, 31, 32]
);
}
#[test]
fn every_row_and_the_head_stream_in() {
let all = live();
let sample = Sample {
method: SampleMethod::EveryRow,
..Sample::default()
};
let drawn = draw(&table(5_000), &sample, None, None, false, &all).unwrap();
assert_eq!((drawn.total, all.rows.rows()), (Some(5_000), 5_000));
let head = live();
let sample = Sample {
method: SampleMethod::FirstRows,
rows: 120,
..Sample::default()
};
draw(&table(5_000), &sample, Some(5_000), None, false, &head).unwrap();
assert_eq!(
values(&head.rows.in_source_order().unwrap().unwrap()),
(0..120).collect::<Vec<_>>()
);
}
#[test]
fn a_known_total_draws_about_n_and_an_unknown_one_exactly_n() {
let sample = Sample {
rows: 500,
..Sample::default()
};
let known = live();
let drawn = draw(&table(20_000), &sample, Some(20_000), None, false, &known).unwrap();
assert!(drawn.about);
let unknown = live();
let drawn = draw(&table(20_000), &sample, None, None, false, &unknown).unwrap();
assert!(!drawn.about);
assert_eq!((drawn.total, unknown.rows.rows()), (Some(20_000), 500));
}
#[test]
fn low_memory_stops_the_draw_and_keeps_what_it_has() {
let mut low = live();
low.memory = MemoryCheck {
limit: Limit::Available,
probe: Arc::new(|| Some(1)),
};
let sample = Sample {
method: SampleMethod::EveryRow,
..Sample::default()
};
let parts: Vec<LazyFrame> = (0..5).map(|_| table(100_000)).collect();
let lf = concat(parts, UnionArgs::default()).unwrap();
let drawn = draw(&lf, &sample, Some(500_000), None, false, &low).unwrap();
assert!(drawn.cut);
let held = low.rows.rows();
assert!(held > 0 && held < 500_000, "{held}");
let reason = low.rows.stopped().unwrap();
assert!(reason.contains("memory ran low"), "{reason}");
assert!(reason.contains(MEMORY_SETTING), "{reason}");
}
#[test]
fn a_recorded_path_draws_the_same_rows_whatever_is_known_now() {
let sample = Sample {
rows: 500,
..Sample::default()
};
let rows = |live: &Live| values(&live.rows.in_source_order().unwrap().unwrap());
let reservoir = live();
draw(&table(20_000), &sample, None, None, false, &reservoir).unwrap();
let counted = live();
let drawn = draw(
&table(20_000),
&sample,
Some(20_000),
Some(DrawPath::Reservoir),
false,
&counted,
)
.unwrap();
assert_eq!(drawn.path, Some(DrawPath::Reservoir));
assert_eq!(rows(&counted), rows(&reservoir));
let known = live();
let drawn = draw(&table(20_000), &sample, Some(20_000), None, false, &known).unwrap();
assert_eq!(drawn.path, Some(DrawPath::Bernoulli { of: 20_000 }));
let uncounted = live();
let path = drawn.path;
draw(&table(20_000), &sample, None, path, false, &uncounted).unwrap();
assert_eq!(rows(&uncounted), rows(&known));
}
#[test]
fn a_reservoir_past_the_memory_stops_and_keeps_what_it_holds() {
let check = MemoryCheck {
limit: Limit::Fixed(1),
probe: Arc::new(|| None),
};
let mut low = live();
low.watch = ReadWatch::judging_held(Arc::new(move |bytes, rows| {
check.holds_too_much(bytes, rows)
}));
let parts: Vec<LazyFrame> = (0..5).map(|_| table(100_000)).collect();
let lf = concat(parts, UnionArgs::default()).unwrap();
let sample = Sample {
rows: 400_000,
..Sample::default()
};
let drawn = draw(&lf, &sample, None, None, false, &low).unwrap();
assert!(drawn.cut);
assert!(low.rows.rows() > 0);
assert!(drawn.total.unwrap() < 500_000, "it stopped early");
let reason = low.rows.stopped().unwrap();
assert!(reason.contains(MEMORY_SETTING), "{reason}");
}
#[test]
fn an_estimate_past_the_room_is_refused_with_the_way_through() {
let check = MemoryCheck {
limit: Limit::Available,
probe: Arc::new(|| Some(4 << 30)),
};
let refused = check.refuses(6 << 30).unwrap();
assert!(refused.contains("available now"), "{refused}");
assert!(refused.contains("Enter again"), "{refused}");
assert!(
refused.contains("-c analysis.sample_memory_limit"),
"{refused}"
);
assert!(check.refuses(1 << 30).is_none());
assert!(MemoryCheck::off().refuses(u64::MAX).is_none());
let fixed = MemoryCheck {
limit: Limit::Fixed(1 << 20),
probe: Arc::new(|| None),
};
assert!(fixed.refuses(2 << 20).unwrap().contains(MEMORY_SETTING));
}
}