fix: canonicalize predicates at backend sinks

This commit is contained in:
Gatefixer
2026-08-06 09:48:18 +00:00
parent 4fe087b78b
commit 8af2c26d54
11 changed files with 173 additions and 67 deletions
+1 -1
View File
@@ -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]
+2 -2
View File
@@ -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")
@@ -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 = [
+40 -10
View File
@@ -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<char> {
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<String> {
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]
+40 -11
View File
@@ -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::<Vec<_>>()
.await
.unwrap();
assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::<usize>(), 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);
}
+36 -24
View File
@@ -1082,11 +1082,12 @@ impl<S: HttpSend> RemoteTable<S> {
}
async fn prepare_query_bodies(&self, query: &AnyQuery) -> Result<Vec<serde_json::Value>> {
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<S: HttpSend> BaseTable for RemoteTable<S> {
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<S: HttpSend> BaseTable for RemoteTable<S> {
Ok(final_analyze)
}
async fn update(&self, update: UpdateBuilder) -> Result<UpdateResult> {
async fn update(&self, mut update: UpdateBuilder) -> Result<UpdateResult> {
update.canonicalize_filter()?;
self.check_mutable().await?;
let request = self
.client
@@ -2346,7 +2348,7 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
async fn delete(&self, predicate: Predicate<'_>) -> Result<DeleteResult> {
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<S: HttpSend> BaseTable for RemoteTable<S> {
async fn merge_insert(
&self,
params: MergeInsertBuilder,
mut params: MergeInsertBuilder,
new_data: Box<dyn RecordBatchReader + Send>,
) -> Result<MergeResult> {
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::<HashMap<_, _>>();
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::<Vec<_>>()
+4 -1
View File
@@ -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(),
}),
+2 -1
View File
@@ -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);
+11 -5
View File
@@ -224,12 +224,17 @@ impl MergeInsertBuilder {
mut self,
new_data: Box<dyn RecordBatchReader + Send>,
) -> Result<MergeResult> {
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<MergeFilter>) -> Result<Option<MergeFilter>> {
@@ -248,9 +253,10 @@ fn canonicalize_merge_filter(filter: Option<MergeFilter>) -> Result<Option<Merge
/// This logic was moved from NativeTable::merge_insert to keep table.rs clean.
pub(crate) async fn execute_merge_insert(
table: &NativeTable,
params: MergeInsertBuilder,
mut params: MergeInsertBuilder,
new_data: Box<dyn RecordBatchReader + Send>,
) -> Result<MergeResult> {
params.canonicalize_filters()?;
match lsm::lsm_dispatch_decision(table, &params).await? {
lsm::LsmDispatch::Lsm(plan) => {
let future =
+23 -5
View File
@@ -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<Self> {
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<DatasetRecordBatchStream> {
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<bool> {
@@ -133,9 +150,10 @@ pub async fn create_plan(
query: &AnyQuery,
options: QueryExecutionOptions,
) -> Result<Arc<dyn ExecutionPlan>> {
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()?;
+12 -5
View File
@@ -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<UpdateResult> {
update.canonicalize_filter()?;
table.dataset.ensure_mutable()?;
// 1. Snapshot the current dataset