From 7d1ef0c043344bc70b8c8d90d6be7223a1de894e Mon Sep 17 00:00:00 2001 From: Gatefixer <312823363+lance-gatefixer[bot]@users.noreply.github.com> Date: Fri, 7 Aug 2026 14:23:17 +0000 Subject: [PATCH] fix(datafusion): use caller session for filter planning --- python/python/tests/test_table_provider.py | 28 +++++++++ rust/lance-datafusion/src/planner.rs | 11 ++++ rust/lance/src/datafusion/dataframe.rs | 6 +- rust/lance/src/dataset/scanner.rs | 51 ++++++++++++--- rust/lance/src/io/exec/filter.rs | 20 +++++- rust/lance/src/io/exec/filtered_read.rs | 73 ++++++++++++++++++---- 6 files changed, 168 insertions(+), 21 deletions(-) diff --git a/python/python/tests/test_table_provider.py b/python/python/tests/test_table_provider.py index 1eddf220dd2..2252f12aa28 100644 --- a/python/python/tests/test_table_provider.py +++ b/python/python/tests/test_table_provider.py @@ -91,3 +91,31 @@ def make_ctx(): result = normalize(ctx.table("ffi_lance_table").limit(1, offset=1).collect()) assert len(result) == 1 assert result["col1"][0].as_py() == 1 + + +def test_custom_udf_filter(tmp_path): + pytest.importorskip("datafusion") + from datafusion import SessionContext, udf + + def is_even(values: pa.Array) -> pa.Array: + return pa.array([value.as_py() % 2 == 0 for value in values], type=pa.bool_()) + + is_even_udf = udf( + is_even, + input_fields=[pa.int64()], + return_field=pa.bool_(), + volatility="stable", + name="is_even", + ) + + dataset = lance.write_dataset(pa.table({"i": [1, 2, 3, 4]}), str(tmp_path)) + provider = FFILanceTableProvider(dataset, with_row_id=True, with_row_addr=True) + + ctx = SessionContext() + ctx.register_table("numbers", provider) + ctx.register_udf(is_even_udf) + + result = normalize( + ctx.sql("SELECT i FROM numbers WHERE i = 2 AND is_even(i)").collect() + ) + assert result["i"].to_pylist() == [2] diff --git a/rust/lance-datafusion/src/planner.rs b/rust/lance-datafusion/src/planner.rs index 5ee19ee2f49..7f4f2cbd376 100644 --- a/rust/lance-datafusion/src/planner.rs +++ b/rust/lance-datafusion/src/planner.rs @@ -17,6 +17,7 @@ use arrow_buffer::OffsetBuffer; use arrow_cast::cast_with_options; use arrow_schema::{DataType as ArrowDataType, Field, SchemaRef, TimeUnit}; use arrow_select::concat::concat; +use datafusion::catalog::Session; use datafusion::common::DFSchema; use datafusion::common::tree_node::{TreeNode, TreeNodeRecursion, TreeNodeVisitor}; use datafusion::config::ConfigOptions; @@ -1035,6 +1036,16 @@ impl Planner { )?) } + /// Create a [`PhysicalExpr`] using the caller's DataFusion session. + pub fn create_physical_expr_with_session( + &self, + expr: &Expr, + session: &dyn Session, + ) -> Result> { + let df_schema = DFSchema::try_from(self.schema.as_ref().clone())?; + Ok(session.create_physical_expr(expr.clone(), &df_schema)?) + } + /// Collect the columns in the expression. /// /// The columns are returned in sorted order. diff --git a/rust/lance/src/datafusion/dataframe.rs b/rust/lance/src/datafusion/dataframe.rs index 7ebc35edbaa..6ff6a09c837 100644 --- a/rust/lance/src/datafusion/dataframe.rs +++ b/rust/lance/src/datafusion/dataframe.rs @@ -112,7 +112,7 @@ impl TableProvider for LanceTableProvider { async fn scan( &self, - _state: &dyn Session, + state: &dyn Session, projection: Option<&Vec>, filters: &[Expr], limit: Option, @@ -161,7 +161,9 @@ impl TableProvider for LanceTableProvider { scan.limit(limit.map(|l| l as i64), None)?; scan.scan_in_order(self.ordered); - scan.create_plan().await.map_err(DataFusionError::from) + scan.create_plan_with_session(state) + .await + .map_err(DataFusionError::from) } // Since we are using datafusion itself to apply the filters it should diff --git a/rust/lance/src/dataset/scanner.rs b/rust/lance/src/dataset/scanner.rs index 335620a83f6..637295be8f1 100644 --- a/rust/lance/src/dataset/scanner.rs +++ b/rust/lance/src/dataset/scanner.rs @@ -17,6 +17,7 @@ use arrow_schema::{DataType, Field as ArrowField, Schema as ArrowSchema, SchemaR use arrow_select::concat::concat_batches; use async_recursion::async_recursion; use chrono::Utc; +use datafusion::catalog::Session; use datafusion::common::{DFSchema, JoinType, NullEquality, exec_datafusion_err}; use datafusion::functions_aggregate; use datafusion::logical_expr::{Expr, ScalarUDF, col, lit}; @@ -522,6 +523,7 @@ impl FilterPlan { &self, input: Arc, scanner: &Scanner, + session: Option<&dyn Session>, ) -> Result> { let mut plan = input; @@ -538,9 +540,12 @@ impl FilterPlan { } if let Some(refine_expr) = &self.expr_filter_plan.refine_expr { - // We create a new planner specific to the node's schema, since - // physical expressions reference column by index rather than by name. - plan = Arc::new(LanceFilterExec::try_new(refine_expr.clone(), plan)?); + plan = Arc::new(match session { + Some(session) => { + LanceFilterExec::try_new_with_session(refine_expr.clone(), plan, session)? + } + None => LanceFilterExec::try_new(refine_expr.clone(), plan)?, + }); } Ok(plan) @@ -2788,8 +2793,22 @@ impl Scanner { /// 3. Sort /// 4. Limit / Offset /// 5. Take remaining columns / Projection + pub fn create_plan(&self) -> BoxFuture<'_, Result>> { + Box::pin(self.create_plan_impl(None)) + } + + pub(crate) fn create_plan_with_session<'a>( + &'a self, + session: &'a dyn Session, + ) -> BoxFuture<'a, Result>> { + Box::pin(self.create_plan_impl(Some(session))) + } + #[instrument(level = "debug", skip_all)] - pub async fn create_plan(&self) -> Result> { + async fn create_plan_impl( + &self, + session: Option<&dyn Session>, + ) -> Result> { log::trace!("creating scanner plan"); self.validate_options()?; @@ -2852,7 +2871,7 @@ impl Scanner { self.take_source(take_op).await? } else { let planned_read = self - .filtered_read_source(&mut filter_plan.expr_filter_plan) + .filtered_read_source(&mut filter_plan.expr_filter_plan, session) .await?; if planned_read.limit_pushed_down { use_limit_node = false; @@ -2897,7 +2916,7 @@ impl Scanner { plan = self.take(plan, pre_filter_projection)?; // Filter - plan = filter_plan.refine_filter(plan, self).await?; + plan = filter_plan.refine_filter(plan, self, session).await?; // Aggregate (if set, applies aggregate and returns early) if let Some(agg) = &self.aggregate { @@ -3124,6 +3143,7 @@ impl Scanner { make_deletions_null: bool, fragments: Option>>, scan_range: Option>, + session: Option<&dyn Session>, ) -> Result> { // Kept for the overlay stale-Take path below, which re-evaluates blocked stale rows. let user_projection = projection.clone(); @@ -3168,6 +3188,10 @@ impl Scanner { read_options = read_options.with_only_indexed_fragments(); } + if let Some(session) = session { + read_options = read_options.with_physical_filters(session)?; + } + // Mask data overlay files: a row with an overlay committed after an index it relies on // touched an indexed field can no longer be trusted to that index. Block just those rows // from the index result (their fragments stay indexed, so non-stale rows keep the index) @@ -3217,7 +3241,12 @@ impl Scanner { .await?; let planner = Planner::new(stale_node.schema()); let optimized_filter = planner.optimize_expr(filter.clone())?; - let filtered = Arc::new(LanceFilterExec::try_new(optimized_filter, stale_node)?); + let filtered = Arc::new(match session { + Some(session) => { + LanceFilterExec::try_new_with_session(optimized_filter, stale_node, session)? + } + None => LanceFilterExec::try_new(optimized_filter, stale_node)?, + }); let stale_path: Arc = Arc::new(project(filtered, plan.schema().as_ref())?); @@ -3231,6 +3260,7 @@ impl Scanner { // Helper function for filtered read // // Delegates to legacy or new filtered read based on dataset storage version + #[allow(clippy::too_many_arguments)] async fn filtered_read( &self, filter_plan: &ExprFilterPlan, @@ -3239,6 +3269,7 @@ impl Scanner { fragments: Option>>, scan_range: Option>, is_prefilter: bool, + session: Option<&dyn Session>, ) -> Result { // Use legacy path if dataset uses legacy storage format if self.dataset.is_legacy_storage() { @@ -3260,6 +3291,7 @@ impl Scanner { make_deletions_null, fragments, scan_range, + session, ) .await?; Ok(PlannedFilteredScan { @@ -3322,6 +3354,7 @@ impl Scanner { async fn filtered_read_source( &self, filter_plan: &mut ExprFilterPlan, + session: Option<&dyn Session>, ) -> Result { log::trace!("source is a filtered read"); @@ -3375,6 +3408,7 @@ impl Scanner { self.fragments.clone().map(Arc::new), scan_range, /*is_prefilter= */ false, + session, ) .await } @@ -4606,6 +4640,7 @@ impl Scanner { Some(Arc::new(fragments)), None, /*is_prefilter=*/ true, + None, ) .await?; if let Some(refine_expr) = filter_plan.refine_expr.as_ref() { @@ -4896,6 +4931,7 @@ impl Scanner { self.fragments.clone().map(Arc::new), None, /*is_prefilter= */ true, + None, ) .await?; @@ -6191,6 +6227,7 @@ impl Scanner { Some(fragments), None, /*is_prefilter= */ true, + None, ) .await?; Ok(PreFilterSource::FilteredRowIds(plan)) diff --git a/rust/lance/src/io/exec/filter.rs b/rust/lance/src/io/exec/filter.rs index 71f1a5b2b4a..f6d3005065e 100644 --- a/rust/lance/src/io/exec/filter.rs +++ b/rust/lance/src/io/exec/filter.rs @@ -3,7 +3,7 @@ use std::sync::Arc; -use datafusion::{execution::TaskContext, logical_expr::Expr}; +use datafusion::{catalog::Session, execution::TaskContext, logical_expr::Expr}; use datafusion_physical_plan::{ DisplayAs, DisplayFormatType, ExecutionPlan, PlanProperties, SendableRecordBatchStream, Statistics, filter::FilterExec, metrics::MetricsSet, @@ -31,6 +31,24 @@ impl LanceFilterExec { pub fn try_new(expr: Expr, input: Arc) -> Result { let planner = Planner::new(input.schema()); let predicate = planner.create_physical_expr(&expr)?; + Self::try_new_with_predicate(expr, predicate, input) + } + + pub fn try_new_with_session( + expr: Expr, + input: Arc, + session: &dyn Session, + ) -> Result { + let planner = Planner::new(input.schema()); + let predicate = planner.create_physical_expr_with_session(&expr, session)?; + Self::try_new_with_predicate(expr, predicate, input) + } + + fn try_new_with_predicate( + expr: Expr, + predicate: Arc, + input: Arc, + ) -> Result { let filter_exec = FilterExec::try_new(predicate.clone(), input)?; Ok(Self { expr, diff --git a/rust/lance/src/io/exec/filtered_read.rs b/rust/lance/src/io/exec/filtered_read.rs index dddf18df4bd..8b1e016e261 100644 --- a/rust/lance/src/io/exec/filtered_read.rs +++ b/rust/lance/src/io/exec/filtered_read.rs @@ -13,6 +13,7 @@ use arrow_array::cast::AsArray; use arrow_array::types::UInt64Type; use arrow_array::{Array, BooleanArray, RecordBatch, UInt32Array}; use arrow_schema::{Schema as ArrowSchema, SchemaRef}; +use datafusion::catalog::Session; use datafusion::common::runtime::SpawnedTask; use datafusion::common::stats::Precision; use datafusion::error::{DataFusionError, Result as DataFusionResult}; @@ -125,6 +126,7 @@ struct ScopedFragmentRead { // An in-memory filter to apply after reading the fragment (whatever couldn't be // pushed down into the index query) filter: Option, + physical_filter: Option>, priority: u32, scan_scheduler: Arc, } @@ -861,6 +863,9 @@ impl FilteredReadStream { // Get filter for this fragment (convert Arc back to Expr) let filter = plan.filters.get(&fragment_id).map(|f| (**f).clone()); + let physical_filter = filter + .as_ref() + .and_then(|filter| options.physical_filter(filter)); scoped_fragments.push(ScopedFragmentRead { fragment: Arc::new(FileFragment::new(dataset.clone(), fragment.clone())), @@ -870,6 +875,7 @@ impl FilteredReadStream { batch_size: default_batch_size, file_reader_options: options.file_reader_options.clone(), filter, + physical_filter, priority: priority as u32, scan_scheduler: scan_scheduler.clone(), }); @@ -1315,15 +1321,18 @@ impl FilteredReadStream { // the row ids are not contiguous fragment_read_task.ranges.sort_by_key(|r| r.start); - let physical_filter = fragment_read_task - .filter - .map(|filter| { - let planner = Planner::new(public_blob_v2_binary_projection_schema( - fragment_read_task.projection.as_ref(), - )); - planner.create_physical_expr(&filter) - }) - .transpose()?; + let physical_filter = match fragment_read_task.physical_filter { + Some(filter) => Some(filter), + None => fragment_read_task + .filter + .map(|filter| { + let planner = Planner::new(public_blob_v2_binary_projection_schema( + fragment_read_task.projection.as_ref(), + )); + planner.create_physical_expr(&filter) + }) + .transpose()?, + }; // We are going to count the fragment as scanned on the first batch we // read. This might miss empty fragments, but we assume that wouldn't be @@ -1511,6 +1520,7 @@ pub struct FilteredReadOptions { /// result to avoid applying this (and instead only apply the refine filter) but in some cases /// the index result does not cover all fragments or is not exact. pub full_filter: Option, + physical_filters: Vec<(Expr, Arc)>, /// The threading mode to use for the scan pub threading_mode: FilteredReadThreadingMode, /// The size of the I/O buffer to use for the scan @@ -1550,6 +1560,7 @@ impl FilteredReadOptions { projection, refine_filter: None, full_filter: None, + physical_filters: Vec::new(), io_buffer_size_bytes: None, only_indexed_fragments: false, overlay_block: None, @@ -1683,6 +1694,7 @@ impl FilteredReadOptions { "refine_filter is set but full_filter is not".into(), )); } + self.physical_filters.clear(); self.refine_filter = refine_filter; self.full_filter = full_filter; Ok(self) @@ -1690,17 +1702,54 @@ impl FilteredReadOptions { /// An alternative to [`Self::with_filter`] to set the filters from a FilterPlan if you already have one pub fn with_filter_plan(mut self, filter_plan: FilterPlan) -> Self { + self.physical_filters.clear(); self.refine_filter = filter_plan.refine_expr; self.full_filter = filter_plan.full_expr; self } + /// Plan configured filters with the supplied DataFusion session. + pub(crate) fn with_physical_filters(mut self, session: &dyn Session) -> Result { + for filter in [&self.full_filter, &self.refine_filter] + .into_iter() + .flatten() + { + if self + .physical_filters + .iter() + .any(|(planned_filter, _)| planned_filter == filter) + { + continue; + } + + let filter_columns = Planner::column_names_in_expr(filter); + let projection = self + .projection + .clone() + .union_columns(filter_columns, OnMissing::Error)?; + let schema = public_blob_v2_binary_projection_schema(&projection); + let physical_filter = + Planner::new(schema).create_physical_expr_with_session(filter, session)?; + self.physical_filters + .push((filter.clone(), physical_filter)); + } + Ok(self) + } + + fn physical_filter(&self, filter: &Expr) -> Option> { + self.physical_filters + .iter() + .find(|(planned_filter, _)| planned_filter == filter) + .map(|(_, physical_filter)| physical_filter.clone()) + } + /// Specify the projection to use for the scan /// /// If the row id or row address are requested then they will be placed at the end /// of the output schema. If both are requested then the row id will come before /// the row address. pub fn with_projection(mut self, projection: Projection) -> Self { + self.physical_filters.clear(); self.projection = projection; self } @@ -2953,8 +3002,10 @@ impl ExecutionPlan for FilteredReadExec { let read_schema = public_blob_v2_binary_projection_schema(&read_projection); - let planner = Arc::new(Planner::new(read_schema.clone())); - let physical_filter = planner.create_physical_expr(filter)?; + let physical_filter = match self.options.physical_filter(filter) { + Some(physical_filter) => physical_filter, + None => Planner::new(read_schema.clone()).create_physical_expr(filter)?, + }; let mock_input = Arc::new(Self::try_new( self.dataset.clone(),