polars-ooc 0.55.2

Out-of-core processing support for Polars
Documentation
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, LazyLock, RwLock};

use polars_async::ASYNC;
use polars_async::executor::TaskPriority;
use polars_config::config;
use polars_utils::total_ord::TotalOrd;
use polars_utils::with_drop::WithDrop;
use tokio::sync::{Mutex as AsyncMutex, Semaphore as AsyncSemaphore};

// How much worse than the best achieved (sample) score are we willing to look
// for spillables.
const EXPLORE_BEYOND_BEST_SCORE_THRESHOLD: f64 = 20.0;

const MAX_PARALLEL_SPILL_TASKS: usize = 64;

use crate::WeakSpillContext;
use crate::spill_context::UNEXPLORED_SCORE;
use crate::spill_token::{DynSpillToken, TrySpillError};

static MEMORY_MANAGER: LazyLock<MemoryManager> = LazyLock::new(MemoryManager::new);

/// Return a reference to the global [`MemoryManager`].
pub fn memory_manager() -> &'static MemoryManager {
    &MEMORY_MANAGER
}

pub struct MemoryManager {
    contexts: RwLock<Vec<WeakSpillContext>>,
    finding_spill_lock: AsyncMutex<()>,
    spill_semaphore: Arc<AsyncSemaphore>,
    est_spill_in_progress: AtomicU64,
}

impl MemoryManager {
    fn new() -> Self {
        Self {
            contexts: RwLock::new(Vec::new()),
            finding_spill_lock: AsyncMutex::new(()),
            spill_semaphore: Arc::new(AsyncSemaphore::new(MAX_PARALLEL_SPILL_TASKS)),
            est_spill_in_progress: AtomicU64::new(0),
        }
    }

    fn should_spill(&self) -> bool {
        let usage = crate::estimate_memory_usage();
        let likely_dealt_with = self.est_spill_in_progress.load(Ordering::Relaxed);
        usage.saturating_sub(likely_dealt_with) > config().ooc_memory_budget_bytes()
    }

    fn clean_contexts(&self) {
        if let Ok(mut ctxs) = self.contexts.try_write() {
            ctxs.retain(|ctx| !ctx.is_dead());
        }
    }

    pub(crate) fn register_ctx(&self, ctx: WeakSpillContext) {
        self.contexts.write().unwrap().push(ctx);
    }

    #[inline(always)]
    pub async fn spill(&self) {
        if self.should_spill() {
            self.do_spill().await
        }
    }

    #[inline(always)]
    pub fn spill_blocking(&self) {
        if self.should_spill() {
            self.do_spill_blocking()
        }
    }

    #[inline(never)]
    #[cold]
    fn do_spill_blocking(&self) {
        ASYNC.block_in_place_on(self.do_spill())
    }

    #[inline(never)]
    #[cold]
    async fn do_spill(&self) {
        while self.should_spill() {
            let Some((ctx, spillables)) = self.find_spillables().await else {
                return;
            };

            let successful_spill = Arc::new(WithDrop::new(
                (AtomicBool::new(false), ctx.clone()),
                move |(success, weak_ctx)| {
                    if !success.load(Ordering::Relaxed) {
                        if let Some(strong) = weak_ctx.upgrade() {
                            strong.stats().finish_exploration_event(false);
                        }
                    }
                },
            ));

            for (spillable, reg_id, sz) in spillables {
                let permit = self.spill_semaphore.clone().acquire_owned().await.unwrap();

                let successful_spill = successful_spill.clone();
                let ctx = ctx.clone();

                polars_async::executor::spawn(TaskPriority::High, async move {
                    // Spill, or reinsert if a failure.
                    match spillable.try_spill(ctx.clone(), reg_id) {
                        Ok(spill_success) => {
                            if spill_success.await {
                                if !successful_spill.0.swap(true, Ordering::Relaxed) {
                                    if let Some(strong) = ctx.upgrade() {
                                        strong.stats().finish_exploration_event(true);
                                    }
                                }
                            } else {
                                ctx.0.reinsert(&spillable, reg_id, ctx.1);
                            }
                        },
                        Err(TrySpillError::Pinned) => {
                            ctx.0.reinsert(&spillable, reg_id, ctx.1);
                        },
                        Err(TrySpillError::AlreadySpilled) => {},
                    }

                    MEMORY_MANAGER
                        .est_spill_in_progress
                        .fetch_sub(sz as u64, Ordering::Relaxed);

                    drop(permit);
                });
            }
        }
    }

    #[inline(never)]
    #[cold]
    async fn find_spillables(
        &self,
    ) -> Option<(WeakSpillContext, Vec<(Arc<dyn DynSpillToken>, u32, usize)>)> {
        // TODO: don't block here under a certain memory threshold.
        let finding_spill_guard = self.finding_spill_lock.lock().await;

        // TODO: don't loop over all contexts here, keep track of good ones and inspect those plus a couple random ones.
        let contexts = self.contexts.read().unwrap();
        let mut has_dead_context = false;
        let mut live_contexts = Vec::new();
        let mut rng = rand::rng();
        for ctx in contexts.iter() {
            if ctx.is_dead() {
                has_dead_context = true;
                continue;
            };

            // Thompson sampling.
            let score_sample = ctx.0.stats().sample_score(&mut rng);
            assert!(!score_sample.is_nan());
            live_contexts.push((ctx.clone(), score_sample));
        }
        drop(contexts);

        // Find the best context and loop over its candidates. For each
        // candidate we check if it can be spilled else we reinsert it.
        let min_spill = config().ooc_spill_min_bytes();
        live_contexts.sort_by(|a, b| a.1.tot_cmp(&b.1).reverse());
        let best_explored_score = live_contexts
            .iter()
            .map(|(_ctx, score)| *score)
            .find(|s| *s < UNEXPLORED_SCORE)
            .unwrap_or_default();

        let mut out = None;
        for (ctx, score) in live_contexts {
            // Refuse to consider contexts which are significantly worse than
            // the best already-explored one.
            if score * EXPLORE_BEYOND_BEST_SCORE_THRESHOLD < best_explored_score {
                break;
            }

            let Some(strong) = ctx.upgrade() else {
                continue;
            };
            strong.stats().start_exploration_event();

            let mut total_est_spill = 0;
            let mut candidates = Vec::new();
            for (cand, reg_id) in ctx.0.pop() {
                if cand.can_spill()
                    && let Some(sz) = cand.estimate_byte_size()
                    && sz as u64 >= min_spill
                {
                    total_est_spill += sz as u64;
                    candidates.push((cand, reg_id, sz));
                } else {
                    if !cand.is_spilled_or_dropped() {
                        ctx.0.reinsert(&cand, reg_id, ctx.1);
                    }
                }
            }

            if candidates.is_empty() {
                strong.stats().finish_exploration_event(false);
            } else {
                // Increment the spill-in-progress to avoid eager over-spilling.
                self.est_spill_in_progress
                    .fetch_add(total_est_spill, Ordering::Relaxed);
                out = Some((ctx, candidates));
                break;
            }
        }

        drop(finding_spill_guard);
        if has_dead_context {
            self.clean_contexts();
        }
        out
    }
}