From 8af2c26d54df5fd052719d02037d990478a06b0f Mon Sep 17 00:00:00 2001 From: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Thu, 6 Aug 2026 09:48:18 +0000 Subject: [PATCH] fix: canonicalize predicates at backend sinks --- nodejs/examples/basic.test.ts | 2 +- python/python/tests/docs/test_basic.py | 4 +- python/python/tests/docs/test_guide_tables.py | 4 +- rust/lancedb/src/expr/sql.rs | 50 ++++++++++++---- rust/lancedb/src/query.rs | 51 ++++++++++++---- rust/lancedb/src/remote/table.rs | 60 +++++++++++-------- rust/lancedb/src/table.rs | 5 +- rust/lancedb/src/table/delete.rs | 3 +- rust/lancedb/src/table/merge.rs | 16 +++-- rust/lancedb/src/table/query.rs | 28 +++++++-- rust/lancedb/src/table/update.rs | 17 ++++-- 11 files changed, 173 insertions(+), 67 deletions(-) diff --git a/nodejs/examples/basic.test.ts b/nodejs/examples/basic.test.ts index b56bb2f95..e45fd622a 100644 --- a/nodejs/examples/basic.test.ts +++ b/nodejs/examples/basic.test.ts @@ -170,7 +170,7 @@ test("basic table examples", async () => { // --8<-- [end:create_index] // --8<-- [start:delete_rows] - await tbl.delete('item = "fizz"'); + await tbl.delete("item = 'fizz'"); // --8<-- [end:delete_rows] // --8<-- [start:drop_table] diff --git a/python/python/tests/docs/test_basic.py b/python/python/tests/docs/test_basic.py index 2a824371f..35d7aac10 100644 --- a/python/python/tests/docs/test_basic.py +++ b/python/python/tests/docs/test_basic.py @@ -105,7 +105,7 @@ def test_quickstart(tmp_path): tbl.create_index(num_sub_vectors=1) # --8<-- [end:create_index] # --8<-- [start:delete_rows] - tbl.delete('item = "fizz"') + tbl.delete("item = 'fizz'") # --8<-- [end:delete_rows] # --8<-- [start:drop_table] db.drop_table("my_table") @@ -201,7 +201,7 @@ async def test_quickstart_async(tmp_path): await tbl.create_index("vector") # --8<-- [end:create_index_async] # --8<-- [start:delete_rows_async] - await tbl.delete('item = "fizz"') + await tbl.delete("item = 'fizz'") # --8<-- [end:delete_rows_async] # --8<-- [start:drop_table_async] await db.drop_table("my_table_async") diff --git a/python/python/tests/docs/test_guide_tables.py b/python/python/tests/docs/test_guide_tables.py index 9ae86d167..dab8d43e9 100644 --- a/python/python/tests/docs/test_guide_tables.py +++ b/python/python/tests/docs/test_guide_tables.py @@ -266,7 +266,7 @@ def test_table(): tbl.add(pydantic_model_items) # --8<-- [end:add_table_from_pydantic] # --8<-- [start:delete_row] - tbl.delete('item = "fizz"') + tbl.delete("item = 'fizz'") # --8<-- [end:delete_row] # --8<-- [start:delete_specific_row] data = [ @@ -538,7 +538,7 @@ async def test_table_async(): await async_tbl.add(pydantic_model_items) # --8<-- [end:add_table_async_from_pydantic] # --8<-- [start:delete_row_async] - await async_tbl.delete('item = "fizz"') + await async_tbl.delete("item = 'fizz'") # --8<-- [end:delete_row_async] # --8<-- [start:delete_specific_row_async] data = [ diff --git a/rust/lancedb/src/expr/sql.rs b/rust/lancedb/src/expr/sql.rs index d5c17819d..24a676485 100644 --- a/rust/lancedb/src/expr/sql.rs +++ b/rust/lancedb/src/expr/sql.rs @@ -1,14 +1,16 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright The LanceDB Authors +use std::any::TypeId; + use datafusion_common::ScalarValue; use datafusion_common::tree_node::{Transformed, TreeNode, TreeNodeRecursion}; use datafusion_expr::Expr; use datafusion_sql::sqlparser::{ - dialect::GenericDialect, + dialect::{Dialect as SqlParserDialect, GenericDialect}, tokenizer::{Token, Tokenizer}, }; -use datafusion_sql::unparser::{self, dialect::Dialect}; +use datafusion_sql::unparser::{self, dialect::Dialect as UnparserDialect}; /// Unparser dialect that matches the quoting style expected by the Lance SQL /// parser. Lance uses backtick (`` ` ``) as the only delimited-identifier @@ -23,7 +25,7 @@ use datafusion_sql::unparser::{self, dialect::Dialect}; /// lower-case by the SQL parser, which would break case-sensitive schemas). struct LanceSqlDialect; -impl Dialect for LanceSqlDialect { +impl UnparserDialect for LanceSqlDialect { fn identifier_quote_style(&self, identifier: &str) -> Option { let needs_quote = identifier.chars().any(|c| c.is_ascii_uppercase()) || !identifier @@ -34,15 +36,40 @@ impl Dialect for LanceSqlDialect { } } +/// Lance's tokenizer dialect with SQL-standard double-quoted identifiers added. +/// +/// Keep this deliberately small: Lance's parser wraps `GenericDialect` and +/// delegates only identifier recognition, leaving every other dialect option at +/// its default. In particular, `/*! ... */` remains an ordinary block comment. +#[derive(Debug, Default)] +struct PredicateDialect(GenericDialect); + +impl SqlParserDialect for PredicateDialect { + fn dialect(&self) -> TypeId { + self.0.dialect() + } + + fn is_identifier_start(&self, ch: char) -> bool { + self.0.is_identifier_start(ch) + } + + fn is_identifier_part(&self, ch: char) -> bool { + self.0.is_identifier_part(ch) + } + + fn is_delimited_identifier_start(&self, ch: char) -> bool { + ch == '"' || ch == '`' + } +} + /// Canonicalize a raw SQL predicate for Lance's parser. /// -/// Lance delegates SQL lexing to [`GenericDialect`] except that it accepts only -/// backticks for delimited identifiers and historically interprets double-quoted -/// tokens as string literals. Tokenizing with the generic dialect lets us rewrite -/// only SQL-standard double-quoted identifier tokens while preserving string -/// literals, comments, and every other token according to the same lexical rules. +/// Lance wraps [`GenericDialect`] for identifier recognition while retaining the +/// default dialect behavior for every other lexical option. [`PredicateDialect`] +/// mirrors that contract and additionally recognizes `"` as an identifier +/// delimiter, allowing this function to rewrite only those identifier tokens. pub fn canonicalize_sql_predicate(predicate: &str) -> crate::Result { - let dialect = GenericDialect; + let dialect = PredicateDialect::default(); let tokens = Tokenizer::new(&dialect, predicate) .with_unescape(false) .tokenize() @@ -175,7 +202,7 @@ mod tests { } #[test] - fn preserves_literals_and_comments_using_generic_dialect_rules() { + fn preserves_literals_and_comments_using_lance_dialect_rules() { let predicate = r#"path = '\' AND "PartyAbbrev" = 'D' -- unmatched " in comment"#; assert_eq!( canonicalize_sql_predicate(predicate).unwrap(), @@ -184,6 +211,9 @@ mod tests { let predicate = r#"id = 1 /* unmatched " in block comment */"#; assert_eq!(canonicalize_sql_predicate(predicate).unwrap(), predicate); + + let predicate = r#"id = 1 /*! OR "PartyAbbrev" = 'D' */"#; + assert_eq!(canonicalize_sql_predicate(predicate).unwrap(), predicate); } #[test] diff --git a/rust/lancedb/src/query.rs b/rust/lancedb/src/query.rs index a4456008f..64f4e6a6a 100644 --- a/rust/lancedb/src/query.rs +++ b/rust/lancedb/src/query.rs @@ -1942,9 +1942,35 @@ mod tests { 2 ); + // Public BaseTable dispatch cannot bypass canonicalization. + let query = AnyQuery::Query(QueryRequest { + filter: Some(QueryFilter::Sql(r#""PartyAbbrev" = 'D'"#.to_string())), + ..Default::default() + }); + let batches = table + .base_table() + .query(&query, Default::default()) + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::(), 2); + assert_eq!( + table + .base_table() + .count_rows(Some(crate::table::Filter::Sql( + r#""PartyAbbrev" = 'D'"#.to_string(), + ))) + .await + .unwrap(), + 2 + ); + for predicate in [ r#"id = 1 -- unmatched " in a valid SQL comment"#, r#"id = 1 /* unmatched " in a valid SQL comment */"#, + r#"id = 1 /*! OR "PartyAbbrev" = 'D' */"#, r#"path = '\' AND "PartyAbbrev" = 'D'"#, ] { let batches = table @@ -1971,11 +1997,12 @@ mod tests { .unwrap(); let mut merge = table.merge_insert(&["id"]); merge.when_not_matched_by_source_delete(Some(r#""PartyAbbrev" = 'D'"#.to_string())); - let result = merge - .execute(Box::new(RecordBatchIterator::new( - vec![Ok(source)], - schema.clone(), - ))) + let result = table + .base_table() + .merge_insert( + merge, + Box::new(RecordBatchIterator::new(vec![Ok(source)], schema.clone())), + ) .await .unwrap(); assert_eq!(result.num_deleted_rows, 1); @@ -2003,13 +2030,11 @@ mod tests { 1 ); - table + let update = table .update() .only_if(r#""PartyAbbrev" = 'R'"#) - .column("PartyAbbrev", "'X'") - .execute() - .await - .unwrap(); + .column("PartyAbbrev", "'X'"); + table.base_table().update(update).await.unwrap(); assert_eq!( table .count_rows(Some(r#""PartyAbbrev" = 'X'"#.to_string())) @@ -2018,7 +2043,11 @@ mod tests { 2 ); - let result = table.delete(r#""PartyAbbrev" = 'X'"#).await.unwrap(); + let result = table + .base_table() + .delete(crate::table::Predicate::String(r#""PartyAbbrev" = 'X'"#)) + .await + .unwrap(); assert_eq!(result.num_deleted_rows, 2); assert_eq!(table.count_rows(None).await.unwrap(), 1); } diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index 5fefabeb2..5628d5acf 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -1082,11 +1082,12 @@ impl RemoteTable { } async fn prepare_query_bodies(&self, query: &AnyQuery) -> Result> { + let query = query.canonicalized()?; let version = self.current_version().await; let mut base_body = serde_json::json!({ "version": version }); self.apply_branch_body(&mut base_body); - match query { + match &query { AnyQuery::Query(query) => { let mut body = base_body.clone(); self.apply_query_params(&mut body, query)?; @@ -2068,7 +2069,7 @@ impl BaseTable for RemoteTable { let mut body = if let Some(filter) = filter { let filter_sql = match filter { - Filter::Sql(sql) => sql.clone(), + Filter::Sql(sql) => crate::expr::canonicalize_sql_predicate(&sql)?, Filter::Datafusion(expr) => expr_to_sql_string(&expr)?, }; serde_json::json!({ "predicate": filter_sql, "version": version }) @@ -2302,7 +2303,8 @@ impl BaseTable for RemoteTable { Ok(final_analyze) } - async fn update(&self, update: UpdateBuilder) -> Result { + async fn update(&self, mut update: UpdateBuilder) -> Result { + update.canonicalize_filter()?; self.check_mutable().await?; let request = self .client @@ -2346,7 +2348,7 @@ impl BaseTable for RemoteTable { async fn delete(&self, predicate: Predicate<'_>) -> Result { self.check_mutable().await?; let predicate_sql = match predicate { - Predicate::String(s) => s.to_string(), + Predicate::String(s) => crate::expr::canonicalize_sql_predicate(s)?, Predicate::Expr(expr) => expr_to_sql_string(expr)?, }; let mut body = serde_json::json!({ "predicate": predicate_sql }); @@ -2394,9 +2396,10 @@ impl BaseTable for RemoteTable { async fn merge_insert( &self, - params: MergeInsertBuilder, + mut params: MergeInsertBuilder, new_data: Box, ) -> Result { + params.canonicalize_filters()?; self.check_mutable().await?; let timeout = params.timeout; @@ -3190,13 +3193,17 @@ mod tests { ); assert_eq!( request.body().unwrap().as_bytes().unwrap(), - br#"{"predicate":"a > 10","version":null}"# + br#"{"predicate":"`A` > 10","version":null}"# ); http::Response::builder().status(200).body("42").unwrap() }); - let count = table.count_rows(Some("a > 10".into())).await.unwrap(); + let count = table + .base_table() + .count_rows(Some(Filter::Sql(r#""A" > 10"#.into()))) + .await + .unwrap(); assert_eq!(count, 42); } @@ -3588,7 +3595,7 @@ mod tests { assert_eq!(expression, "b - 1"); let only_if = value.get("predicate").unwrap().as_str().unwrap(); - assert_eq!(only_if, "b > 10"); + assert_eq!(only_if, "`B` > 10"); } if old_server { @@ -3604,14 +3611,12 @@ mod tests { } }); - let result = table + let update = table .update() .column("a", "a + 1") .column("b", "b - 1") - .only_if("b > 10") - .execute() - .await - .unwrap(); + .only_if(r#""B" > 10"#); + let result = table.base_table().update(update).await.unwrap(); assert_eq!(result.version, if old_server { 0 } else { 43 }); assert_eq!(result.rows_updated, if old_server { 0 } else { 5 }); @@ -3695,10 +3700,10 @@ mod tests { let params = request.url().query_pairs().collect::>(); assert_eq!(params["on"], "some_col"); - assert_eq!(params["when_matched_update_all"], "false"); + assert_eq!(params["when_matched_update_all"], "true"); assert_eq!(params["when_not_matched_insert_all"], "false"); assert_eq!(params["when_not_matched_by_source_delete"], "false"); - assert!(!params.contains_key("when_matched_update_all_filt")); + assert_eq!(params["when_matched_update_all_filt"], "target.`A` > 0"); assert!(!params.contains_key("when_not_matched_by_source_delete_filt")); assert!(!params.contains_key("use_index")); @@ -3715,11 +3720,9 @@ mod tests { } }); - let result = table - .merge_insert(&["some_col"]) - .execute(data) - .await - .unwrap(); + let mut merge = table.merge_insert(&["some_col"]); + merge.when_matched_update_all(Some(r#"target."A" > 0"#.into())); + let result = table.base_table().merge_insert(merge, data).await.unwrap(); assert_eq!(result.version, if old_server { 0 } else { 43 }); if !old_server { @@ -3781,7 +3784,7 @@ mod tests { let body = request.body().unwrap().as_bytes().unwrap(); let body: serde_json::Value = serde_json::from_slice(body).unwrap(); let predicate = body.get("predicate").unwrap().as_str().unwrap(); - assert_eq!(predicate, "id in (1, 2, 3)"); + assert_eq!(predicate, "`ID` in (1, 2, 3)"); if old_server { http::Response::builder().status(200).body("{}").unwrap() @@ -3796,7 +3799,11 @@ mod tests { } }); - let result = table.delete("id in (1, 2, 3)").await.unwrap(); + let result = table + .base_table() + .delete(Predicate::String(r#""ID" in (1, 2, 3)"#)) + .await + .unwrap(); assert_eq!(result.version, if old_server { 0 } else { 43 }); } @@ -3888,6 +3895,7 @@ mod tests { let body = request.body().unwrap().as_bytes().unwrap(); let body: serde_json::Value = serde_json::from_slice(body).unwrap(); let expected_body = serde_json::json!({ + "filter": "`A` > 0", "k": isize::MAX as usize, "prefilter": true, "vector": [], // Empty vector means no vector query. @@ -3903,9 +3911,13 @@ mod tests { .unwrap() }); + let query = AnyQuery::Query(QueryRequest { + filter: Some(QueryFilter::Sql(r#""A" > 0"#.into())), + ..Default::default() + }); let data = table - .query() - .execute() + .base_table() + .query(&query, Default::default()) .await .unwrap() .collect::>() diff --git a/rust/lancedb/src/table.rs b/rust/lancedb/src/table.rs index 15e7e6890..e4d895a46 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -2966,7 +2966,10 @@ impl BaseTable for NativeTable { let dataset = self.dataset.get().await?; match filter { None => Ok(dataset.count_rows(None).await?), - Some(Filter::Sql(sql)) => Ok(dataset.count_rows(Some(sql)).await?), + Some(Filter::Sql(sql)) => { + let sql = crate::expr::canonicalize_sql_predicate(&sql)?; + Ok(dataset.count_rows(Some(sql)).await?) + } Some(Filter::Datafusion(_)) => Err(Error::NotSupported { message: "Datafusion filters are not yet supported".to_string(), }), diff --git a/rust/lancedb/src/table/delete.rs b/rust/lancedb/src/table/delete.rs index 8f11ee019..cb1da03ae 100644 --- a/rust/lancedb/src/table/delete.rs +++ b/rust/lancedb/src/table/delete.rs @@ -31,8 +31,9 @@ pub(crate) async fn execute_delete( table.dataset.ensure_mutable()?; match predicate { Predicate::String(s) => { + let predicate = crate::expr::canonicalize_sql_predicate(s)?; let mut dataset = (*table.dataset.get().await?).clone(); - let delete_result = dataset.delete(s).boxed().await?; + let delete_result = dataset.delete(&predicate).boxed().await?; let num_deleted_rows = delete_result.num_deleted_rows; let version = dataset.version().version; table.dataset.update(dataset); diff --git a/rust/lancedb/src/table/merge.rs b/rust/lancedb/src/table/merge.rs index 6e703ee35..df378c999 100644 --- a/rust/lancedb/src/table/merge.rs +++ b/rust/lancedb/src/table/merge.rs @@ -224,12 +224,17 @@ impl MergeInsertBuilder { mut self, new_data: Box, ) -> Result { - self.when_matched_update_all_filt = - canonicalize_merge_filter(self.when_matched_update_all_filt)?; - self.when_not_matched_by_source_delete_filt = - canonicalize_merge_filter(self.when_not_matched_by_source_delete_filt)?; + self.canonicalize_filters()?; self.table.clone().merge_insert(self, new_data).await } + + pub(crate) fn canonicalize_filters(&mut self) -> Result<()> { + self.when_matched_update_all_filt = + canonicalize_merge_filter(self.when_matched_update_all_filt.take())?; + self.when_not_matched_by_source_delete_filt = + canonicalize_merge_filter(self.when_not_matched_by_source_delete_filt.take())?; + Ok(()) + } } fn canonicalize_merge_filter(filter: Option) -> Result> { @@ -248,9 +253,10 @@ fn canonicalize_merge_filter(filter: Option) -> Result, ) -> Result { + params.canonicalize_filters()?; match lsm::lsm_dispatch_decision(table, ¶ms).await? { lsm::LsmDispatch::Lsm(plan) => { let future = diff --git a/rust/lancedb/src/table/query.rs b/rust/lancedb/src/table/query.rs index 9feb9d5ab..4267b8811 100644 --- a/rust/lancedb/src/table/query.rs +++ b/rust/lancedb/src/table/query.rs @@ -45,6 +45,22 @@ impl AnyQuery { Self::VectorQuery(query) => &query.base, } } + + fn base_mut(&mut self) -> &mut QueryRequest { + match self { + Self::Query(query) => query, + Self::VectorQuery(query) => &mut query.base, + } + } + + /// Canonicalize any raw SQL filter immediately before backend dispatch. + pub(crate) fn canonicalized(&self) -> Result { + let mut query = self.clone(); + if let Some(QueryFilter::Sql(predicate)) = &mut query.base_mut().filter { + *predicate = crate::expr::canonicalize_sql_predicate(predicate)?; + } + Ok(query) + } } //Decide between namespace or local @@ -53,15 +69,16 @@ pub async fn execute_query( query: &AnyQuery, options: QueryExecutionOptions, ) -> Result { + let query = query.canonicalized()?; // QueryTable pushdown runs the query server-side, but only on the main // branch: the namespace request carries no branch yet, so a branch handle // must fall through to local execution. - if can_execute_namespace_query(table, query).await? + if can_execute_namespace_query(table, &query).await? && let Some(ref namespace_client) = table.namespace_client { - return execute_namespace_query(table, namespace_client.clone(), query, options).await; + return execute_namespace_query(table, namespace_client.clone(), &query, options).await; } - execute_generic_query(table, query, options).await + execute_generic_query(table, &query, options).await } async fn can_execute_namespace_query(table: &NativeTable, query: &AnyQuery) -> Result { @@ -133,9 +150,10 @@ pub async fn create_plan( query: &AnyQuery, options: QueryExecutionOptions, ) -> Result> { + let query = query.canonicalized()?; let query = match query { - AnyQuery::VectorQuery(query) => query.clone(), - AnyQuery::Query(query) => VectorQueryRequest::from_plain_query(query.clone()), + AnyQuery::VectorQuery(query) => query, + AnyQuery::Query(query) => VectorQueryRequest::from_plain_query(query), }; query.base.check_filter()?; diff --git a/rust/lancedb/src/table/update.rs b/rust/lancedb/src/table/update.rs index 99e3cde89..9da040a1a 100644 --- a/rust/lancedb/src/table/update.rs +++ b/rust/lancedb/src/table/update.rs @@ -68,20 +68,27 @@ impl UpdateBuilder { message: "at least one column must be specified in an update operation".to_string(), }) } else { - self.filter = self - .filter - .map(|predicate| crate::expr::canonicalize_sql_predicate(&predicate)) - .transpose()?; + self.canonicalize_filter()?; self.parent.clone().update(self).await } } + + pub(crate) fn canonicalize_filter(&mut self) -> Result<()> { + self.filter = self + .filter + .take() + .map(|predicate| crate::expr::canonicalize_sql_predicate(&predicate)) + .transpose()?; + Ok(()) + } } /// Internal implementation of the update logic pub(crate) async fn execute_update( table: &NativeTable, - update: UpdateBuilder, + mut update: UpdateBuilder, ) -> Result { + update.canonicalize_filter()?; table.dataset.ensure_mutable()?; // 1. Snapshot the current dataset