fix: canonicalize predicates across table operations

This commit is contained in:
Gatefixer
2026-08-06 08:44:30 +00:00
parent 00b1663732
commit 4fe087b78b
6 changed files with 189 additions and 114 deletions
+1 -1
View File
@@ -19,8 +19,8 @@
mod sql;
pub(crate) use sql::canonicalize_sql_predicate;
pub use sql::expr_to_sql_string;
pub(crate) use sql::normalize_sql_filter;
use std::sync::Arc;
+51 -98
View File
@@ -4,6 +4,10 @@
use datafusion_common::ScalarValue;
use datafusion_common::tree_node::{Transformed, TreeNode, TreeNodeRecursion};
use datafusion_expr::Expr;
use datafusion_sql::sqlparser::{
dialect::GenericDialect,
tokenizer::{Token, Tokenizer},
};
use datafusion_sql::unparser::{self, dialect::Dialect};
/// Unparser dialect that matches the quoting style expected by the Lance SQL
@@ -30,96 +34,34 @@ impl Dialect for LanceSqlDialect {
}
}
/// Translate SQL-standard double-quoted identifiers into the backtick-quoted
/// identifiers expected by Lance's SQL parser.
/// Canonicalize a raw SQL predicate for Lance's parser.
///
/// Lance historically interpreted double-quoted values as string literals.
/// Rewriting them at the query boundary avoids silently evaluating a predicate
/// such as `"mixedCase" = 'value'` as a comparison between two literals. String
/// contents and existing backtick-quoted identifiers are left unchanged.
pub fn normalize_sql_filter(filter: &str) -> crate::Result<String> {
#[derive(Clone, Copy, PartialEq, Eq)]
enum Quote {
None,
Single,
Backtick,
Double,
}
/// 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.
pub fn canonicalize_sql_predicate(predicate: &str) -> crate::Result<String> {
let dialect = GenericDialect;
let tokens = Tokenizer::new(&dialect, predicate)
.with_unescape(false)
.tokenize()
.map_err(|err| crate::Error::InvalidInput {
message: format!("invalid SQL predicate: {err}"),
})?;
let mut normalized = String::with_capacity(filter.len());
let mut chars = filter.chars().peekable();
let mut quote = Quote::None;
while let Some(ch) = chars.next() {
match quote {
Quote::None => match ch {
'\'' => {
normalized.push(ch);
quote = Quote::Single;
}
'`' => {
normalized.push(ch);
quote = Quote::Backtick;
}
'"' => {
normalized.push('`');
quote = Quote::Double;
}
_ => normalized.push(ch),
},
Quote::Single => {
normalized.push(ch);
if ch == '\\' {
if let Some(escaped) = chars.next() {
normalized.push(escaped);
}
} else if ch == '\'' {
if chars.peek() == Some(&'\'') {
normalized.push(chars.next().expect("peeked character must exist"));
} else {
quote = Quote::None;
}
}
Ok(tokens
.into_iter()
.map(|token| match token {
Token::Word(word) if word.quote_style == Some('"') => {
// with_unescape(false) retains doubled double quotes. Decode
// those before escaping any backticks for Lance's delimiter.
let identifier = word.value.replace("\"\"", "\"").replace('`', "``");
format!("`{identifier}`")
}
Quote::Backtick => {
normalized.push(ch);
if ch == '`' {
if chars.peek() == Some(&'`') {
normalized.push(chars.next().expect("peeked character must exist"));
} else {
quote = Quote::None;
}
}
}
Quote::Double => {
if ch == '"' {
if chars.peek() == Some(&'"') {
// SQL escapes a double quote within an identifier by
// doubling it. A quote needs no escaping inside Lance's
// backtick-delimited form.
normalized.push('"');
chars.next();
} else {
normalized.push('`');
quote = Quote::None;
}
} else if ch == '`' {
// Lance escapes a backtick within an identifier by doubling it.
normalized.push_str("``");
} else {
normalized.push(ch);
}
}
}
}
if quote == Quote::Double {
return Err(crate::Error::InvalidInput {
message: "unterminated double-quoted identifier in SQL filter".to_string(),
});
}
Ok(normalized)
other => other.to_string(),
})
.collect())
}
/// Prefix for placeholder strings inserted in place of binary literals. Chosen
@@ -208,34 +150,45 @@ pub fn expr_to_sql_string(expr: &Expr) -> crate::Result<String> {
#[cfg(test)]
mod tests {
use super::normalize_sql_filter;
use super::canonicalize_sql_predicate;
#[test]
fn normalizes_double_quoted_identifiers() {
assert_eq!(
normalize_sql_filter(r#""PartyAbbrev" = 'D'"#).unwrap(),
canonicalize_sql_predicate(r#""PartyAbbrev" = 'D'"#).unwrap(),
"`PartyAbbrev` = 'D'"
);
assert_eq!(
normalize_sql_filter(r#""MetaData"."userId" = 5"#).unwrap(),
canonicalize_sql_predicate(r#""MetaData"."userId" = 5"#).unwrap(),
"`MetaData`.`userId` = 5"
);
assert_eq!(normalize_sql_filter(r#""a""b" = 1"#).unwrap(), "`a\"b` = 1");
assert_eq!(
canonicalize_sql_predicate(r#""a""b" = 1"#).unwrap(),
"`a\"b` = 1"
);
}
#[test]
fn preserves_quotes_inside_literals_and_backticks() {
let filter = r#"name = 'Alice "Ace"' AND `quoted"field` = 1"#;
assert_eq!(normalize_sql_filter(filter).unwrap(), filter);
assert_eq!(canonicalize_sql_predicate(filter).unwrap(), filter);
}
#[test]
fn preserves_literals_and_comments_using_generic_dialect_rules() {
let predicate = r#"path = '\' AND "PartyAbbrev" = 'D' -- unmatched " in comment"#;
assert_eq!(
canonicalize_sql_predicate(predicate).unwrap(),
r#"path = '\' AND `PartyAbbrev` = 'D' -- unmatched " in comment"#
);
let predicate = r#"id = 1 /* unmatched " in block comment */"#;
assert_eq!(canonicalize_sql_predicate(predicate).unwrap(), predicate);
}
#[test]
fn rejects_unterminated_double_quoted_identifier() {
let error = normalize_sql_filter(r#""PartyAbbrev = 'D'"#).unwrap_err();
assert!(
error
.to_string()
.contains("unterminated double-quoted identifier")
);
let error = canonicalize_sql_predicate(r#""PartyAbbrev = 'D'"#).unwrap_err();
assert!(matches!(error, crate::Error::InvalidInput { .. }));
}
}
+102 -11
View File
@@ -917,7 +917,7 @@ impl QueryRequest {
/// by [`Self::check_filter`].
pub(crate) fn add_filter(&mut self, new: QueryFilter) {
let new = match new {
QueryFilter::Sql(filter) => match crate::expr::normalize_sql_filter(&filter) {
QueryFilter::Sql(filter) => match crate::expr::canonicalize_sql_predicate(&filter) {
Ok(filter) => QueryFilter::Sql(filter),
Err(err) => {
self.filter_error = Some(err.to_string());
@@ -1662,8 +1662,8 @@ mod tests {
use super::*;
use arrow::{array::downcast_array, compute::concat_batches, datatypes::Int32Type};
use arrow_array::{
FixedSizeListArray, Float32Array, Int32Array, RecordBatch, StringArray, cast::AsArray,
types::Float32Type,
FixedSizeListArray, Float32Array, Int32Array, RecordBatch, RecordBatchIterator,
StringArray, cast::AsArray, types::Float32Type,
};
use arrow_schema::{DataType, Field as ArrowField, Schema as ArrowSchema};
use futures::{StreamExt, TryStreamExt};
@@ -1902,18 +1902,22 @@ mod tests {
}
#[tokio::test]
async fn test_double_quoted_filter_identifier() {
async fn test_double_quoted_predicates_across_table_operations() {
let tmp_dir = tempdir().unwrap();
let dataset_path = tmp_dir.path().join("test.lance");
let uri = dataset_path.to_str().unwrap();
let schema = Arc::new(ArrowSchema::new(vec![ArrowField::new(
"PartyAbbrev",
DataType::Utf8,
false,
)]));
let schema = Arc::new(ArrowSchema::new(vec![
ArrowField::new("id", DataType::Int32, false),
ArrowField::new("PartyAbbrev", DataType::Utf8, false),
ArrowField::new("path", DataType::Utf8, false),
]));
let batch = RecordBatch::try_new(
schema,
vec![Arc::new(StringArray::from(vec!["D", "R", "R", "D"]))],
schema.clone(),
vec![
Arc::new(Int32Array::from(vec![1, 2, 3, 4])),
Arc::new(StringArray::from(vec!["D", "R", "R", "D"])),
Arc::new(StringArray::from(vec!["\\", "\\", "x", "x"])),
],
)
.unwrap();
@@ -1930,6 +1934,93 @@ mod tests {
.unwrap();
assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::<usize>(), 2);
assert_eq!(
table
.count_rows(Some(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#"path = '\' AND "PartyAbbrev" = 'D'"#,
] {
let batches = table
.query()
.only_if(predicate)
.execute()
.await
.unwrap()
.try_collect::<Vec<_>>()
.await
.unwrap();
assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::<usize>(), 1);
}
// The same canonical predicate contract applies to both merge filters.
let source = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Int32Array::from(vec![1, 2, 3])),
Arc::new(StringArray::from(vec!["D", "R", "R"])),
Arc::new(StringArray::from(vec!["\\", "\\", "x"])),
],
)
.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(),
)))
.await
.unwrap();
assert_eq!(result.num_deleted_rows, 1);
let source = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Int32Array::from(vec![1, 2, 3])),
Arc::new(StringArray::from(vec!["U", "U", "U"])),
Arc::new(StringArray::from(vec!["\\", "\\", "x"])),
],
)
.unwrap();
let mut merge = table.merge_insert(&["id"]);
merge.when_matched_update_all(Some(r#"target."PartyAbbrev" = 'D'"#.to_string()));
merge
.execute(Box::new(RecordBatchIterator::new(vec![Ok(source)], schema)))
.await
.unwrap();
assert_eq!(
table
.count_rows(Some(r#""PartyAbbrev" = 'U'"#.to_string()))
.await
.unwrap(),
1
);
table
.update()
.only_if(r#""PartyAbbrev" = 'R'"#)
.column("PartyAbbrev", "'X'")
.execute()
.await
.unwrap();
assert_eq!(
table
.count_rows(Some(r#""PartyAbbrev" = 'X'"#.to_string()))
.await
.unwrap(),
2
);
let result = table.delete(r#""PartyAbbrev" = 'X'"#).await.unwrap();
assert_eq!(result.num_deleted_rows, 2);
assert_eq!(table.count_rows(None).await.unwrap(), 1);
}
#[tokio::test]
+11 -2
View File
@@ -1074,7 +1074,10 @@ impl Table {
///
/// * `filter` if present, only count rows matching the filter
pub async fn count_rows(&self, filter: Option<String>) -> Result<usize> {
self.inner.count_rows(filter.map(Filter::Sql)).await
let filter = filter
.map(|predicate| crate::expr::canonicalize_sql_predicate(&predicate).map(Filter::Sql))
.transpose()?;
self.inner.count_rows(filter).await
}
/// Names of the blob v2 columns in this table, in declaration order.
@@ -1274,7 +1277,13 @@ impl Table {
/// # });
/// ```
pub async fn delete(&self, predicate: impl Into<Predicate<'_>>) -> Result<DeleteResult> {
self.inner.delete(predicate.into()).await
match predicate.into() {
Predicate::String(predicate) => {
let predicate = crate::expr::canonicalize_sql_predicate(predicate)?;
self.inner.delete(Predicate::String(&predicate)).await
}
predicate @ Predicate::Expr(_) => self.inner.delete(predicate).await,
}
}
/// Create an index on the provided column(s).
+19 -1
View File
@@ -220,11 +220,29 @@ impl MergeInsertBuilder {
///
/// Returns version and statistics about the merge operation including the number of rows
/// inserted, updated, and deleted.
pub async fn execute(self, new_data: Box<dyn RecordBatchReader + Send>) -> Result<MergeResult> {
pub async fn execute(
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.table.clone().merge_insert(self, new_data).await
}
}
fn canonicalize_merge_filter(filter: Option<MergeFilter>) -> Result<Option<MergeFilter>> {
filter
.map(|filter| match filter {
MergeFilter::Sql(predicate) => {
crate::expr::canonicalize_sql_predicate(&predicate).map(MergeFilter::Sql)
}
filter @ MergeFilter::Expr(_) => Ok(filter),
})
.transpose()
}
/// Internal implementation of the merge insert logic
///
/// This logic was moved from NativeTable::merge_insert to keep table.rs clean.
+5 -1
View File
@@ -62,12 +62,16 @@ impl UpdateBuilder {
}
/// Executes the update operation.
pub async fn execute(self) -> Result<UpdateResult> {
pub async fn execute(mut self) -> Result<UpdateResult> {
if self.columns.is_empty() {
Err(Error::InvalidInput {
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.parent.clone().update(self).await
}
}