diff --git a/Cargo.lock b/Cargo.lock index 8fa95485e..4c499b697 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2431,7 +2431,6 @@ dependencies = [ "serde_json", "sketches-ddsketch", "structopt", - "sysinfo", "tempfile", "tokio", "tonic", diff --git a/benchmarks/Cargo.toml b/benchmarks/Cargo.toml index 7b19f0ea9..00d7ac359 100644 --- a/benchmarks/Cargo.toml +++ b/benchmarks/Cargo.toml @@ -34,7 +34,6 @@ tempfile = "3" [dev-dependencies] criterion = "0.5" -sysinfo = "0.30" [build-dependencies] built = { version = "0.8", features = ["git2", "chrono"] } diff --git a/benchmarks/benches/broadcast_cache_scenarios.rs b/benchmarks/benches/broadcast_cache_scenarios.rs index d18791c85..dcf89fcab 100644 --- a/benchmarks/benches/broadcast_cache_scenarios.rs +++ b/benchmarks/benches/broadcast_cache_scenarios.rs @@ -5,6 +5,8 @@ use datafusion::arrow::record_batch::RecordBatch; use datafusion::common::Statistics; use datafusion::common::tree_node::TreeNodeRecursion; use datafusion::error::Result; +use datafusion::execution::memory_pool::{MemoryPool, PeakRecordingPool, UnboundedMemoryPool}; +use datafusion::execution::runtime_env::RuntimeEnvBuilder; use datafusion::execution::{SendableRecordBatchStream, TaskContext}; use datafusion::physical_expr::{EquivalenceProperties, PhysicalExpr}; use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType}; @@ -12,16 +14,11 @@ use datafusion::physical_plan::stream::RecordBatchStreamAdapter; use datafusion::physical_plan::{ DisplayAs, DisplayFormatType, ExecutionPlan, Partitioning, PlanProperties, }; -use datafusion::prelude::SessionContext; +use datafusion::prelude::{SessionConfig, SessionContext}; use datafusion_distributed::BroadcastExec; use futures::{StreamExt, stream}; -use std::sync::{ - Arc, - atomic::{AtomicBool, AtomicU64, Ordering}, -}; -use std::thread; +use std::sync::Arc; use std::time::{Duration, Instant}; -use sysinfo::{System, get_current_pid}; use tokio::runtime::Builder as RuntimeBuilder; #[derive(Clone, Copy)] @@ -47,56 +44,22 @@ struct Scenario { consumers: Vec, } -struct PeakRssSampler { - stop: Arc, - peak_kb: Arc, - handle: Option>, -} - -impl PeakRssSampler { - fn start(interval: Duration) -> Self { - let stop = Arc::new(AtomicBool::new(false)); - let peak_kb = Arc::new(AtomicU64::new(0)); - let stop_clone = Arc::clone(&stop); - let peak_clone = Arc::clone(&peak_kb); - let handle = thread::spawn(move || { - let mut sys = System::new(); - let pid = get_current_pid().expect("pid"); - while !stop_clone.load(Ordering::Relaxed) { - sys.refresh_process(pid); - if let Some(proc) = sys.process(pid) { - let mem_kb = proc.memory(); - peak_clone.fetch_max(mem_kb, Ordering::Relaxed); - } - thread::sleep(interval); - } - }); - Self { - stop, - peak_kb, - handle: Some(handle), - } - } - - fn stop(mut self) -> u64 { - self.stop.store(true, Ordering::Relaxed); - if let Some(handle) = self.handle.take() { - let _ = handle.join(); - } - self.peak_kb.load(Ordering::Relaxed) - } -} - #[derive(Debug)] struct SyntheticExec { schema: SchemaRef, partitions: usize, batches: Arc>>, + interval: Option, properties: Arc, } impl SyntheticExec { - fn new(schema: SchemaRef, partitions: usize, batches: Arc>>) -> Self { + fn new( + schema: SchemaRef, + partitions: usize, + batches: Arc>>, + interval: Option, + ) -> Self { let properties = Arc::new(PlanProperties::new( EquivalenceProperties::new(Arc::clone(&schema)), Partitioning::UnknownPartitioning(partitions), @@ -107,6 +70,7 @@ impl SyntheticExec { schema, partitions, batches, + interval, properties, } } @@ -158,12 +122,29 @@ impl ExecutionPlan for SyntheticExec { let batches = Arc::clone(&self.batches); let len = batches.len(); - let stream = stream::iter((0..len).map(move |idx| { - let batch = &batches[idx]; - Ok(batch.as_ref().clone()) - })); - - Ok(Box::pin(RecordBatchStreamAdapter::new(schema, stream))) + match self.interval { + None => { + let stream = stream::iter((0..len).map(move |idx| { + let batch = &batches[idx]; + Ok(batch.as_ref().clone()) + })); + Ok(Box::pin(RecordBatchStreamAdapter::new(schema, stream))) + } + Some(interval) => { + let stream = futures::stream::unfold(0usize, move |idx| { + let batches = Arc::clone(&batches); + async move { + if idx >= len { + return None; + } + tokio::time::sleep(interval).await; + let batch = batches[idx].as_ref().clone(); + Some((Ok(batch), idx + 1)) + } + }); + Ok(Box::pin(RecordBatchStreamAdapter::new(schema, stream))) + } + } } fn partition_statistics(&self, _partition: Option) -> Result> { @@ -197,23 +178,15 @@ async fn consume_partition( async fn run_scenario( scenario: &Scenario, - schema: Arc, - batches: Arc>>, + input: Arc, task_ctx: Arc, - sample_rss: bool, -) -> Result<(Duration, u64)> { - let input: Arc = Arc::new(SyntheticExec::new( - Arc::clone(&schema), - scenario.input_partitions, - batches, - )); - + recording_pool: Option>, +) -> Result<(Duration, Option<(usize, usize)>)> { let broadcast = Arc::new(BroadcastExec::new( Arc::clone(&input), scenario.consumer_tasks, )); - let sampler = sample_rss.then(|| PeakRssSampler::start(Duration::from_millis(25))); let start = Instant::now(); let mut join_set = tokio::task::JoinSet::new(); @@ -232,8 +205,50 @@ async fn run_scenario( } let elapsed = start.elapsed(); - let peak_kb = sampler.map(|s| s.stop()).unwrap_or(0); - Ok((elapsed, peak_kb)) + let memory = recording_pool.map(|pool| (pool.peak_reserved(), pool.reserved())); + Ok((elapsed, memory)) +} + +fn memory_context() -> Result<(Arc, Arc)> { + let inner_pool: Arc = Arc::new(UnboundedMemoryPool::default()); + let recording_pool = Arc::new(PeakRecordingPool::new(inner_pool)); + let runtime = RuntimeEnvBuilder::new() + .with_memory_pool(Arc::clone(&recording_pool) as Arc) + .build()?; + let task_ctx = + SessionContext::new_with_config_rt(SessionConfig::new(), Arc::new(runtime)).task_ctx(); + Ok((task_ctx, recording_pool)) +} + +fn synthetic_input( + scenario: &Scenario, + schema: &SchemaRef, + interval: Option, +) -> Arc { + let batches = (0..scenario.num_batches) + .map(|_| { + let array = UInt8Array::from(vec![0u8; scenario.rows_per_batch]); + let batch = + RecordBatch::try_new(Arc::clone(schema), vec![Arc::new(array)]).expect("batch"); + Arc::new(batch) + }) + .collect::>(); + Arc::new(SyntheticExec::new( + Arc::clone(schema), + scenario.input_partitions, + Arc::new(batches), + interval, + )) +} + +fn prebuilt_input(scenario: &Scenario, schema: &SchemaRef) -> Arc { + synthetic_input(scenario, schema, None) +} + +fn paced_input(scenario: &Scenario, schema: &SchemaRef) -> Arc { + // Consumers are started before the first batch and receive a scheduling window between + // batches. This is intentionally outside the timing measurement. + synthetic_input(scenario, schema, Some(Duration::from_millis(1))) } fn all_fast_consumers(output_partitions: usize) -> Vec { @@ -326,18 +341,8 @@ fn scenario_matrix() -> Vec { scenarios } -fn verbose_enabled() -> bool { - match std::env::var("BROADCAST_BENCH_VERBOSE") { - Ok(val) => { - let val = val.to_ascii_lowercase(); - val == "1" || val == "true" || val == "yes" - } - Err(_) => false, - } -} - -fn rss_enabled() -> bool { - match std::env::var("BROADCAST_BENCH_RSS") { +fn memory_enabled() -> bool { + match std::env::var("BROADCAST_BENCH_MEMORY") { Ok(val) => { let val = val.to_ascii_lowercase(); val == "1" || val == "true" || val == "yes" @@ -353,6 +358,8 @@ fn runtime_threads() -> Option { .filter(|threads| *threads > 0) } +const DIAGNOSTIC_RUNS: usize = 5; + fn bench_broadcast_cache(c: &mut Criterion) { let mut rt_builder = RuntimeBuilder::new_multi_thread(); if let Some(threads) = runtime_threads() { @@ -362,8 +369,7 @@ fn bench_broadcast_cache(c: &mut Criterion) { let mut group = c.benchmark_group("broadcast_cache_scenarios"); group.sample_size(10); - let verbose = verbose_enabled(); - let sample_rss = rss_enabled(); + let sample_memory = memory_enabled(); let task_ctx = SessionContext::new().task_ctx(); for scenario in scenario_matrix() { @@ -372,60 +378,69 @@ fn bench_broadcast_cache(c: &mut Criterion) { DataType::UInt8, false, )])); - let batches = (0..scenario.num_batches) - .map(|_| { - let data = vec![0u8; scenario.rows_per_batch]; - let array = UInt8Array::from(data); - let batch = RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(array)]) - .expect("batch"); - Arc::new(batch) - }) - .collect::>(); - let batches = Arc::new(batches); - + let input = prebuilt_input(&scenario, &schema); + let mut memory_peaks = Vec::new(); + let mut memory_residuals = Vec::new(); group.bench_function(BenchmarkId::new("scenario", scenario.name), |b| { b.iter_custom(|iters| { let mut total = Duration::ZERO; - let mut peaks = Vec::with_capacity(iters as usize); - for i in 0..iters { - let (elapsed, peak_kb) = rt + for _ in 0..iters { + let (elapsed, _) = rt .block_on(run_scenario( &scenario, - Arc::clone(&schema), - Arc::clone(&batches), + Arc::clone(&input), Arc::clone(&task_ctx), - sample_rss, + None, )) .expect("scenario"); - if verbose || sample_rss { - eprintln!( - "scenario={} iter={} peak_rss_kb={} elapsed_ms={}", - scenario.name, - i, - peak_kb, - elapsed.as_millis() - ); - } - peaks.push(peak_kb); total += elapsed; } - if sample_rss && !peaks.is_empty() { - peaks.sort_unstable(); - let min = peaks[0]; - let max = peaks[peaks.len() - 1]; - let median = peaks[peaks.len() / 2]; - eprintln!( - "scenario={} peak_rss_kb[min/median/max]={}/{}/{} runs={}", - scenario.name, - min, - median, - max, - peaks.len() - ); - } total }); }); + + if sample_memory { + for _ in 0..DIAGNOSTIC_RUNS { + let (diagnostic_task_ctx, pool) = memory_context().expect("memory context"); + let (_, memory) = rt + .block_on(run_scenario( + &scenario, + paced_input(&scenario, &schema), + diagnostic_task_ctx, + Some(pool), + )) + .expect("diagnostic scenario"); + if let Some((peak_reserved, residual_reserved)) = memory { + memory_peaks.push(peak_reserved); + memory_residuals.push(residual_reserved); + } + } + } + + if sample_memory && !memory_peaks.is_empty() { + memory_peaks.sort_unstable(); + memory_residuals.sort_unstable(); + let peak_min = memory_peaks[0]; + let peak_median = memory_peaks[memory_peaks.len() / 2]; + let peak_max = memory_peaks[memory_peaks.len() - 1]; + let residual_min = memory_residuals[0]; + let residual_median = memory_residuals[memory_residuals.len() / 2]; + let residual_max = memory_residuals[memory_residuals.len() - 1]; + eprintln!( + concat!( + "scenario={} reserved_bytes[peak_min/median/max]={}/{}/{} ", + "end_min/median/max={}/{}/{} runs={}" + ), + scenario.name, + peak_min, + peak_median, + peak_max, + residual_min, + residual_median, + residual_max, + memory_peaks.len() + ); + } } group.finish(); diff --git a/src/execution_plans/broadcast.rs b/src/execution_plans/broadcast.rs index 4d06d4d09..c5f278778 100644 --- a/src/execution_plans/broadcast.rs +++ b/src/execution_plans/broadcast.rs @@ -1,10 +1,10 @@ use crate::common::{OnceLockResult, require_one_child}; -use crossbeam_queue::SegQueue; use datafusion::arrow::datatypes::SchemaRef; +use datafusion::arrow::record_batch::RecordBatch; use datafusion::common::runtime::SpawnedTask; use datafusion::common::tree_node::TreeNodeRecursion; use datafusion::error::{DataFusionError, Result}; -use datafusion::execution::memory_pool::MemoryConsumer; +use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation}; use datafusion::execution::{SendableRecordBatchStream, TaskContext}; use datafusion::physical_expr::PhysicalExpr; use datafusion::physical_plan::stream::RecordBatchStreamAdapter; @@ -12,11 +12,15 @@ use datafusion::physical_plan::{ DisplayAs, DisplayFormatType, ExecutionPlan, Partitioning, PlanProperties, internal_err, }; use futures::{Stream, StreamExt}; +use std::collections::VecDeque; use std::fmt::Formatter; use std::pin::Pin; use std::sync::{Arc, Mutex, OnceLock}; use std::task::{Context, Poll}; use tokio_stream::wrappers::WatchStream; +use tokio_util::sync::CancellationToken; + +const RECLAIM_INTERVAL: usize = 8; /// [ExecutionPlan] that scales up partitions for network broadcasting. /// @@ -77,7 +81,9 @@ pub struct BroadcastExec { queues: Vec>, } -type StreamAndTask = (SegQueue, Arc>); +type BroadcastMessage = + std::result::Result<(RecordBatch, Arc), Arc>; +type StreamAndTask = (BroadcastReaders, Arc>); impl BroadcastExec { pub fn new(input: Arc, consumer_task_count: usize) -> Self { @@ -112,6 +118,18 @@ impl BroadcastExec { } } +impl Drop for BroadcastExec { + fn drop(&mut self) { + // The last plan reference is gone, so its unclaimed virtual partitions cannot + // execute. Active streams may still hold their producer task and queue alive. + for queue in &self.queues { + if let Some(Ok((readers, _task))) = queue.get() { + readers.release_pending(); + } + } + } +} + impl DisplayAs for BroadcastExec { fn fmt_as(&self, _t: DisplayFormatType, f: &mut Formatter) -> std::fmt::Result { let input_partition_count = self.input_partition_count(); @@ -160,57 +178,61 @@ impl ExecutionPlan for BroadcastExec { partition: usize, context: Arc, ) -> Result { - let real_partition = partition % self.input_partition_count(); + let input_partition_count = self.input_partition_count(); + let real_partition = partition % input_partition_count; + let consumer_task = partition / input_partition_count; let input = Arc::clone(&self.input); let queue_or_err = self.queues[real_partition].get_or_init(|| { - let queue = BroadcastQueue::new(); - let consumers = SegQueue::new(); - for _ in 0..self.consumer_task_count { - consumers.push(Box::pin(RecordBatchStreamAdapter::new( - self.schema(), - queue.new_consumer().map(|msg| match msg { - Ok((batch, _reservation)) => Ok(batch), - Err(e) => Err(DataFusionError::Shared(e)), - }), - )) as SendableRecordBatchStream); - } + let queue = BroadcastQueue::new(self.consumer_task_count); + let readers = queue.readers(); let pool = Arc::clone(context.memory_pool()); let mut stream = input.execute(real_partition, context).map_err(Arc::new)?; + let cancel = queue.shared.cancel.clone(); let task = SpawnedTask::spawn(async move { let mem_consumer = MemoryConsumer::new(format!("BroadcastExec[{real_partition}]")); - while let Some(msg) = stream.next().await { + loop { + let msg = tokio::select! { + _ = cancel.cancelled() => break, + msg = stream.next() => msg, + }; + let Some(msg) = msg else { break }; match msg { Ok(record_batch) => { let reservation = mem_consumer.clone_with_new_id().register(&pool); reservation.grow(record_batch.get_array_memory_size()); - queue.push(Ok((record_batch, Arc::new(reservation)))); + if !queue.push(Ok((record_batch, Arc::new(reservation)))) { + // If there are no remaining readers, short-circuit. + break; + } } Err(err) => { - queue.push(Err(Arc::new(err))); + let _ = queue.push(Err(Arc::new(err))); break; } } } }); - Ok::<_, Arc>((consumers, Arc::new(task))) + Ok::<_, Arc>((readers, Arc::new(task))) }); let (consumer, task) = match queue_or_err { - Ok((consumers, task)) => (consumers.pop(), Arc::clone(task)), + Ok((readers, task)) => (readers.claim(consumer_task)?, Arc::clone(task)), Err(err) => return Err(DataFusionError::Shared(Arc::clone(err))), }; - let Some(consumer) = consumer else { - return internal_err!("Too many consumers for real partition {real_partition}"); - }; Ok(Box::pin(RecordBatchStreamAdapter::new( self.schema(), - consumer.inspect(move |_| { - let _ = &task; - }), + consumer + .map(|msg| match msg { + Ok((batch, _reservation)) => Ok(batch), + Err(e) => Err(DataFusionError::Shared(e)), + }) + .inspect(move |_| { + let _ = &task; + }), ))) } @@ -219,92 +241,338 @@ impl ExecutionPlan for BroadcastExec { } } -#[derive(Debug, Clone, Copy)] -struct BroadcastState { - len: usize, +/// Represents the queue for a single real partition, which multiple virtual partitions may share. +/// Assume we have 4 consumer tasks each represented as a reader, this could create a situation as +/// such: +/// +/// ```text +/// +/// base_sequence tail_sequence +/// │ │ +/// ▼ ▼ +/// ┌ ─ ─┌────┬────┬────┐ ┌────┬────┐ +/// entries: e0 │ e1 │ e2 │ e3 │ ... │eN-1│ eN │ +/// └ ─ ─└────┴────┴────┘ └────┴────┘ +/// ▲ ▲ +/// ┌──────┘ ┌────────────────┴────┐ +/// │ │ │ +/// ┌──────────┬──────────┬──────────┬──────────┐ +/// readers: │ r0: │ r1: │ r2: │ r3: │ +/// │Reading(2)│Reading(N)│ Released │Reading(N)│ +/// └──────────┴──────────┴──────────┴──────────┘ +/// ``` +/// +/// The `base_sequence` represents the first retained entry while `tail_sequence` indicates the next +/// append position. In this example, entry `e0` is shown using a dashed line because it has already +/// been evicted and `base_sequence` now points at `e1`. Every retaining reader has advanced past +/// `e1`, so `e1` is reclaimable but remains temporarily retained until an append reaches the next +/// `RECLAIM_INTERVAL` boundary. At that boundary, the queue removes entries before the smallest +/// `next_sequence` (the most lagging retaining reader, `r0` here) and advances `base_sequence`: +/// +/// ```text +/// base_sequence tail_sequence +/// │ │ +/// ▼ ▼ +/// ┌ ─ ─┌ ─ ─┌────┬────┐ ┌────┬────┬────┐ +/// entries: e0 │ e1 │ e2 │ e3 │ ... │eN-1│ eN │eN+1│ +/// └ ─ ─└ ─ ─└────┴────┘ └────┴────┴────┘ +/// ▲ ▲ ▲ +/// ┌──────┘ ┌────────────────┘ │ +/// │ │ │ +/// ┌──────────┬──────────┬──────────┬────────────┐ +/// readers: │ r0: │ r1: │ r2: │ r3: │ +/// │Reading(2)│Reading(N)│ Released │Reading(N+1)│ +/// └──────────┴──────────┴──────────┴────────────┘ +/// ``` +#[derive(Debug)] +struct QueueState { + entries: VecDeque, + base_sequence: usize, + tail_sequence: usize, + readers: Box<[ReaderSlot]>, + retaining_readers: usize, closed: bool, } +#[derive(Debug)] +enum ReaderSlot { + Pending, + Reading { next_sequence: usize }, + Released, +} + +/// Shared state and signals for one real input partition. +/// +/// The producer task owns the input stream and queue handle while every consumer stream shares the +/// same queue state but has its own reader cursor and notification receiver. There are two flows +/// this is responsible for: queue updates and cancellation. +/// +/// ## Queue Push Flow: +/// +/// ```text +/// ┌───────────────────reads──────────────────┐ +/// │ ┌─────────(each acquire lock)─────────┐ │ +/// │ │ ┌─────────────────────────────────┐ │ │ +/// │ │ │ │ │ │ +/// │ │ │ ┌────────────────┐ │ │ │ +/// ┌─────────────────────────┐ │ │ │ │ │ │ │ │ +/// │BroadcastShared │ │ │ │ ┌────▶│ Consumer 0 │──┘ │ │ +/// │ ┌─────────────────────┐◀┼────────────────────┘ │ │ │ │ │ │ │ +/// ┌───────┼▶│ Mutex(QueueState) │◀┼───────────────────────┘ │ │ └────────────────┘ │ │ +/// push │ └─────────────────────┘◀┼─────────────────────────┘ │ │ │ +/// (acquires lock)│ │notify │ │ │ │ +/// ┌────────────────┐ │ │ ▼ │ │ ┌────────────────┐ │ │ +/// │ │ │ │ ┌─────────────────────┐ │ notify ┌───────────────┐ │ │ │ │ │ +/// │ Producer │────────┘ │ │ Sender │ ├──signals──▶│ Watch Channel │─notify──▶│ Consumer 1 │────┘ │ +/// │ │ │ └─────────────────────┘ │ └───────────────┘ signal │ │ │ +/// └────────────────┘ │ ┌─────────────────────┐ │ │ └────────────────┘ │ +/// │ │ CancellationToken │ │ │ │ +/// │ └─────────────────────┘ │ │ ... │ +/// └─────────────────────────┘ │ │ +/// │ ┌────────────────┐ │ +/// │ │ │ │ +/// └────▶│ Consumer N │──────┘ +/// │ │ +/// └────────────────┘ +/// ``` +/// +/// A consumer stream has access to the shared `Arc`, but only `queue_state` is locked so its +/// `poll_next` holds that mutex while it reads an entry. The producer follows the same rule, +/// it appends under the mutex, then sends the notification after unlocking. +/// +/// ## Cancellation flow: +/// +/// ```text +/// ┌────────────────┐ +/// │ │───────────────────────┐ ┌─────────────────────────┐ +/// │ Consumer 0 │─────────────┐ │ │BroadcastShared │ ┌────────────close ───────┐ +/// │ │◀───────┐ │ │ release reader │ ┌─────────────────────┐ │ │ (acquires lock) │ +/// └────────────────┘ │ │ └────(acquires lock)───┼▶│ Mutex(QueueState) │◀┼───────┘ │ +/// │ │ │ └─────────────────────┘ │ ┌────────────────┐ │ +/// │ │ │ │ notify on │ │ │ │ +/// ┌────────────────┐ │ │ │ ▼ close │ ┌──────▶│ Producer │──┘ +/// │ │ │ │ ┌───────────────┐ notify │ ┌─────────────────────┐ │ │ │ │ +/// │ Consumer 1 │◀────notify──┼───│ Watch Channel │◀──signals──┼─│ Sender │ │ cancel └────────────────┘ +/// │ │ signal │ └───────────────┘ │ └─────────────────────┘ │ signal +/// └────────────────┘ │ │ │ ┌─────────────────────┐ │ │ +/// │ └────────────cancel──────────────┼▶│ CancellationToken │─┼─────┘ +/// ... │ │ └─────────────────────┘ │ +/// │ └─────────────────────────┘ +/// ┌────────────────┐ │ +/// │ │ │ +/// │ Consumer N │◀───────┘ +/// │ │ +/// └────────────────┘ +/// ``` +/// +/// In this case, the consumer initiates the action by mutating the queue state to release itself. +/// Only in the case that the last reader has dropped the consumer will also set the +/// `CancellationToken` to tell the producer to close the queue. +/// +/// Also, a consumer stream may outlive the `BroadcastExec`. In this case, the plan's `Drop` releases +/// only still `Pending` readers while active `BroadcastConsumer`s keep the shared state alive, and +/// keeps the producer task alive until those streams finish or are dropped. +#[derive(Debug)] +struct BroadcastShared { + queue_state: Mutex>, + notify: tokio::sync::watch::Sender<()>, + cancel: CancellationToken, +} + #[derive(Debug)] struct BroadcastQueue { - entries: Arc>>, - notify: tokio::sync::watch::Sender, + shared: Arc>, +} + +#[derive(Debug)] +struct BroadcastReaders { + shared: Arc>, } impl BroadcastQueue { - fn new() -> Self { - let (notify, _rx) = tokio::sync::watch::channel(BroadcastState { - len: 0, - closed: false, - }); + fn new(expected_readers: usize) -> Self { + let (notify, _rx) = tokio::sync::watch::channel(()); Self { - entries: Arc::new(Mutex::new(vec![])), - notify, + shared: Arc::new(BroadcastShared { + queue_state: Mutex::new(QueueState { + entries: VecDeque::new(), + base_sequence: 0, + tail_sequence: 0, + readers: (0..expected_readers).map(|_| ReaderSlot::Pending).collect(), + retaining_readers: expected_readers, + closed: false, + }), + notify, + cancel: CancellationToken::new(), + }), } } - fn new_consumer(&self) -> BroadcastConsumer { - let rx = self.notify.subscribe(); - let state = *rx.borrow(); - BroadcastConsumer { - index: 0, - entries: Arc::clone(&self.entries), - notify: WatchStream::new(rx), - state, + fn readers(&self) -> BroadcastReaders { + BroadcastReaders { + shared: Arc::clone(&self.shared), } } - fn push(&self, entry: T) { - let len = { - let mut entries = self.entries.lock().unwrap(); - entries.push(entry); - entries.len() + /// Appends a value to the entry queue and increments `tail_sequence`. Every `RECLAIM_INTERVAL` + /// calls, this checks for reclaimable entries in the queue. + /// + /// This method will not append the value and returns `false` if no retaining readers remain. + fn push(&self, value: T) -> bool { + let reclaimed = { + let mut queue_state = self.shared.queue_state.lock().unwrap(); + + if queue_state.retaining_readers == 0 { + return false; + } + + queue_state.entries.push_back(value); + queue_state.tail_sequence += 1; + if queue_state.tail_sequence.is_multiple_of(RECLAIM_INTERVAL) { + Self::reclaim_processed_entries(&mut queue_state) + } else { + Vec::new() + } }; - let mut state = *self.notify.borrow(); - state.len = len; - let _ = self.notify.send(state); + + drop(reclaimed); + + self.shared.notify.send_replace(()); + true + } + + /// Frees all entries in the queue that have been processed and updates `base_sequence` to point + /// at the first non-freeable position. + fn reclaim_processed_entries(queue_state: &mut QueueState) -> Vec { + let minimum_sequence = queue_state + .readers + .iter() + .filter_map(|reader| match reader { + ReaderSlot::Pending => Some(0), + ReaderSlot::Reading { next_sequence } => Some(*next_sequence), + ReaderSlot::Released => None, + }) + .min() + .unwrap_or(queue_state.tail_sequence); + + let reclaim_count = minimum_sequence - queue_state.base_sequence; + queue_state.base_sequence = minimum_sequence; + queue_state.entries.drain(..reclaim_count).collect() } } impl Drop for BroadcastQueue { fn drop(&mut self) { - let mut state = *self.notify.borrow(); - state.closed = true; - let _ = self.notify.send(state); + { + let mut state = self.shared.queue_state.lock().unwrap(); + state.closed = true; + } + self.shared.notify.send_replace(()); + } +} + +impl BroadcastReaders { + /// Creates a new `BroadcastConsumer` for a given consumer task and claims its reader slot, + /// starting at sequence zero. + /// + /// Returns an error if the consumer task is out of range or if the same reader is claimed more + /// than once. + fn claim(&self, consumer_task: usize) -> Result> { + let rx = self.shared.notify.subscribe(); + let mut state = self.shared.queue_state.lock().unwrap(); + let Some(reader) = state.readers.get_mut(consumer_task) else { + return internal_err!("broadcast consumer {consumer_task} is out of range"); + }; + match reader { + ReaderSlot::Pending => *reader = ReaderSlot::Reading { next_sequence: 0 }, + ReaderSlot::Reading { .. } | ReaderSlot::Released => { + return internal_err!("broadcast consumer {consumer_task} cannot execute twice"); + } + } + Ok(BroadcastConsumer { + consumer_id: consumer_task, + shared: Arc::clone(&self.shared), + notify: WatchStream::new(rx), + }) + } + + /// Sets all `ReaderSlot::Pending` readers to `ReaderSlot::Released`. This also cleans up newly + /// freeable entries and cancels the producer if all readers are released. + fn release_pending(&self) { + let (no_readers_remain, reclaimed) = { + let mut state = self.shared.queue_state.lock().unwrap(); + let mut released = 0; + for reader in &mut state.readers { + if matches!(reader, ReaderSlot::Pending) { + *reader = ReaderSlot::Released; + released += 1; + } + } + state.retaining_readers -= released; + let reclaimed = BroadcastQueue::::reclaim_processed_entries(&mut state); + (state.retaining_readers == 0, reclaimed) + }; + + drop(reclaimed); + + if no_readers_remain { + self.shared.cancel.cancel(); + } } } /// A consumer stream that reads from the broadcast queue. -struct BroadcastConsumer { - index: usize, - entries: Arc>>, - notify: WatchStream, - state: BroadcastState, +struct BroadcastConsumer { + consumer_id: usize, + shared: Arc>, + notify: WatchStream<()>, } impl Stream for BroadcastConsumer { type Item = T; + /// Poll the next value from the stream reading from the shared entry queue. fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { loop { - if self.index < self.state.len { - let entry = self.entries.lock().unwrap().get(self.index).cloned(); - if let Some(v) = entry { - self.index += 1; - return Poll::Ready(Some(v)); + let (value, closed) = { + let mut state = self.shared.queue_state.lock().unwrap(); + let next_sequence = match state.readers[self.consumer_id] { + ReaderSlot::Reading { next_sequence } => next_sequence, + ReaderSlot::Released => return Poll::Ready(None), + ReaderSlot::Pending => unreachable!("an unclaimed consumer was polled"), + }; + if next_sequence < state.tail_sequence { + let offset = next_sequence + .checked_sub(state.base_sequence) + .expect("broadcast consumer fell behind evicted entries"); + let value = state + .entries + .get(offset) + .expect("broadcast entry sequence was not retained") + .clone(); + state.readers[self.consumer_id] = ReaderSlot::Reading { + next_sequence: next_sequence + 1, + }; + (Some(value), false) + } else { + (None, state.closed) } + }; + + if let Some(value) = value { + return Poll::Ready(Some(value)); } - if self.state.closed { + if closed { + self.shared.release(self.consumer_id); return Poll::Ready(None); } match Pin::new(&mut self.notify).poll_next(cx) { - Poll::Ready(Some(state)) => { - self.state = state; - } + Poll::Ready(Some(_)) => continue, Poll::Ready(None) => { - self.state.closed = true; + self.shared.release(self.consumer_id); + return Poll::Ready(None); } Poll::Pending => return Poll::Pending, } @@ -312,6 +580,33 @@ impl Stream for BroadcastConsumer { } } +impl BroadcastShared { + fn release(&self, consumer_id: usize) { + let (no_readers_remain, reclaimed) = { + let mut state = self.queue_state.lock().unwrap(); + if matches!(state.readers[consumer_id], ReaderSlot::Released) { + return; + } + state.readers[consumer_id] = ReaderSlot::Released; + state.retaining_readers -= 1; + let reclaimed = BroadcastQueue::::reclaim_processed_entries(&mut state); + (state.retaining_readers == 0, reclaimed) + }; + + drop(reclaimed); + + if no_readers_remain { + self.cancel.cancel(); + } + } +} + +impl Drop for BroadcastConsumer { + fn drop(&mut self) { + self.shared.release(self.consumer_id); + } +} + #[cfg(test)] mod tests { use super::*; @@ -337,6 +632,114 @@ mod tests { } } + fn buffered_len(queue: &BroadcastQueue) -> usize { + queue.shared.queue_state.lock().unwrap().entries.len() + } + + fn sequence_bounds(queue: &BroadcastQueue) -> (usize, usize) { + let state = queue.shared.queue_state.lock().unwrap(); + (state.base_sequence, state.tail_sequence) + } + + #[tokio::test] + async fn broadcast_queue_evicts_consumed_prefix() { + let queue = BroadcastQueue::new(2); + let mut consumer0 = queue.readers().claim(0).expect("consumer 0 registration"); + let mut consumer1 = queue.readers().claim(1).expect("consumer 1 registration"); + + queue.push(10); + queue.push(20); + assert_eq!(buffered_len(&queue), 2); + + assert_eq!(consumer0.next().await, Some(10)); + assert_eq!(buffered_len(&queue), 2); + assert_eq!(consumer1.next().await, Some(10)); + assert_eq!(sequence_bounds(&queue), (0, 2)); + + assert_eq!(consumer0.next().await, Some(20)); + assert_eq!(buffered_len(&queue), 2); + assert_eq!(consumer1.next().await, Some(20)); + assert_eq!(buffered_len(&queue), 2); + + // The eighth append observes that consumers have advanced and removes the consumed prefix. + for value in [30, 40, 50, 60, 70, 80] { + queue.push(value); + } + assert_eq!(sequence_bounds(&queue), (2, 8)); + assert_eq!(buffered_len(&queue), 6); + + drop(consumer0); + drop(consumer1); + assert_eq!(buffered_len(&queue), 0); + } + + #[tokio::test] + async fn broadcast_queue_drop_releases_unread_entries() { + let queue = BroadcastQueue::new(2); + let mut consumer0 = queue.readers().claim(0).expect("consumer 0 registration"); + let mut consumer1 = queue.readers().claim(1).expect("consumer 1 registration"); + + queue.push(10); + queue.push(20); + assert_eq!(consumer0.next().await, Some(10)); + drop(consumer0); + assert!(!queue.shared.cancel.is_cancelled()); + + // The dropped consumer must no longer pin the unread suffix, including entries produced + // after the first reader is released. + queue.push(30); + assert_eq!(consumer1.next().await, Some(10)); + assert_eq!(consumer1.next().await, Some(20)); + assert_eq!(consumer1.next().await, Some(30)); + drop(consumer1); + assert!(queue.shared.cancel.is_cancelled()); + assert_eq!(buffered_len(&queue), 0); + assert!(!queue.push(20)); + } + + #[tokio::test] + async fn broadcast_queue_releases_reader_at_eof() { + let queue = BroadcastQueue::new(1); + let shared = Arc::clone(&queue.shared); + let mut consumer = queue.readers().claim(0).expect("consumer registration"); + + assert!(queue.push(10)); + assert_eq!(consumer.next().await, Some(10)); + drop(queue); + + assert_eq!(consumer.next().await, None); + assert_eq!(shared.queue_state.lock().unwrap().entries.len(), 0); + assert!(shared.cancel.is_cancelled()); + assert_eq!(consumer.next().await, None); + } + + #[tokio::test] + async fn broadcast_exec_releases_pending_readers_when_plan_drops() -> Result<()> { + let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])); + let input = Arc::new(MockExec::new_partitioned(vec![vec![]], Arc::clone(&schema))); + let broadcast = BroadcastExec::new(input, 2); + let task_ctx = SessionContext::new().task_ctx(); + + let stream = broadcast.execute(0, task_ctx)?; + let shared = Arc::clone( + &broadcast.queues[0] + .get() + .expect("initialized queue") + .as_ref() + .expect("queue initialization") + .0 + .shared, + ); + + drop(broadcast); + assert_eq!(shared.queue_state.lock().unwrap().retaining_readers, 1); + drop(stream); + assert_eq!(shared.queue_state.lock().unwrap().retaining_readers, 0); + assert!(shared.cancel.is_cancelled()); + + Ok(()) + } + #[tokio::test] async fn broadcast_exec_reuses_queue_for_virtual_partitions() -> Result<()> { let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])); @@ -427,6 +830,23 @@ mod tests { Ok(()) } + #[tokio::test] + async fn broadcast_exec_claims_virtual_partition_once() -> Result<()> { + let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])); + let input = Arc::new(MockExec::new_partitioned(vec![vec![]], Arc::clone(&schema))); + let broadcast = BroadcastExec::new(input, 2); + let task_ctx = SessionContext::new().task_ctx(); + + let _stream = broadcast.execute(0, Arc::clone(&task_ctx))?; + let err = broadcast + .execute(0, task_ctx) + .err() + .expect("duplicate execute"); + assert!(err.to_string().contains("consumer 0 cannot execute twice")); + + Ok(()) + } + #[tokio::test] async fn broadcast_exec_queue_survives_cancellation() -> Result<()> { let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)]));