use std::{fmt, sync::Arc};
use datafusion::{
error::{DataFusionError, Result as DfResult},
execution::{
memory_pool::{MemoryLimit, MemoryPool, MemoryReservation},
runtime_env::{RuntimeEnv, RuntimeEnvBuilder},
},
prelude::{SessionConfig, SessionContext},
};
use crate::memory::ConnectionMemoryBudget;
#[derive(Debug)]
struct ConnectionBudgetPool {
budget: Arc<ConnectionMemoryBudget>,
}
const POOL_NAME: &str = "ConnectionBudgetPool";
const PARTIAL_AGG_SKIP_PROBE_RATIO: f64 = 0.5;
impl fmt::Display for ConnectionBudgetPool {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(POOL_NAME)
}
}
impl MemoryPool for ConnectionBudgetPool {
fn name(&self) -> &str {
POOL_NAME
}
fn grow(&self, _reservation: &MemoryReservation, additional: usize) {
self.budget.grow_unchecked(additional);
}
fn shrink(&self, _reservation: &MemoryReservation, shrink: usize) {
self.budget.release(shrink);
}
fn try_grow(&self, _reservation: &MemoryReservation, additional: usize) -> DfResult<()> {
self.budget.try_grow(additional).map_err(|over| {
DataFusionError::ResourcesExhausted(format!("during SQL query, {over}"))
})
}
fn reserved(&self) -> usize {
self.budget.used()
}
fn memory_limit(&self) -> MemoryLimit {
match self.budget.limit() {
Some(limit) => MemoryLimit::Finite(limit),
None => MemoryLimit::Infinite,
}
}
}
fn budgeted_runtime(budget: &Arc<ConnectionMemoryBudget>) -> DfResult<Arc<RuntimeEnv>> {
let pool: Arc<dyn MemoryPool> = Arc::new(ConnectionBudgetPool {
budget: Arc::clone(budget),
});
RuntimeEnvBuilder::new().with_memory_pool(pool).build_arc()
}
pub(crate) fn budgeted_session_context(
budget: &Arc<ConnectionMemoryBudget>,
) -> DfResult<SessionContext> {
let mut config = SessionConfig::new();
config.options_mut().optimizer.expand_views_at_output = true;
config
.options_mut()
.execution
.skip_partial_aggregation_probe_ratio_threshold = PARTIAL_AGG_SKIP_PROBE_RATIO;
Ok(SessionContext::new_with_config_rt(
config,
budgeted_runtime(budget)?,
))
}
#[cfg(test)]
mod tests {
use datafusion::{
arrow::{
array::{Array, StringArray},
datatypes::{DataType, Field, Schema},
record_batch::RecordBatch,
},
datasource::MemTable,
execution::memory_pool::MemoryConsumer,
physical_plan::{ExecutionPlan, collect},
};
use tokio::runtime::Runtime;
use super::*;
#[test]
fn measured_pool_never_refuses() {
let budget = ConnectionMemoryBudget::measured();
let pool: Arc<dyn MemoryPool> = Arc::new(ConnectionBudgetPool {
budget: Arc::clone(&budget),
});
let res = MemoryConsumer::new("t").register(&pool);
res.try_grow(1 << 30).expect("measured never refuses");
assert_eq!(pool.reserved(), 1 << 30);
assert!(matches!(pool.memory_limit(), MemoryLimit::Infinite));
}
#[test]
fn bounded_pool_refuses_past_the_gate() {
let budget = ConnectionMemoryBudget::with_limit(1000);
let pool: Arc<dyn MemoryPool> = Arc::new(ConnectionBudgetPool {
budget: Arc::clone(&budget),
});
assert!(matches!(pool.memory_limit(), MemoryLimit::Finite(900)));
let res = MemoryConsumer::new("t").register(&pool);
res.try_grow(900).expect("exactly at the gate fits");
res.try_grow(1)
.expect_err("one byte over the gate is refused");
assert_eq!(pool.reserved(), 900);
res.shrink(900);
assert_eq!(pool.reserved(), 0);
}
#[test]
fn pools_over_one_budget_share_the_counter() {
let budget = ConnectionMemoryBudget::with_limit(1000); let pool_a: Arc<dyn MemoryPool> = Arc::new(ConnectionBudgetPool {
budget: Arc::clone(&budget),
});
let pool_b: Arc<dyn MemoryPool> = Arc::new(ConnectionBudgetPool {
budget: Arc::clone(&budget),
});
let a = MemoryConsumer::new("a").register(&pool_a);
let b = MemoryConsumer::new("b").register(&pool_b);
a.try_grow(600).expect("fits");
b.try_grow(300).expect("600 + 300 = 900 fits");
b.try_grow(1).expect_err("the two pools share one ceiling");
}
#[test]
fn grow_charges_past_the_limit_for_unspillable_reservations() {
let budget = ConnectionMemoryBudget::with_limit(1000); let pool: Arc<dyn MemoryPool> = Arc::new(ConnectionBudgetPool { budget });
let res = MemoryConsumer::new("must-succeed").register(&pool);
res.grow(5000); assert_eq!(pool.reserved(), 5000);
}
fn total_spill_count(plan: &dyn ExecutionPlan) -> usize {
let here = plan.metrics().and_then(|m| m.spill_count()).unwrap_or(0);
here + plan
.children()
.iter()
.map(|c| total_spill_count(c.as_ref()))
.sum::<usize>()
}
fn is_sorted(batches: &[RecordBatch]) -> bool {
let mut prev: Option<String> = None;
for b in batches {
let col = b
.column(0)
.as_any()
.downcast_ref::<StringArray>()
.expect("string column");
for i in 0..col.len() {
let v = col.value(i).to_string();
if prev.as_ref().is_some_and(|p| &v < p) {
return false;
}
prev = Some(v);
}
}
true
}
async fn run_sorted_query(
schema: &Arc<Schema>,
batches: &[RecordBatch],
configured_bytes: u64,
reservation_bytes: usize,
) -> (usize, bool, usize) {
let budget = ConnectionMemoryBudget::with_limit(configured_bytes);
let runtime = budgeted_runtime(&budget).expect("runtime env");
let mut cfg = SessionConfig::new();
cfg.options_mut().execution.sort_spill_reservation_bytes = reservation_bytes;
cfg.options_mut().execution.target_partitions = 1;
let ctx = SessionContext::new_with_config_rt(cfg, runtime);
let table =
MemTable::try_new(Arc::clone(schema), vec![batches.to_vec()]).expect("memtable");
ctx.register_table("t", Arc::new(table)).expect("register");
let df = ctx.sql("SELECT s FROM t ORDER BY s").await.expect("plan");
let plan = df.create_physical_plan().await.expect("physical plan");
let out = collect(Arc::clone(&plan), ctx.task_ctx())
.await
.expect("collect");
let rows: usize = out.iter().map(|b| b.num_rows()).sum();
(rows, is_sorted(&out), total_spill_count(plan.as_ref()))
}
#[test]
fn bounded_pool_spills_a_large_sort_and_returns_correct_results() {
const ROWS: usize = 512 * 1024; const CHUNK: usize = 8 * 1024; const RESERVATION: usize = 20 * 1024 * 1024;
let filler = "x".repeat(56);
let schema = Arc::new(Schema::new(vec![Field::new("s", DataType::Utf8, false)]));
let mut batches = Vec::new();
for start in (0..ROWS).step_by(CHUNK) {
let vals: Vec<String> = (start..(start + CHUNK).min(ROWS))
.map(|i| format!("{:08}{filler}", ROWS - 1 - i))
.collect();
batches.push(
RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(StringArray::from(vals))])
.expect("batch"),
);
}
let runtime = Runtime::new().expect("tokio runtime");
runtime.block_on(async {
let (rows, sorted, spills) =
run_sorted_query(&schema, &batches, 32 * 1024 * 1024, RESERVATION).await;
assert_eq!(rows, ROWS, "every row comes back despite spilling");
assert!(sorted, "spilled sort still returns rows in order");
assert!(spills > 0, "the sort must spill under the tight budget");
let (rows, sorted, spills) =
run_sorted_query(&schema, &batches, 256 * 1024 * 1024, RESERVATION).await;
assert_eq!(rows, ROWS);
assert!(sorted);
assert_eq!(spills, 0, "a generous budget sorts in memory, no spill");
});
}
#[test]
fn budget_counter_climbs_per_query_and_returns_to_baseline() {
const GROUPS: usize = 50 * 1024;
let schema = Arc::new(Schema::new(vec![Field::new("s", DataType::Utf8, false)]));
let vals: Vec<String> = (0..GROUPS).map(|i| format!("key-{i}")).collect();
let batch =
RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(StringArray::from(vals))])
.expect("batch");
let budget = ConnectionMemoryBudget::measured();
let runtime_env = Runtime::new().expect("tokio runtime");
runtime_env.block_on(async {
let mut prev_peak = 0;
for i in 0..3 {
let rt = budgeted_runtime(&budget).expect("runtime env");
let mut cfg = SessionConfig::new();
cfg.options_mut().execution.target_partitions = 1;
let ctx = SessionContext::new_with_config_rt(cfg, rt);
let table = MemTable::try_new(Arc::clone(&schema), vec![vec![batch.clone()]])
.expect("memtable");
ctx.register_table("t", Arc::new(table)).expect("register");
let rows: usize = ctx
.sql("SELECT s, COUNT(*) FROM t GROUP BY s")
.await
.expect("plan")
.collect()
.await
.expect("collect")
.iter()
.map(|b| b.num_rows())
.sum();
assert_eq!(rows, GROUPS);
let peak = budget.peak();
assert!(peak > 0, "query {i} reserved against the budget");
assert_eq!(budget.used(), 0, "query {i} freed all its reservations");
if i > 0 {
assert_eq!(
peak, prev_peak,
"identical queries hold a flat peak (no leak)"
);
}
prev_peak = peak;
}
});
}
}