diff --git a/vortex-duckdb/src/convert/expr.rs b/vortex-duckdb/src/convert/expr.rs index 51aeebebb6d..554b6a52ede 100644 --- a/vortex-duckdb/src/convert/expr.rs +++ b/vortex-duckdb/src/convert/expr.rs @@ -27,16 +27,15 @@ use vortex::error::VortexExpect; use vortex::error::VortexResult; use vortex::error::vortex_bail; use vortex::error::vortex_ensure; -use vortex::error::vortex_err; use vortex::expr::Expression; use vortex::expr::and_collect; use vortex::expr::byte_length; use vortex::expr::cast; use vortex::expr::col; use vortex::expr::get_item; +use vortex::expr::in_list; use vortex::expr::is_not_null; use vortex::expr::is_null; -use vortex::expr::list_contains; use vortex::expr::list_length; use vortex::expr::lit; use vortex::expr::not; @@ -433,13 +432,20 @@ pub fn can_push_expression(value: &duckdb::ExpressionRef) -> bool { // columns are native. } ExpressionClass::BoundOperator(op) => { + if matches!( + op.op, + DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_COMPARE_IN + | DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_COMPARE_NOT_IN + ) { + let mut children = op.children(); + return children.next().is_some_and(can_push_expression) + && children.all(|child| matches!(child.as_class(), Some(BoundConstant(_)))); + } if !matches!( op.op, DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_OPERATOR_NOT | DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_OPERATOR_IS_NULL | DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_OPERATOR_IS_NOT_NULL - | DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_COMPARE_IN - | DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_COMPARE_NOT_IN ) { return false; } @@ -729,12 +735,7 @@ fn try_from_compare_in( let Some(value) = try_from_expression_inner(c, ctx)? else { return Ok(None); }; - Ok(Some( - value - .as_opt::() - .ok_or_else(|| vortex_err!("cannot have a non literal in a in_list"))? - .clone(), - )) + Ok(value.as_opt::().cloned()) }) .collect::>>>()? else { @@ -743,10 +744,10 @@ fn try_from_compare_in( let list = Scalar::list( Arc::new(list_elements[0].dtype().clone()), list_elements, - Nullability::Nullable, + Nullability::NonNullable, ); - let expr = list_contains(lit(list), element); + let expr = in_list(element, lit(list)); Ok(Some(if not_in { not(expr) } else { expr })) } diff --git a/vortex-duckdb/src/convert/table_filter.rs b/vortex-duckdb/src/convert/table_filter.rs index b527e9576e2..8cfba75dd53 100644 --- a/vortex-duckdb/src/convert/table_filter.rs +++ b/vortex-duckdb/src/convert/table_filter.rs @@ -15,9 +15,9 @@ use vortex::error::vortex_err; use vortex::expr::Expression; use vortex::expr::and_collect; use vortex::expr::get_item; +use vortex::expr::in_list; use vortex::expr::is_not_null; use vortex::expr::is_null; -use vortex::expr::list_contains; use vortex::expr::lit; use vortex::expr::or_collect; use vortex::scalar::Scalar; @@ -89,8 +89,8 @@ pub fn try_from_table_filter( "IN filter must have at least one value" ); let dtype = scalars[0].dtype().clone(); - let list_scalar = Scalar::list(Arc::new(dtype), scalars, Nullability::Nullable); - list_contains(lit(list_scalar), col.clone()) + let list_scalar = Scalar::list(Arc::new(dtype), scalars, Nullability::NonNullable); + in_list(col.clone(), lit(list_scalar)) } TableFilterClass::Dynamic(dynamic) => { let op = match dynamic.operator { diff --git a/vortex-duckdb/src/e2e_test/vortex_scan_test.rs b/vortex-duckdb/src/e2e_test/vortex_scan_test.rs index e8c40346d24..a85c95c3dcf 100644 --- a/vortex-duckdb/src/e2e_test/vortex_scan_test.rs +++ b/vortex-duckdb/src/e2e_test/vortex_scan_test.rs @@ -18,6 +18,7 @@ use jiff::Zoned; use jiff::tz; use jiff::tz::TimeZone; use num_traits::AsPrimitive; +use rstest::rstest; use tempfile::NamedTempFile; use vortex::array::IntoArray; use vortex::array::VortexSessionExecute; @@ -280,6 +281,101 @@ fn test_vortex_scan_integers_between() { assert_eq!(sum, 43); } +fn assert_membership_filter_pushed( + conn: &Connection, + table: &str, + predicate: &str, + expected: i64, +) -> Result<()> { + let query = format!("SELECT count(*) FROM {table} WHERE {predicate}"); + let result = conn.query(&query)?; + let chunk = result.into_iter().next().unwrap(); + assert_eq!(chunk.get_vector(0).as_slice_with_len::(1), [expected]); + + let mut plan = String::new(); + for mut chunk in conn.query(&format!("EXPLAIN {query}"))? { + let len = chunk.len().as_(); + let vector = chunk.get_vector_mut(1); + for value in unsafe { vector.as_slice_mut::(len) } { + plan.push_str(&String::from_duckdb_value(value)); + } + } + assert!( + plan.contains("list_contains"), + "missing membership filter:\n{plan}" + ); + assert!(!plan.contains("FILTER"), "filter was not pushed:\n{plan}"); + Ok(()) +} + +#[rstest] +#[case::in_list("number IN (1, 3)", 2)] +#[case::not_in_list("number NOT IN (1, 3)", 2)] +#[case::null_last("number IN (1, NULL)", 1)] +#[case::null_first("number IN (NULL, 1)", 1)] +#[case::not_in_null("number NOT IN (1, NULL)", 0)] +#[case::unknown("(number IN (1, NULL)) IS NULL", 4)] +#[case::known("(number IN (NULL, 1)) IS NOT NULL", 1)] +fn test_membership_filter_pushdown( + #[case] predicate: &str, + #[case] expected: i64, + #[values(false, true)] through_view: bool, +) -> Result<()> { + let file = RUNTIME.block_on(async { + let numbers = + PrimitiveArray::from_option_iter([Some(1i32), Some(2), Some(3), Some(4), None]); + write_single_column_vortex_file("number", numbers).await + }); + let conn = database_connection(); + let file_table = format!("'{}'", file.path().to_string_lossy()); + let table = if through_view { + conn.query(&format!( + "CREATE VIEW members AS SELECT * FROM {file_table}" + ))?; + "members" + } else { + file_table.as_str() + }; + assert_membership_filter_pushed(&conn, table, predicate, expected) +} + +#[test] +fn test_large_membership_filter_pushdown() -> Result<()> { + let file = RUNTIME.block_on(async { + let numbers = PrimitiveArray::from_option_iter([Some(1i32), Some(42), Some(1001), None]); + write_single_column_vortex_file("number", numbers).await + }); + let conn = database_connection(); + let members = (0..1000) + .map(|value| value.to_string()) + .collect::>() + .join(","); + assert_membership_filter_pushed( + &conn, + &format!("'{}'", file.path().to_string_lossy()), + &format!("number IN ({members})"), + 2, + ) +} + +#[test] +fn test_membership_with_nonconstant_list_stays_in_duckdb() -> Result<()> { + let file = RUNTIME.block_on(async { + let numbers = buffer![1i32, 2, 3]; + write_single_column_vortex_file("number", numbers).await + }); + let conn = database_connection(); + // Exercise the unsupported IN expression itself, before DuckDB expands it to comparisons. + conn.query("SET disabled_optimizers = 'in_clause'")?; + let result = conn.query(&format!( + "SELECT count(*) FROM '{}' WHERE number IN (number, 99)", + file.path().to_string_lossy(), + ))?; + let chunk = result.into_iter().next().unwrap(); + assert_eq!(chunk.get_vector(0).as_slice_with_len::(1), [3]); + Ok(()) +} + #[test] fn test_issue_5927_not_in_does_not_panic() { let file = RUNTIME.block_on(async {