use std::pin::Pin;
use std::task::{Context, Poll};
use arrow::array::RecordBatch;
use datafusion_common::Result;
use datafusion_physical_plan::metrics::PruningMetrics;
use datafusion_pruning::FilePruner;
use futures::{Stream, StreamExt, ready};
pub(super) struct EarlyStoppingStream<S> {
done: bool,
file_pruner: FilePruner,
files_ranges_pruned_statistics: PruningMetrics,
inner: S,
}
impl<S> EarlyStoppingStream<S> {
pub(super) fn new(
stream: S,
file_pruner: FilePruner,
files_ranges_pruned_statistics: PruningMetrics,
) -> Self {
Self {
done: false,
inner: stream,
file_pruner,
files_ranges_pruned_statistics,
}
}
}
impl<S> EarlyStoppingStream<S>
where
S: Stream<Item = Result<RecordBatch>> + Unpin,
{
fn check_prune(&mut self, input: Result<RecordBatch>) -> Result<Option<RecordBatch>> {
let batch = input?;
if self.file_pruner.should_prune()? {
self.files_ranges_pruned_statistics.add_pruned(1);
self.files_ranges_pruned_statistics.subtract_matched(1);
self.done = true;
Ok(None)
} else {
Ok(Some(batch))
}
}
}
impl<S> Stream for EarlyStoppingStream<S>
where
S: Stream<Item = Result<RecordBatch>> + Unpin,
{
type Item = Result<RecordBatch>;
fn poll_next(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
if self.done {
return Poll::Ready(None);
}
match ready!(self.inner.poll_next_unpin(cx)) {
None => {
self.done = true;
Poll::Ready(None)
}
Some(input_batch) => {
let output = self.check_prune(input_batch);
Poll::Ready(output.transpose())
}
}
}
}