Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 13 additions & 6 deletions docs/source/contributor-guide/development.md
Original file line number Diff line number Diff line change
Expand Up @@ -43,12 +43,19 @@ onto a tokio worker thread and batches are delivered to the executor thread via
The executor thread parks in `blocking_recv()` until the next batch is ready. This avoids
busy-polling on I/O-bound workloads.

**JVM data source path (ScanExec present):** The executor thread calls `block_on()` and polls the
DataFusion stream directly, interleaving `pull_input_batches()` calls on `Poll::Pending` to feed
data from the JVM into ScanExec operators.

In both cases, DataFusion operators execute on **tokio worker threads**, not on the Spark executor
task thread. All Spark tasks on an executor share one tokio runtime.
**JVM data source path (ScanExec or ShuffleScanExec present):** The executor thread calls
`block_on()` and polls the DataFusion stream directly. On `Poll::Pending` it calls
`pull_input_batches()` to feed data from the JVM into the ScanExec and ShuffleScanExec operators,
whose streams register the poll's waker and are woken by the refill. The stream is polled with a
waker that also sets a flag. If the flag is still clear when the pull returns, the stream is
waiting on native I/O, and the thread parks until a waker fires instead of busy-polling. If it is
set, the loop polls again. It checks the flag rather than trusting the thread's parker because the
pull can run another Comet plan on the same thread, and that plan's `block_on()` shares the parker
and can consume the wake-up meant for the outer loop.

On the async I/O path, DataFusion operators execute on **tokio worker threads**. On the JVM data
source path, `block_on()` polls them on the Spark executor task thread, and any tasks they spawn
run on the shared runtime. All Spark tasks on an executor share one tokio runtime.

### Rules for native code

Expand Down
276 changes: 223 additions & 53 deletions native/core/src/execution/jni_api.rs
Original file line number Diff line number Diff line change
Expand Up @@ -77,8 +77,7 @@ use datafusion_spark::function::string::substring::SparkSubstring;
use datafusion_spark::function::url::try_url_decode::TryUrlDecode as SparkTryUrlDecode;
use datafusion_spark::function::url::url_decode::UrlDecode as SparkUrlDecode;
use datafusion_spark::function::url::url_encode::UrlEncode as SparkUrlEncode;
use futures::poll;
use futures::stream::StreamExt;
use futures::stream::{Stream, StreamExt};
use futures::FutureExt;
use jni::objects::JByteBuffer;
use jni::sys::{jlongArray, JNI_FALSE};
Expand All @@ -95,8 +94,13 @@ use parking_lot::Mutex;
use prost::Message;
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::{Duration, Instant};
use std::{sync::Arc, task::Poll};
use std::{
future::poll_fn,
sync::Arc,
task::{Context, Poll, Wake, Waker},
};
use tokio::runtime::{Handle, Runtime};
use tokio::sync::mpsc;

Expand Down Expand Up @@ -508,8 +512,6 @@ struct ExecutionContext {
pub metrics_update_interval: Option<Duration>,
// The last update time of metrics
pub metrics_last_update_time: Instant,
/// Counter to avoid checking time on every poll iteration (reduces syscalls)
pub poll_count_since_metrics_check: u32,
/// The time it took to create the native plan and configure the context
pub plan_creation_time: Duration,
/// DataFusion SessionContext
Expand Down Expand Up @@ -725,7 +727,6 @@ pub unsafe extern "system" fn Java_org_apache_comet_Native_createPlan(
metrics,
metrics_update_interval,
metrics_last_update_time: Instant::now(),
poll_count_since_metrics_check: 0,
plan_creation_time,
session_ctx: session,
debug_native,
Expand Down Expand Up @@ -1021,6 +1022,67 @@ fn pull_input_batches(exec_context: &mut ExecutionContext) -> Result<(), CometEr
})
}

/// Forwards a wake-up to the `block_on` task and records that it happened, so `next_batch` can
/// tell whether the stream was woken even if something else took the wake-up from the thread's
/// parker.
struct WakeFlag {
woken: AtomicBool,
parent: Waker,
}

impl Wake for WakeFlag {
fn wake(self: Arc<Self>) {
self.wake_by_ref()
}

fn wake_by_ref(self: &Arc<Self>) {
self.woken.store(true, Ordering::Release);
self.parent.wake_by_ref();
}
}

/// Drives `stream` to its next item. JVM-fed scans return `Pending` until `on_pending` refills
/// them, so every pending poll runs it, and the refill wakes the stream.
///
/// Each poll gives the stream a `WakeFlag` waker. If the flag is still clear when `on_pending`
/// returns, the stream is waiting on native I/O and `block_on` parks until it completes.
/// Otherwise `next_batch` wakes `block_on` itself so that it polls again at once. It can't rely on
/// the wake-up that set the flag, because `on_pending` can run another Comet plan on this thread,
/// as when a native writer's input is itself native. That plan's `block_on` shares this thread's
/// parker, which holds a single wake-up, and its park can take this one.
///
/// It polls again by yielding to `block_on` rather than looping here, so that each poll starts
/// with a fresh coop budget. A stream that has spent its budget wakes itself and returns
/// `Pending`, and a loop here would only get past that because `block_in_place` happens to leave
/// this thread's budget unconstrained.
async fn next_batch<S>(
stream: &mut S,
mut on_pending: impl FnMut() -> Result<(), CometError>,
) -> Result<Option<RecordBatch>, CometError>
where
S: Stream<Item = DataFusionResult<RecordBatch>> + Unpin,
{
poll_fn(|cx| {
let flag = Arc::new(WakeFlag {
woken: AtomicBool::new(false),
parent: cx.waker().clone(),
});
let waker = Waker::from(Arc::clone(&flag));
if let Poll::Ready(item) = stream.poll_next_unpin(&mut Context::from_waker(&waker)) {
return Poll::Ready(Ok(item.transpose()?));
}
// `on_pending` calls into the JVM, which can run another Comet plan on this thread.
// `block_in_place` exits the runtime context so that plan's `block_on` doesn't panic.
tokio::task::block_in_place(&mut on_pending)?;
if flag.woken.load(Ordering::Acquire) {
// Poll again at once: a nested `block_on` may have taken the wake-up.
cx.waker().wake_by_ref();
}
Poll::Pending
})
.await
}
Comment on lines +1044 to 1084

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Skipping the park after a pull is correct for the nesting we know about. It depends on every call that can run another Comet plan on this thread reporting itself through on_pending's return value, though. The closure in executePlan also calls update_metrics_on_interval, which calls into the JVM and doesn't report it. That's safe today because a metrics update doesn't run a plan, but nothing in the code enforces it, and a future JNI call in the closure would bring the hang back without failing a test. The pull isn't the only way into the JVM from this thread either. The threading section of development.md notes that memory pool operations call acquireMemory() over JNI on whatever thread the operator runs on, and on this path that's inside the stream poll, where no return value from on_pending can report it. Spark can make other consumers in the task spill to satisfy that request. Every Comet spill() returns 0 today, so this doesn't nest a plan now, but the loop's safety still depends on that staying true.

The lost wake comes from tokio keeping one wake token per thread. CachedParkThread::block_on polls and then parks on the thread-local CURRENT_PARKER with no per-call state (park.rs), while the block_in_place docs present Handle::block_on inside block_in_place as supported. The std Wake docs call out this case for their block_on example: "production-grade implementations will also need to handle intermediate calls to thread::unpark as well as nested invocations."

Could next_batch track its own wake-ups instead? It can poll the stream with a waker that sets a flag owned by this call and then forwards to the block_on waker. After the pull, it returns Pending only if the flag is still unset. A nested block_on can consume the thread's token but not the flag, so the loop no longer needs to know which calls can nest, and get_next_batch can keep the Result<(), CometError> signature from #6092. The cost is one Arc per next_batch call, which is once per output batch rather than once per poll. Roughly:

struct WakeFlag {
    woken: AtomicBool,
    parent: Waker,
}

impl Wake for WakeFlag {
    fn wake(self: Arc<Self>) {
        self.wake_by_ref()
    }

    fn wake_by_ref(self: &Arc<Self>) {
        self.woken.store(true, Ordering::Release);
        self.parent.wake_by_ref();
    }
}

async fn next_batch<S>(
    stream: &mut S,
    mut on_pending: impl FnMut() -> Result<(), CometError>,
) -> Result<Option<RecordBatch>, CometError>
where
    S: Stream<Item = DataFusionResult<RecordBatch>> + Unpin,
{
    poll_fn(|cx| {
        let flag = Arc::new(WakeFlag {
            woken: AtomicBool::new(false),
            parent: cx.waker().clone(),
        });
        let waker = Waker::from(Arc::clone(&flag));
        let mut stream_cx = Context::from_waker(&waker);
        loop {
            match stream.poll_next_unpin(&mut stream_cx) {
                Poll::Ready(item) => return Poll::Ready(Ok(item.transpose()?)),
                Poll::Pending => {
                    tokio::task::block_in_place(&mut on_pending)?;
                    if !flag.woken.swap(false, Ordering::AcqRel) {
                        return Poll::Pending;
                    }
                }
            }
        }
    })
    .await
}

This also replaces park_until_woken. I tried it in a local checkout of this branch (keeping the bool closures and ignoring the value). The three next_batch tests and refill_wakes_the_pending_poll_and_eof_stays_buffered pass. I also changed the nested block_on test's closure to return Ok(false), which is a nested plan that doesn't get reported. The sketch still returns in 0.10 s. The head commit fails after 10.0 s with "the park after the pull lost the stream's wake-up". If you go this way, the nested test's closure could return Ok(()), so that it covers a nested plan without depending on how the pull reports it. The paragraph this PR adds to development.md would then describe the flag instead of the rule about skipping the park after a JNI call.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks, this is a better design, and I've switched to it in 398b3f9. get_next_batch and pull_input_batches return () again, the nested block_on test's closure returns Ok(()), and the development.md paragraph describes the flag.

One change from your sketch: when the flag is set, next_batch wakes the block_on waker and returns Pending instead of looping inside the poll_fn, so each poll starts with a fresh coop budget. A stream that has spent its budget wakes itself and returns Pending, which sets the flag. The loop only gets past that because block_in_place leaves the budget unconstrained on a thread that isn't a tokio worker, since its Reset guard restores the budget only when there's a worker context. In a standalone tokio 1.53.1 program with block_in_place taken out, so the spent budget stayed spent, the loop re-polled more than a million times without progress, and the yielding version finished in 8 polls.

On acquireMemory inside the poll: a nested plan can't run there, because tokio panics with "Cannot start a runtime from within a runtime" on a block_on outside block_in_place. A spill that ran a Comet plan would error out rather than hang, so on_pending is the only place a nested plan can run today. If that changes, the flag also records a wake that arrives during the poll.

I also moved the elapsed-time assertion into the shared test helper, because the timeout's timer polls the future again when it fires, and that can finish it. Without the check, a WakeFlag that didn't forward to block_on passed the native I/O test in 10 s.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks, this is a better design, and I've switched to it in 398b3f9.

I checked the new version locally. With the flag.woken check forced to false, next_batch_polls_again_when_a_nested_block_on_took_the_wake_up fails after 10.0 s with "a wake was lost, and only the timeout's timer woke the task", and at the head commit it passes.

On acquireMemory inside the poll: a nested plan can't run there, because tokio panics with "Cannot start a runtime from within a runtime" on a block_on outside block_in_place.

That matches what I see. With block_in_place(&mut on_pending) replaced by a plain on_pending(), the same test panics with "Cannot start a runtime from within a runtime", so a nested plan inside the poll errors out rather than hanging.


/// Accept serialized query plan and the addresses of Arrow Arrays from Spark,
/// then execute the query. Return addresses of arrow vector.
/// # Safety
Expand Down Expand Up @@ -1158,54 +1220,31 @@ pub unsafe extern "system" fn Java_org_apache_comet_Native_executePlan(
}
}

// ScanExec path: busy-poll to interleave JVM batch pulls with stream polling
get_runtime().block_on(async {
loop {
let next_item = exec_context.stream.as_mut().unwrap().next();
let poll_output = poll!(next_item);

// Only check time/tracing every 100 polls to reduce overhead
exec_context.poll_count_since_metrics_check += 1;
if exec_context.poll_count_since_metrics_check >= 100 {
exec_context.poll_count_since_metrics_check = 0;
if let Some(interval) = exec_context.metrics_update_interval {
let now = Instant::now();
if now - exec_context.metrics_last_update_time >= interval {
update_metrics(env, exec_context)?;
exec_context.metrics_last_update_time = now;
}
}
if exec_context.tracing_enabled {
log_memory_usage(
&exec_context.tracing_memory_metric_name,
total_reserved_for_thread(exec_context.rust_thread_id) as u64,
);
}
}

match poll_output {
Poll::Ready(Some(output)) => {
return prepare_output(
env,
array_addrs,
schema_addrs,
output?,
exec_context.debug_native,
);
}
Poll::Ready(None) => {
log_plan_metrics(exec_context, stage_id, partition);
return Ok(-1);
}
Poll::Pending => {
// JNI call to pull batches from JVM into ScanExec operators.
// block_in_place lets tokio move other tasks off this worker
// while we wait for JVM data.
tokio::task::block_in_place(|| pull_input_batches(exec_context))?;
}
}
// ScanExec path: JVM-fed scans return `Pending` until `pull_input_batches` refills
// them and wakes the stream. A poll that is still pending, with nothing having woken
// the stream by the end of the pull, waits on native I/O, and `next_batch` parks
// until it completes.
let mut stream = exec_context.stream.take().unwrap();
let next = get_runtime().block_on(next_batch(&mut stream, || {
pull_input_batches(exec_context)?;
update_metrics_on_interval(env, exec_context)
}));
exec_context.stream = Some(stream);
let next = next?;
update_metrics_on_interval(env, exec_context)?;
match next {
Some(batch) => prepare_output(
env,
array_addrs,
schema_addrs,
batch,
exec_context.debug_native,
),
None => {
log_plan_metrics(exec_context, stage_id, partition);
Ok(-1)
}
})
}
});

if exec_context.tracing_enabled {
Expand Down Expand Up @@ -1267,6 +1306,30 @@ pub extern "system" fn Java_org_apache_comet_Native_releasePlan(
})
}

/// Runs `update_metrics` once the configured interval has passed and, with tracing on, samples
/// this thread's pool reservation at the same cadence.
fn update_metrics_on_interval(
env: &mut Env,
exec_context: &mut ExecutionContext,
) -> CometResult<()> {
let Some(interval) = exec_context.metrics_update_interval else {
return Ok(());
};
let now = Instant::now();
if now - exec_context.metrics_last_update_time < interval {
return Ok(());
}
update_metrics(env, exec_context)?;
exec_context.metrics_last_update_time = now;
if exec_context.tracing_enabled {
log_memory_usage(
&exec_context.tracing_memory_metric_name,
total_reserved_for_thread(exec_context.rust_thread_id) as u64,
);
}
Ok(())
}

fn update_metrics(env: &mut Env, exec_context: &mut ExecutionContext) -> CometResult<()> {
if let Some(native_query) = &exec_context.root_op {
let metrics = exec_context.metrics.as_obj();
Expand Down Expand Up @@ -1829,15 +1892,22 @@ pub unsafe extern "system" fn Java_org_apache_comet_Native_columnarToRowClose(
#[cfg(test)]
mod tests {
use super::*;
use crate::execution::operators::InputBatch;
use crate::execution::planner::TEST_EXEC_CONTEXT_ID;
use arrow::array::{ArrayRef, Int32Array};
use arrow::datatypes::{DataType, Field, Schema};
use datafusion::execution::memory_pool::{
MemoryConsumer, MemoryReservation, UnboundedMemoryPool,
};
use datafusion::execution::FunctionRegistry;
use datafusion::execution::TaskContext;
use datafusion::logical_expr::ReturnFieldArgs;
use datafusion::physical_plan::ExecutionPlan;
use datafusion_comet_proto::spark_expression;
use datafusion_comet_proto::spark_expression::{AggExpr, Count, Expr, Sum};
use datafusion_comet_proto::spark_operator::{HashAggregate, ShuffleWriter};
use std::cell::Cell;
use std::future::Future;

#[test]
fn skip_partial_eligibility_is_fail_closed() {
Expand Down Expand Up @@ -2314,4 +2384,104 @@ mod tests {
assert_eq!(ret.data_type(), &DataType::Int32, "length({input})");
}
}

fn single_worker_runtime() -> Runtime {
tokio::runtime::Builder::new_multi_thread()
.worker_threads(1)
.enable_all()
.build()
.unwrap()
}

/// Fails when a wake is lost. The timeout keeps a lost wake from hanging the suite, but when
/// its timer fires it polls `future` again, which can finish it, so this also fails when
/// `future` took more than five seconds.
async fn without_a_lost_wake<F: Future>(future: F) -> F::Output {
let start = Instant::now();
let output = tokio::time::timeout(Duration::from_secs(10), future)
.await
.expect("timed out: a wake was lost");
let elapsed = start.elapsed();
assert!(
elapsed < Duration::from_secs(5),
"took {elapsed:?}: a wake was lost, and only the timeout's timer woke the task"
);
output
}

#[test]
fn next_batch_parks_while_the_stream_waits_on_native_io() {
let batch = RecordBatch::new_empty(Arc::new(Schema::empty()));
let mut stream = futures::stream::once(async move {
tokio::time::sleep(Duration::from_millis(50)).await;
Ok::<_, DataFusionError>(batch)
})
.boxed();
let mut pulls = 0;
let next = single_worker_runtime()
.block_on(without_a_lost_wake(next_batch(&mut stream, || {
// Every JVM-fed scan already holds a batch, so the pull wakes nothing.
pulls += 1;
Ok(())
})))
.unwrap();
assert!(next.is_some());
assert!(
pulls < 5,
"the loop pulled {pulls} times during one 50 ms wait"
);
}

#[test]
fn next_batch_resumes_on_a_refill_and_stops_pulling_after_eof() {
let mut scan =
ScanExec::new(TEST_EXEC_CONTEXT_ID, None, "", vec![DataType::Int32]).unwrap();
let mut stream = scan.execute(0, Arc::new(TaskContext::default())).unwrap();
let column: ArrayRef = Arc::new(Int32Array::from(vec![1, 2, 3]));
let mut inputs = vec![InputBatch::new(vec![column], Some(3)), InputBatch::EOF].into_iter();
let pulls = Cell::new(0);
let mut pull = || {
pulls.set(pulls.get() + 1);
if let Some(input) = inputs.next() {
scan.set_input_batch(input);
}
Ok::<(), CometError>(())
};
// Only the refill's wake gets the stream polled again.
single_worker_runtime().block_on(without_a_lost_wake(async {
let first = next_batch(&mut stream, &mut pull).await.unwrap();
assert_eq!(first.unwrap().num_rows(), 3);
assert_eq!(pulls.get(), 1);
assert!(next_batch(&mut stream, &mut pull).await.unwrap().is_none());
assert_eq!(pulls.get(), 2);
assert!(next_batch(&mut stream, &mut pull).await.unwrap().is_none());
assert_eq!(pulls.get(), 2);
}));
}

/// A pull that runs another Comet plan on this thread, whose `block_on` parks until after the
/// stream's native I/O has completed. The nested park takes the I/O's wake-up from the
/// thread's parker, so `next_batch` has to have seen the wake some other way, or it parks
/// until `without_a_lost_wake`'s timer wakes it.
#[test]
fn next_batch_polls_again_when_a_nested_block_on_took_the_wake_up() {
let runtime = single_worker_runtime();
let handle = runtime.handle().clone();
let batch = RecordBatch::new_empty(Arc::new(Schema::empty()));
let mut stream = futures::stream::once(async move {
tokio::time::sleep(Duration::from_millis(20)).await;
Ok::<_, DataFusionError>(batch)
})
.boxed();
let mut pulls = 0;
let next = runtime
.block_on(without_a_lost_wake(next_batch(&mut stream, || {
pulls += 1;
handle.block_on(tokio::time::sleep(Duration::from_millis(100)));
Ok(())
})))
.unwrap();
assert!(next.is_some());
assert_eq!(pulls, 1);
}
}
Loading
Loading