nathanb9 commented on code in PR #25580:
URL: https://github.com/apache/datafusion/pull/25580#discussion_r4141520894


##########
datafusion/sqllogictest/test_files/cte_materialized.slt:
##########
@@ -0,0 +1,437 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+

Review Comment:
   I think there's a pre-existing issue that shows up with `WITH RECURSIVE ... 
AS MATERIALIZED`, and I cut #25893 for it. If the recursive term widens the 
type (e.g. `n + 0.5`), it gets silently cast back to the static term's type, 
and the query can loop forever. It happens on main with or without 
`MATERIALIZED`, so it's not from this PR. Just a heads-up in case you add 
recursive tests that change types.



##########
datafusion/sqllogictest/test_files/cte_materialized.slt:
##########
@@ -0,0 +1,437 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+
+# Tests for `WITH x AS MATERIALIZED (...)`: the CTE body runs once and every
+# reference reads its buffered output.
+
+# sqlparser only parses `[NOT] MATERIALIZED` for the PostgreSQL dialect.
+statement ok
+set datafusion.sql_parser.dialect = 'PostgreSQL';
+
+statement ok
+set datafusion.explain.logical_plan_only = false;
+
+statement ok
+CREATE TABLE t(k INT, v INT) AS VALUES (1, 10), (2, 20), (3, 30), (4, 40);
+
+# A volatile body is evaluated once, so both references see the same value.
+query I
+WITH r AS MATERIALIZED (SELECT random() AS x)
+SELECT count(DISTINCT x) FROM (SELECT x FROM r UNION ALL SELECT x FROM r)
+----
+1
+
+# Without MATERIALIZED the body is inlined and evaluated once per reference.
+query I
+WITH r AS NOT MATERIALIZED (SELECT random() AS x)

Review Comment:
   FYI: Postgres ignores `NOT MATERIALIZED` when the CTE has volatile functions 
if we are trying to follow that closely. But to me it also makes sense to allow 
this control for a user



##########
datafusion/physical-plan/src/materialized_cte.rs:
##########
@@ -0,0 +1,829 @@
+// Licensed to the Apache Software Foundation (ASF) under one
+// or more contributor license agreements.  See the NOTICE file
+// distributed with this work for additional information
+// regarding copyright ownership.  The ASF licenses this file
+// to you under the Apache License, Version 2.0 (the
+// "License"); you may not use this file except in compliance
+// with the License.  You may obtain a copy of the License at
+//
+//   http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing,
+// software distributed under the License is distributed on an
+// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+// KIND, either express or implied.  See the License for the
+// specific language governing permissions and limitations
+// under the License.
+
+//! Execution plans for `WITH x AS MATERIALIZED (...)`.
+//!
+//! [`MaterializedCteExec`] owns the CTE body and the continuation (the rest of
+//! the query). Every [`MaterializedCteScanExec`] in the continuation shares 
one
+//! [`MaterializedCteBuffer`] with it. The first consumer to execute runs every
+//! partition of the body to completion and buffers the output; every scan then
+//! replays that buffer.
+//!
+//! The body is fully materialized before any scan yields a row. This is a
+//! pipeline break, but it cannot deadlock when two consumers of the same CTE
+//! are read at different rates (for example the build and probe side of one
+//! hash join), which a bounded fan-out channel can.
+//!
+//! Buffered batches are accounted in the memory pool. When a reservation
+//! fails, the rest of that partition is written to a spill file, and the scans
+//! read the in-memory prefix and then the spill files, so row order within a
+//! partition is kept.
+//!
+//! A scan reads a spill file one batch at a time, without read-ahead. After
+//! the body is buffered, memory for one decoded batch per concurrent reader is
+//! reserved. If the pool cannot grant it, in-memory batches are moved to spill
+//! files until it can.
+
+use std::fmt;
+use std::pin::Pin;
+use std::sync::Arc;
+use std::task::{Context, Poll};
+
+use arrow::datatypes::SchemaRef;
+use arrow::record_batch::RecordBatch;
+use datafusion_common::config::ConfigOptions;
+use datafusion_common::tree_node::TreeNodeRecursion;
+use datafusion_common::{Result, Statistics, internal_err};
+use datafusion_common_runtime::JoinSet;
+use datafusion_execution::memory_pool::{MemoryConsumer, MemoryReservation};
+use datafusion_execution::{SendableRecordBatchStream, SpillFile, TaskContext};
+use datafusion_physical_expr::{EquivalenceProperties, Partitioning, 
PhysicalExpr};
+use futures::stream::BoxStream;
+use futures::{Stream, StreamExt, TryStreamExt};
+use parking_lot::Mutex;
+
+use crate::coop::cooperative;
+use crate::execution_plan::{
+    Boundedness, CardinalityEffect, EmissionType, ExecutionPlan, 
ExecutionPlanProperties,
+    PlanProperties, SchedulingType,
+};
+use crate::filter_pushdown::{
+    ChildFilterDescription, ChildPushdownResult, FilterDescription, 
FilterPushdownPhase,
+    FilterPushdownPropagation,
+};
+use crate::joins::utils::{OnceAsync, OnceFut};
+use crate::metrics::{
+    BaselineMetrics, ExecutionPlanMetricsSet, MetricBuilder, MetricsSet, 
SpillMetrics,
+};
+use crate::spill::spill_manager::SpillManager;
+use crate::statistics::{ChildStats, StatisticsArgs};
+use crate::stream::RecordBatchStreamAdapter;
+use crate::{
+    ChildrenPropertiesMode, DisplayAs, DisplayFormatType, 
ReplaceChildrenOptions,
+};
+
+/// The buffered output of one body partition: an in-memory prefix followed by
+/// spill files with the remaining batches, in order.
+struct BufferedPartition {
+    batches: Vec<RecordBatch>,
+    spill_files: Vec<Arc<dyn SpillFile>>,
+    /// The memory size of the largest batch in `spill_files` when decoded.
+    max_spilled_batch_size: usize,
+    /// Holds `batches`. Released when the buffer is dropped.
+    reservation: MemoryReservation,
+}
+
+impl BufferedPartition {
+    /// Move in-memory batches from the end of `batches` to a new spill file,
+    /// until at least `bytes` are released or no in-memory batch is left.
+    /// The new file is read before the existing spill files, so the order of
+    /// the partition is kept.
+    fn spill_in_memory_suffix(
+        &mut self,
+        bytes: usize,
+        spill_manager: &SpillManager,
+    ) -> Result<()> {
+        let mut split = self.batches.len();
+        let mut released = 0;
+        while split > 0 && released < bytes {
+            split -= 1;
+            released += self.batches[split].get_array_memory_size();
+        }
+        let suffix = self.batches.split_off(split);
+        let mut file = 
spill_manager.create_in_progress_file("MaterializedCte")?;
+        for batch in &suffix {
+            let size = file.append_batch(batch)?;
+            self.max_spilled_batch_size = 
self.max_spilled_batch_size.max(size);
+        }
+        if let Some(file) = file.finish()? {
+            self.spill_files.insert(0, file);
+        }
+        self.reservation.shrink(released);
+        Ok(())
+    }
+}
+
+/// The materialized output of a CTE body.
+struct MaterializedOutput {
+    partitions: Vec<BufferedPartition>,
+    spill_manager: SpillManager,
+    /// Memory for the batches that the scans decode from the spill files.
+    /// Released when the buffer is dropped.
+    _replay_reservation: MemoryReservation,
+}
+
+/// The memory that the scans need to read the spill files of `partitions`.
+///
+/// Partition `p` of a scan with `n` partitions reads the body partitions
+/// `p, p + n, ...` one after the other, with at most one decoded batch in
+/// memory. All scan partitions can read at the same time.
+fn replay_memory(partitions: &[BufferedPartition], scan_partitions: &[usize]) 
-> usize {
+    scan_partitions
+        .iter()
+        .map(|&n| {
+            (0..n)
+                .map(|p| {
+                    partitions
+                        .iter()
+                        .skip(p)
+                        .step_by(n)
+                        .map(|b| b.max_spilled_batch_size)
+                        .max()
+                        .unwrap_or(0)
+                })
+                .sum::<usize>()
+        })
+        .sum()
+}
+
+/// Grow `reservation` to the memory the scans need to replay `partitions`.
+/// If the pool cannot grant it, move in-memory batches to spill files, from
+/// the partition that holds the most memory first, until it can.
+fn reserve_replay_memory(
+    partitions: &mut [BufferedPartition],
+    scan_partitions: &[usize],
+    spill_manager: &SpillManager,
+    reservation: &mut MemoryReservation,
+) -> Result<()> {
+    loop {
+        let required = replay_memory(partitions, scan_partitions);
+        let Err(e) = reservation.try_resize(required) else {
+            return Ok(());
+        };
+        let Some(partition) = partitions
+            .iter_mut()
+            .filter(|p| !p.batches.is_empty())
+            .max_by_key(|p| p.reservation.size())
+        else {
+            return Err(e);
+        };
+        partition.spill_in_memory_suffix(
+            required.saturating_sub(reservation.size()),
+            spill_manager,
+        )?;
+    }
+}
+
+/// State shared by one [`MaterializedCteExec`] and all its
+/// [`MaterializedCteScanExec`]s.
+pub struct MaterializedCteBuffer {
+    id: u64,
+    name: String,
+    /// The body to run. [`MaterializedCteExec`] updates it every time it is
+    /// rebuilt (for example by a physical optimizer rule), so it is the body 
of
+    /// the final plan. A scan can run the body before the owning
+    /// `MaterializedCteExec` executes, for example from a scalar subquery.
+    body: Mutex<Option<Arc<dyn ExecutionPlan>>>,
+    output: Mutex<Arc<OnceAsync<MaterializedOutput>>>,
+    /// The partition count of every scan bound to this buffer.
+    scan_partitions: Mutex<Vec<usize>>,
+    metrics: ExecutionPlanMetricsSet,
+}
+
+impl fmt::Debug for MaterializedCteBuffer {
+    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
+        f.debug_struct("MaterializedCteBuffer")
+            .field("name", &self.name)
+            .finish_non_exhaustive()
+    }
+}
+
+impl MaterializedCteBuffer {
+    pub fn new(id: u64, name: impl Into<String>) -> Self {
+        Self {
+            id,
+            name: name.into(),
+            body: Mutex::new(None),
+            output: Mutex::new(Arc::default()),
+            scan_partitions: Mutex::new(vec![]),
+            metrics: ExecutionPlanMetricsSet::new(),
+        }
+    }
+
+    /// The id that binds scans to this buffer.
+    pub fn id(&self) -> u64 {
+        self.id
+    }
+
+    fn set_body(&self, body: Arc<dyn ExecutionPlan>) {
+        *self.body.lock() = Some(body);
+    }
+
+    fn reset(&self) {
+        *self.output.lock() = Arc::default();
+    }
+
+    /// Return a future that resolves once the body has been fully buffered.
+    /// The body runs at most once, whichever consumer asks first.
+    fn materialize(
+        self: &Arc<Self>,
+        context: Arc<TaskContext>,
+    ) -> Result<OnceFut<MaterializedOutput>> {
+        let Some(body) = self.body.lock().clone() else {
+            return internal_err!("MaterializedCte {} has no body", self.name);
+        };
+        let once = Arc::clone(&self.output.lock());
+        let name = self.name.clone();
+        let scan_partitions = self.scan_partitions.lock().clone();
+        let metrics = self.metrics.clone();
+        once.try_once(move || {
+            Ok(buffer_body(name, body, scan_partitions, context, metrics))
+        })
+    }
+}
+
+async fn buffer_body(
+    name: String,
+    body: Arc<dyn ExecutionPlan>,
+    scan_partitions: Vec<usize>,
+    context: Arc<TaskContext>,
+    metrics: ExecutionPlanMetricsSet,
+) -> Result<MaterializedOutput> {
+    let spill_manager = SpillManager::new(
+        context.runtime_env(),
+        SpillMetrics::new(&metrics, 0),
+        body.schema(),
+    )
+    .with_compression_type(context.session_config().spill_compression());
+    let buffered_rows = 
MetricBuilder::new(&metrics).global_counter("buffered_rows");
+
+    let mut join_set = JoinSet::new();
+    for partition in 0..body.output_partitioning().partition_count() {
+        let stream = body.execute(partition, Arc::clone(&context))?;
+        let reservation =
+            
MemoryConsumer::new(format!("MaterializedCte[{name}][{partition}]"))
+                .with_can_spill(true)
+                .register(context.memory_pool());
+        let spill_manager = spill_manager.clone();
+        let buffered_rows = buffered_rows.clone();
+        join_set.spawn(async move {
+            let result =
+                buffer_partition(stream, reservation, &spill_manager, 
&buffered_rows)
+                    .await;
+            (partition, result)
+        });
+    }
+
+    let mut partitions = Vec::with_capacity(join_set.len());
+    while let Some(joined) = join_set.join_next().await {
+        match joined {
+            Ok((partition, result)) => partitions.push((partition, result?)),
+            Err(e) if e.is_panic() => 
std::panic::resume_unwind(e.into_panic()),
+            Err(e) => return internal_err!("MaterializedCte task failed: {e}"),
+        }
+    }
+    partitions.sort_by_key(|(partition, _)| *partition);
+    let mut partitions: Vec<_> = partitions.into_iter().map(|(_, p)| 
p).collect();
+
+    let mut replay_reservation =
+        MemoryConsumer::new(format!("MaterializedCte[{name}] replay"))
+            .register(context.memory_pool());
+    reserve_replay_memory(
+        &mut partitions,
+        &scan_partitions,
+        &spill_manager,
+        &mut replay_reservation,
+    )?;
+    Ok(MaterializedOutput {
+        partitions,
+        spill_manager,
+        _replay_reservation: replay_reservation,
+    })
+}
+
+async fn buffer_partition(
+    mut stream: SendableRecordBatchStream,
+    reservation: MemoryReservation,
+    spill_manager: &SpillManager,
+    buffered_rows: &crate::metrics::Count,
+) -> Result<BufferedPartition> {
+    let mut batches = vec![];
+    let mut spill = None;
+    let mut max_spilled_batch_size = 0;
+    while let Some(batch) = stream.next().await.transpose()? {
+        buffered_rows.add(batch.num_rows());
+        // Once a partition spills, every later batch goes to the same file,
+        // so replay keeps the order of the partition.
+        if spill.is_none() && 
reservation.try_grow(batch.get_array_memory_size()).is_ok()
+        {
+            batches.push(batch);
+            continue;
+        }
+        if spill.is_none() {
+            spill = 
Some(spill_manager.create_in_progress_file("MaterializedCte")?);
+        }
+        if let Some(file) = spill.as_mut() {
+            let size = file.append_batch(&batch)?;
+            max_spilled_batch_size = max_spilled_batch_size.max(size);
+        }
+    }
+    let spill_file = match spill {
+        Some(mut file) => file.finish()?,
+        None => None,
+    };
+    Ok(BufferedPartition {
+        batches,
+        spill_files: spill_file.into_iter().collect(),
+        max_spilled_batch_size,
+        reservation,
+    })
+}
+
+/// Stream the buffered partitions `partition, partition + n, ...` of `output`,
+/// where `n` is the number of output partitions of the scan.
+fn replay(
+    output: Arc<MaterializedOutput>,
+    partition: usize,
+    output_partitions: usize,
+) -> Result<ReplayStream> {
+    let mut streams = vec![];
+    for buffered in output
+        .partitions
+        .iter()
+        .skip(partition)
+        .step_by(output_partitions)
+    {
+        let schema = Arc::clone(output.spill_manager.schema());
+        let batches = buffered.batches.clone();
+        streams.push(Box::pin(RecordBatchStreamAdapter::new(
+            schema,
+            futures::stream::iter(batches.into_iter().map(Ok)),
+        )) as SendableRecordBatchStream);
+        // Read without read-ahead, so that the decoded batch fits in the
+        // replay reservation.
+        for file in &buffered.spill_files {
+            streams.push(output.spill_manager.read_spill_as_stream_unbuffered(
+                Arc::clone(file),
+                Some(buffered.max_spilled_batch_size),
+            )?);
+        }
+    }
+    Ok(ReplayStream {
+        stream: futures::stream::iter(streams).flatten().boxed(),
+        _output: output,
+    })
+}
+
+/// The replay of one scan partition.
+///
+/// Holds the [`MaterializedOutput`], and with it the memory reservations of
+/// the buffered batches and of the spill reads, until the replay is dropped.
+/// The scan can outlive its plan, for example when the plan is dropped after
+/// `execute`, so the output must not be released earlier.
+struct ReplayStream {
+    stream: BoxStream<'static, Result<RecordBatch>>,
+    _output: Arc<MaterializedOutput>,
+}
+
+impl Stream for ReplayStream {
+    type Item = Result<RecordBatch>;
+
+    fn poll_next(
+        mut self: Pin<&mut Self>,
+        cx: &mut Context<'_>,
+    ) -> Poll<Option<Self::Item>> {
+        self.stream.poll_next_unpin(cx)
+    }
+}
+
+/// Computes a CTE body once and runs the continuation, whose
+/// [`MaterializedCteScanExec`]s read the buffered body output.
+///
+/// Children: `[body, continuation]`. The output is the output of the
+/// continuation.
+#[derive(Debug)]
+pub struct MaterializedCteExec {
+    body: Arc<dyn ExecutionPlan>,
+    continuation: Arc<dyn ExecutionPlan>,
+    buffer: Arc<MaterializedCteBuffer>,
+    cache: Arc<PlanProperties>,
+}
+
+impl MaterializedCteExec {
+    pub fn new(
+        body: Arc<dyn ExecutionPlan>,
+        continuation: Arc<dyn ExecutionPlan>,
+        buffer: Arc<MaterializedCteBuffer>,
+    ) -> Self {
+        buffer.set_body(Arc::clone(&body));
+        let cache = Arc::clone(continuation.properties());
+        Self {
+            body,
+            continuation,
+            buffer,
+            cache,
+        }
+    }
+
+    pub fn buffer(&self) -> &Arc<MaterializedCteBuffer> {
+        &self.buffer
+    }
+}
+
+impl DisplayAs for MaterializedCteExec {
+    fn fmt_as(&self, t: DisplayFormatType, f: &mut fmt::Formatter) -> 
fmt::Result {
+        match t {
+            DisplayFormatType::Default | DisplayFormatType::Verbose => {
+                write!(f, "MaterializedCteExec: name={}", self.buffer.name)
+            }
+            DisplayFormatType::TreeRender => write!(f, "name={}", 
self.buffer.name),
+        }
+    }
+}
+
+impl ExecutionPlan for MaterializedCteExec {
+    fn name(&self) -> &'static str {
+        "MaterializedCteExec"
+    }
+
+    fn properties(&self) -> &Arc<PlanProperties> {
+        &self.cache
+    }
+
+    fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
+        vec![&self.body, &self.continuation]
+    }
+
+    fn apply_expressions(
+        &self,
+        _f: &mut dyn FnMut(&Arc<dyn PhysicalExpr>) -> 
Result<TreeNodeRecursion>,
+    ) -> Result<TreeNodeRecursion> {
+        Ok(TreeNodeRecursion::Continue)
+    }
+
+    fn replace_children(
+        self: Arc<Self>,
+        mut children: Vec<Arc<dyn ExecutionPlan>>,
+        _: ReplaceChildrenOptions,
+    ) -> Result<Arc<dyn ExecutionPlan>> {
+        if children.len() != 2 {
+            return internal_err!("MaterializedCteExec takes 2 children");
+        }
+        let continuation = children.pop().unwrap();
+        let body = children.pop().unwrap();
+        Ok(Arc::new(Self::new(
+            body,
+            continuation,
+            Arc::clone(&self.buffer),
+        )))
+    }
+
+    fn with_new_children(
+        self: Arc<Self>,
+        children: Vec<Arc<dyn ExecutionPlan>>,
+    ) -> Result<Arc<dyn ExecutionPlan>> {
+        self.replace_children(
+            children,
+            ReplaceChildrenOptions::new(ChildrenPropertiesMode::Recompute),
+        )
+    }
+
+    fn reset_state(self: Arc<Self>) -> Result<Arc<dyn ExecutionPlan>> {

Review Comment:
   I think `reset_state` on `MaterializedCteExec` might cause problems. It gets 
called whenever a plan needs to run again from scratch, e.g. by 
`reset_plan_states`, which recursive queries use before each iteration. The 
expectation is that it returns a fresh copy and leaves the original alone.
   
   Here it clears the shared buffer in place, so every copy sees the reset, 
including one that's still running. 
   
   `HashJoinExec` returns a new node: 
https://github.com/apache/datafusion/blob/f427fb788fdb75c607db76063e1794f83e2ff83e/datafusion/physical-plan/src/joins/hash_join/exec.rs#L1681-L1683
 with an empty hash table and leaves the old one untouched. Could this do the 
same, i.e. create a new buffer and point its scans at it?



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to