diff --git a/.bumpversion.toml b/.bumpversion.toml index 1c4aea809..2e0b78bf3 100644 --- a/.bumpversion.toml +++ b/.bumpversion.toml @@ -1,5 +1,5 @@ [tool.bumpversion] -current_version = "0.38.0-beta.10" +current_version = "0.38.0-beta.11" parse = """(?x) (?P0|[1-9]\\d*)\\. (?P0|[1-9]\\d*)\\. diff --git a/Cargo.lock b/Cargo.lock index e27f9f271..822689a0c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5402,7 +5402,7 @@ dependencies = [ [[package]] name = "lancedb" -version = "0.38.0-beta.10" +version = "0.38.0-beta.11" dependencies = [ "ahash", "anyhow", @@ -5490,7 +5490,7 @@ dependencies = [ [[package]] name = "lancedb-nodejs" -version = "0.38.0-beta.10" +version = "0.38.0-beta.11" dependencies = [ "arrow-array", "arrow-buffer", @@ -5515,7 +5515,7 @@ dependencies = [ [[package]] name = "lancedb-python" -version = "0.38.0-beta.10" +version = "0.38.0-beta.11" dependencies = [ "arrow", "async-trait", diff --git a/docs/src/java/java.md b/docs/src/java/java.md index 06dc267f3..1ce012522 100644 --- a/docs/src/java/java.md +++ b/docs/src/java/java.md @@ -14,7 +14,7 @@ Add the following dependency to your `pom.xml`: com.lancedb lancedb-core - 0.38.0-beta.10 + 0.38.0-beta.11 ``` diff --git a/java/lancedb-core/pom.xml b/java/lancedb-core/pom.xml index 25e3b10e3..c6acc7dfe 100644 --- a/java/lancedb-core/pom.xml +++ b/java/lancedb-core/pom.xml @@ -8,7 +8,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.10 + 0.38.0-beta.11 ../pom.xml diff --git a/java/pom.xml b/java/pom.xml index a2cec19c0..0b85b69df 100644 --- a/java/pom.xml +++ b/java/pom.xml @@ -6,7 +6,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.10 + 0.38.0-beta.11 pom ${project.artifactId} LanceDB Java SDK Parent POM diff --git a/nodejs/Cargo.toml b/nodejs/Cargo.toml index fd08e7a5e..6496c6384 100644 --- a/nodejs/Cargo.toml +++ b/nodejs/Cargo.toml @@ -1,7 +1,7 @@ [package] name = "lancedb-nodejs" edition.workspace = true -version = "0.38.0-beta.10" +version = "0.38.0-beta.11" publish = false license.workspace = true description.workspace = true diff --git a/nodejs/__test__/arrow.test.ts b/nodejs/__test__/arrow.test.ts index c5bbbf169..cb56cb5ae 100644 --- a/nodejs/__test__/arrow.test.ts +++ b/nodejs/__test__/arrow.test.ts @@ -6,6 +6,9 @@ import * as arrow17 from "apache-arrow-17"; import * as arrow18 from "apache-arrow-18"; import { + Field as CurrentField, + LargeBinary as CurrentLargeBinary, + Schema as CurrentSchema, Vector as CurrentVector, convertToTable, tableFromIPC as currentTableFromIPC, @@ -36,6 +39,24 @@ function sampleRecords(): Array> { }, ]; } + +it("preserves field metadata from a provided schema", async function () { + const jsonMetadata = new Map([["ARROW:extension:name", "lance.json"]]); + const schema = new CurrentSchema([ + new CurrentField("meta", new CurrentLargeBinary(), true, jsonMetadata), + ]); + + const table = makeArrowTable( + [{ meta: Buffer.from(JSON.stringify({ source: "test" })) }], + { schema }, + ); + + expect(table.schema.fields[0].metadata).toEqual(jsonMetadata); + + const roundTripped = currentTableFromIPC(await fromTableToBuffer(table)); + expect(roundTripped.schema.fields[0].metadata).toEqual(jsonMetadata); +}); + describe.each([arrow15, arrow16, arrow17, arrow18])( "Arrow", ( diff --git a/nodejs/__test__/embedding.test.ts b/nodejs/__test__/embedding.test.ts index 2a8494e0f..45d171a3d 100644 --- a/nodejs/__test__/embedding.test.ts +++ b/nodejs/__test__/embedding.test.ts @@ -187,6 +187,58 @@ describe("embedding functions", () => { const vector0 = JSON.parse(JSON.stringify(arr[0].vector)); expect(vector0).toEqual([1, 2, 3]); }); + it("should append multiple Python embeddings with the same alias", async () => { + @register("python-mock") + // biome-ignore lint/correctness/noUnusedVariables: the decorator registers this class + class MockEmbeddingFunction extends EmbeddingFunction { + ndims() { + return 3; + } + embeddingDataType(): Float { + return new Float32(); + } + async computeQueryEmbeddings(_data: string) { + return [1, 2, 3]; + } + async computeSourceEmbeddings(data: string[]) { + return data.map((value) => + value === "hello world" ? [1, 2, 3] : [4, 5, 6], + ); + } + } + + const metadata = new Map([ + [ + "embedding_functions", + '[{"source_column":"text1","vector_column":"vector1","name":"python-mock","model":{}},{"source_column":"text2","vector_column":"vector2","name":"python-mock","model":{}}]', + ], + ]); + const schema = new Schema( + [ + new Field("text1", new Utf8(), true), + new Field("text2", new Utf8(), true), + new Field( + "vector1", + new FixedSizeList(3, new Field("item", new Float32(), true)), + true, + ), + new Field( + "vector2", + new FixedSizeList(3, new Field("item", new Float32(), true)), + true, + ), + ], + metadata, + ); + + const db = await connect(tmpDir.name); + const table = await db.createEmptyTable("test", schema); + await table.add([{ text1: "hello world", text2: "goodbye world" }]); + + const rows = await table.query().toArray(); + expect(JSON.parse(JSON.stringify(rows[0].vector1))).toEqual([1, 2, 3]); + expect(JSON.parse(JSON.stringify(rows[0].vector2))).toEqual([4, 5, 6]); + }); it("should append generated vectors to a non-nullable schema", async () => { @register("non_nullable_schema_test") diff --git a/nodejs/__test__/table.test.ts b/nodejs/__test__/table.test.ts index dae640850..554c7fcd3 100644 --- a/nodejs/__test__/table.test.ts +++ b/nodejs/__test__/table.test.ts @@ -3561,6 +3561,27 @@ describe("when creating an empty table", () => { expect((actualSchema.fields[1].type as Float64).precision).toBe(2); }); + it("can add and query JSON data", async () => { + const schema = new Schema([ + new Field("id", new Int32(), true), + new Field( + "meta", + new Utf8(), + true, + new Map([["ARROW:extension:name", "arrow.json"]]), + ), + ]); + const table = await con.createEmptyTable("json", schema); + const meta = JSON.stringify({ x: 1 }); + + await table.add([{ id: 1, meta }]); + + const rows = await table.query().toArray(); + expect(rows).toHaveLength(1); + expect(rows[0].id).toBe(1); + expect(rows[0].meta).toBe(meta); + }); + it("can create an empty table from schema that specifies field types by name", async () => { const schemaLike = { fields: [ 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/nodejs/lancedb/schema.ts b/nodejs/lancedb/schema.ts index e4749ef37..3a3ee9316 100644 --- a/nodejs/lancedb/schema.ts +++ b/nodejs/lancedb/schema.ts @@ -406,10 +406,11 @@ function matchingFields(fields: Field[], tree: FieldTree): Field[] { field.name, new Struct(matchingFields(struct.children, value)), field.nullable, + field.metadata, ), ); } else { - matches.push(new Field(field.name, value as DataType, field.nullable)); + matches.push(field); } } return matches; diff --git a/nodejs/npm/darwin-arm64/package.json b/nodejs/npm/darwin-arm64/package.json index ff6347c4d..38b3db7d7 100644 --- a/nodejs/npm/darwin-arm64/package.json +++ b/nodejs/npm/darwin-arm64/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-darwin-arm64", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.11", "os": ["darwin"], "cpu": ["arm64"], "main": "lancedb.darwin-arm64.node", diff --git a/nodejs/npm/linux-arm64-gnu/package.json b/nodejs/npm/linux-arm64-gnu/package.json index ed99fee05..a256cb1ee 100644 --- a/nodejs/npm/linux-arm64-gnu/package.json +++ b/nodejs/npm/linux-arm64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-gnu", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.11", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-gnu.node", diff --git a/nodejs/npm/linux-arm64-musl/package.json b/nodejs/npm/linux-arm64-musl/package.json index 5b215dcc0..567b785b0 100644 --- a/nodejs/npm/linux-arm64-musl/package.json +++ b/nodejs/npm/linux-arm64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-musl", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.11", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-musl.node", diff --git a/nodejs/npm/linux-x64-gnu/package.json b/nodejs/npm/linux-x64-gnu/package.json index e0f5a9f26..4443a2748 100644 --- a/nodejs/npm/linux-x64-gnu/package.json +++ b/nodejs/npm/linux-x64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-gnu", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.11", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-gnu.node", diff --git a/nodejs/npm/linux-x64-musl/package.json b/nodejs/npm/linux-x64-musl/package.json index d42541707..5c0710d56 100644 --- a/nodejs/npm/linux-x64-musl/package.json +++ b/nodejs/npm/linux-x64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-musl", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.11", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-musl.node", diff --git a/nodejs/npm/win32-arm64-msvc/package.json b/nodejs/npm/win32-arm64-msvc/package.json index 496a40720..648c985f1 100644 --- a/nodejs/npm/win32-arm64-msvc/package.json +++ b/nodejs/npm/win32-arm64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-arm64-msvc", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.11", "os": [ "win32" ], diff --git a/nodejs/npm/win32-x64-msvc/package.json b/nodejs/npm/win32-x64-msvc/package.json index 734013343..1f3bfdeb8 100644 --- a/nodejs/npm/win32-x64-msvc/package.json +++ b/nodejs/npm/win32-x64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-x64-msvc", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.11", "os": ["win32"], "cpu": ["x64"], "main": "lancedb.win32-x64-msvc.node", diff --git a/nodejs/package-lock.json b/nodejs/package-lock.json index 732a7c01c..d46f08628 100644 --- a/nodejs/package-lock.json +++ b/nodejs/package-lock.json @@ -1,12 +1,12 @@ { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.11", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.11", "cpu": [ "x64", "arm64" diff --git a/nodejs/package.json b/nodejs/package.json index 262d4c4a7..665e2a522 100644 --- a/nodejs/package.json +++ b/nodejs/package.json @@ -11,7 +11,7 @@ "ann" ], "private": false, - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.11", "main": "dist/index.js", "exports": { ".": "./dist/index.js", diff --git a/python/Cargo.toml b/python/Cargo.toml index 3a8a05522..b97fad0ed 100644 --- a/python/Cargo.toml +++ b/python/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb-python" -version = "0.38.0-beta.10" +version = "0.38.0-beta.11" publish = false edition.workspace = true description = "Python bindings for LanceDB" 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/Cargo.toml b/rust/lancedb/Cargo.toml index 8276e5bb3..a71e3c948 100644 --- a/rust/lancedb/Cargo.toml +++ b/rust/lancedb/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb" -version = "0.38.0-beta.10" +version = "0.38.0-beta.11" edition.workspace = true description = "LanceDB: A serverless, low-latency vector database for AI applications" license.workspace = true diff --git a/rust/lancedb/src/expr.rs b/rust/lancedb/src/expr.rs index 75cce443d..da69914e3 100644 --- a/rust/lancedb/src/expr.rs +++ b/rust/lancedb/src/expr.rs @@ -19,6 +19,7 @@ mod sql; +pub(crate) use sql::canonicalize_sql_predicate; pub use sql::expr_to_sql_string; use std::sync::Arc; diff --git a/rust/lancedb/src/expr/sql.rs b/rust/lancedb/src/expr/sql.rs index 23b89821a..24a676485 100644 --- a/rust/lancedb/src/expr/sql.rs +++ b/rust/lancedb/src/expr/sql.rs @@ -1,10 +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::unparser::{self, dialect::Dialect}; +use datafusion_sql::sqlparser::{ + dialect::{Dialect as SqlParserDialect, GenericDialect}, + tokenizer::{Token, Tokenizer}, +}; +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 @@ -19,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 @@ -30,6 +36,61 @@ 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 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 = PredicateDialect::default(); + let tokens = Tokenizer::new(&dialect, predicate) + .with_unescape(false) + .tokenize() + .map_err(|err| crate::Error::InvalidInput { + message: format!("invalid SQL predicate: {err}"), + })?; + + 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}`") + } + other => other.to_string(), + }) + .collect()) +} + /// Prefix for placeholder strings inserted in place of binary literals. Chosen /// to be extremely unlikely to occur in user data. const BINARY_PLACEHOLDER_PREFIX: &str = "__lancedb_binary_placeholder_"; @@ -113,3 +174,51 @@ pub fn expr_to_sql_string(expr: &Expr) -> crate::Result { } Ok(sql) } + +#[cfg(test)] +mod tests { + use super::canonicalize_sql_predicate; + + #[test] + fn normalizes_double_quoted_identifiers() { + assert_eq!( + canonicalize_sql_predicate(r#""PartyAbbrev" = 'D'"#).unwrap(), + "`PartyAbbrev` = 'D'" + ); + assert_eq!( + canonicalize_sql_predicate(r#""MetaData"."userId" = 5"#).unwrap(), + "`MetaData`.`userId` = 5" + ); + 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!(canonicalize_sql_predicate(filter).unwrap(), filter); + } + + #[test] + 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(), + 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); + + let predicate = r#"id = 1 /*! OR "PartyAbbrev" = 'D' */"#; + assert_eq!(canonicalize_sql_predicate(predicate).unwrap(), predicate); + } + + #[test] + fn rejects_unterminated_double_quoted_identifier() { + let error = canonicalize_sql_predicate(r#""PartyAbbrev = 'D'"#).unwrap_err(); + assert!(matches!(error, crate::Error::InvalidInput { .. })); + } +} diff --git a/rust/lancedb/src/materialized_view.rs b/rust/lancedb/src/materialized_view.rs index b28d52931..08d6c921e 100644 --- a/rust/lancedb/src/materialized_view.rs +++ b/rust/lancedb/src/materialized_view.rs @@ -170,6 +170,15 @@ pub(crate) fn plan( filter: Option<&str>, limit: Option, ) -> Result<(MaterializedViewDefinition, Vec, Lineage)> { + let filter = filter + .map(crate::expr::canonicalize_sql_predicate) + .transpose() + .map_err(|err| match err { + Error::InvalidInput { message } => Error::InvalidInput { + message: format!("invalid view filter: {message}"), + }, + err => err, + })?; let projections: Vec<(String, String)> = if projections.is_empty() { source_schema .fields() @@ -274,7 +283,7 @@ pub(crate) fn plan( declared.push(output); } - if let Some(filter) = filter { + if let Some(filter) = filter.as_deref() { let expr = planner .parse_filter(filter) .map_err(|e| Error::InvalidInput { @@ -314,7 +323,7 @@ pub(crate) fn plan( .into_iter() .map(|(output, expression)| ViewProjection { output, expression }) .collect(), - filter: filter.map(String::from), + filter, limit, inputs, }; diff --git a/rust/lancedb/src/materialized_view/refresh.rs b/rust/lancedb/src/materialized_view/refresh.rs index 735751c27..b967e81f8 100644 --- a/rust/lancedb/src/materialized_view/refresh.rs +++ b/rust/lancedb/src/materialized_view/refresh.rs @@ -46,8 +46,9 @@ use lance_table::format::Fragment; use serde::{Deserialize, Serialize}; use super::{ - INCARNATION_META_KEY, MaterializedViewDefinition, REFRESHED_AT_MS_META_KEY, - SOURCE_ROW_ID_COLUMN, SOURCE_VERSION_META_KEY, + DEFINITION_META_KEY, INCARNATION_META_KEY, MaterializedViewDefinition, + REFRESHED_AT_MS_META_KEY, SOURCE_ROW_ID_COLUMN, SOURCE_VERSION_META_KEY, + definition_to_metadata, }; use crate::database::OpenTableRequest; use crate::table::{NativeTable, NativeTableExt, Table}; @@ -197,8 +198,28 @@ pub(crate) async fn execute_refresh( ), }); } + let definition_changed = + definition.filter != replanned.filter || definition.inputs != replanned.inputs; let definition = &replanned; + // A watermark written for a legacy raw filter certifies the rows that + // filter produced, not the canonical predicate above. Rebuild instead of + // accepting or advancing it, and persist the migrated definition in the + // same metadata commit that certifies the replacement rows. + if definition_changed { + return rebuild( + view_native, + &view_ds, + &source_ds, + source_version, + source_ts, + definition, + true, + expected_incarnation, + ) + .await; + } + let metadata = &view_ds.schema().metadata; let watermark: Option = metadata .get(SOURCE_VERSION_META_KEY) @@ -257,6 +278,7 @@ pub(crate) async fn execute_refresh( source_version, source_ts, definition, + false, expected_incarnation, ) .await @@ -271,6 +293,7 @@ pub(crate) async fn execute_refresh( source_version, source_ts, definition, + false, expected_incarnation, ) .await @@ -683,6 +706,7 @@ async fn incremental( view_ds.clone(), source_version, source_ts, + None, expected_incarnation, ) .await?; @@ -704,6 +728,7 @@ async fn incremental( published, source_version, source_ts, + None, expected_incarnation, ) .await?; @@ -775,6 +800,7 @@ async fn incremental( published, source_version, source_ts, + None, expected_incarnation, ) .await?; @@ -824,12 +850,14 @@ async fn incremental( appended, source_version, source_ts, + None, expected_incarnation, ) .await?; Ok(Some(result)) } +#[allow(clippy::too_many_arguments)] async fn rebuild( view_native: &NativeTable, view_ds: &Dataset, @@ -837,6 +865,7 @@ async fn rebuild( source_version: u64, source_ts: u128, definition: &MaterializedViewDefinition, + persist_definition: bool, expected_incarnation: Option<&str>, ) -> Result { let rows_written = Arc::new(AtomicU64::new(0)); @@ -867,6 +896,7 @@ async fn rebuild( replaced, source_version, source_ts, + persist_definition.then_some(definition), expected_incarnation, ) .await?; @@ -981,6 +1011,7 @@ async fn stamp_watermark( mut dataset: Dataset, source_version: u64, source_ts: u128, + definition: Option<&MaterializedViewDefinition>, expected_incarnation: Option<&str>, ) -> Result { ensure_incarnation(&dataset, expected_incarnation, dataset.uri()).await?; @@ -993,27 +1024,32 @@ async fn stamp_watermark( .get(INCARNATION_META_KEY) .cloned() .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); - dataset - .update_schema_metadata([ - (INCARNATION_META_KEY.to_string(), Some(incarnation)), - ( - SOURCE_VERSION_META_KEY.to_string(), - Some(source_version.to_string()), - ), - ( - SOURCE_VERSION_TS_META_KEY.to_string(), - Some(source_ts.to_string()), - ), - ( - REFRESHED_AT_MS_META_KEY.to_string(), - Some(now_ms().to_string()), - ), - ( - VIEW_VERSION_META_KEY.to_string(), - Some(predicted.to_string()), - ), - ]) - .await?; + let mut metadata = vec![(INCARNATION_META_KEY.to_string(), Some(incarnation))]; + if let Some(definition) = definition { + metadata.push(( + DEFINITION_META_KEY.to_string(), + Some(definition_to_metadata(definition)?), + )); + } + metadata.extend([ + ( + SOURCE_VERSION_META_KEY.to_string(), + Some(source_version.to_string()), + ), + ( + SOURCE_VERSION_TS_META_KEY.to_string(), + Some(source_ts.to_string()), + ), + ( + REFRESHED_AT_MS_META_KEY.to_string(), + Some(now_ms().to_string()), + ), + ( + VIEW_VERSION_META_KEY.to_string(), + Some(predicted.to_string()), + ), + ]); + dataset.update_schema_metadata(metadata).await?; let actual = dataset.version().version; if actual != predicted { return Err(Error::Runtime { @@ -1585,6 +1621,106 @@ mod tests { assert_eq!(read(view.table(), "x").await, vec![20, 40]); } + #[tokio::test] + async fn test_mixed_case_filter_is_canonicalized_for_lineage_and_refresh() { + let conn = connect("memory://").execute().await.unwrap(); + let batch = record_batch!( + ("id", Int32, [1, 2, 3]), + ("PartyAbbrev", Utf8, ["D", "R", "D"]) + ) + .unwrap(); + conn.create_table("src", batch) + .write_options(crate::materialized_view::tests::stable_row_ids()) + .execute() + .await + .unwrap(); + conn.create_materialized_view("democrats", "src") + .select([("id", "id")]) + .only_if(r#""PartyAbbrev" = 'D'"#) + .execute() + .await + .unwrap(); + + // Reopen from schema metadata so these assertions cover the stored + // predicate and lineage, not only the declaration-time handle. + let view = conn.open_materialized_view("democrats").await.unwrap(); + assert_eq!( + view.definition().filter.as_deref(), + Some("`PartyAbbrev` = 'D'") + ); + assert_eq!(view.definition().inputs, ["PartyAbbrev", "id"]); + + let result = view.refresh().execute().await.unwrap(); + assert_eq!(result.rows_written, 2); + assert_eq!(read(view.table(), "id").await, vec![1, 3]); + } + + #[tokio::test] + async fn test_legacy_raw_filter_rebuilds_and_persists_canonical_definition() { + let conn = connect("memory://").execute().await.unwrap(); + let batch = record_batch!( + ("id", Int32, [1, 2, 3]), + ("PartyAbbrev", Utf8, ["D", "R", "D"]) + ) + .unwrap(); + conn.create_table("legacy_src", batch) + .write_options(crate::materialized_view::tests::stable_row_ids()) + .execute() + .await + .unwrap(); + let view = conn + .create_materialized_view("legacy_view", "legacy_src") + .select([("id", "id")]) + .only_if(r#""PartyAbbrev" = 'X'"#) + .execute() + .await + .unwrap(); + assert_eq!(view.refresh().execute().await.unwrap().rows_written, 0); + + // Model a definition and up-to-date watermark written before filter + // canonicalization was applied to materialized views. + let mut legacy = view.definition().clone(); + legacy.filter = Some(r#""PartyAbbrev" = 'D'"#.into()); + legacy.inputs = vec!["id".into()]; + let native = view.table().as_native().unwrap(); + let mut dataset = native.dataset.get().await.unwrap().as_ref().clone(); + let predicted = dataset.version().version + 1; + dataset + .update_schema_metadata([ + ( + DEFINITION_META_KEY.to_string(), + Some(definition_to_metadata(&legacy).unwrap()), + ), + ( + VIEW_VERSION_META_KEY.to_string(), + Some(predicted.to_string()), + ), + ]) + .await + .unwrap(); + native.dataset.update(dataset); + + let reopened = conn.open_materialized_view("legacy_view").await.unwrap(); + let result = reopened.refresh().execute().await.unwrap(); + assert_eq!(result.mode, RefreshMode::Rebuild); + assert_eq!(result.rows_written, 2); + assert_eq!(read(reopened.table(), "id").await, vec![1, 3]); + + // A fresh handle proves the migration was stored alongside the new + // watermark and therefore happens only once. + let migrated = conn.open_materialized_view("legacy_view").await.unwrap(); + assert_eq!( + migrated.definition().filter.as_deref(), + Some("`PartyAbbrev` = 'D'") + ); + assert_eq!(migrated.definition().inputs, ["PartyAbbrev", "id"]); + assert_eq!( + migrated.refresh().execute().await.unwrap().mode, + RefreshMode::NoOp + ); + assert_eq!(read(migrated.table(), "id").await, vec![1, 3]); + } + #[tokio::test] async fn test_append_refreshes_incrementally() { let (_conn, source, view) = refreshed_doubled(vec![1, 2]).await; @@ -2767,7 +2903,7 @@ mod tests { let stale = view_native.dataset.get().await.unwrap().as_ref().clone(); view.table().delete("x = 1").await.unwrap(); - let err = stamp_watermark(view_native, stale, 99, 99, None).await; + let err = stamp_watermark(view_native, stale, 99, 99, None, None).await; assert!(err.is_err()); let result = view.refresh().execute().await.unwrap(); diff --git a/rust/lancedb/src/query.rs b/rust/lancedb/src/query.rs index 2a1283f22..cd346f42e 100644 --- a/rust/lancedb/src/query.rs +++ b/rust/lancedb/src/query.rs @@ -399,6 +399,9 @@ pub trait QueryBase { /// x > 5 OR y = 'test' /// ``` /// + /// Identifiers may be delimited with SQL-standard double quotes or + /// backticks. String literals must use single quotes. + /// /// Filtering performance can often be improved by creating a scalar index /// on the filter column(s). /// @@ -913,6 +916,17 @@ impl QueryRequest { /// use different representations) the error is recorded and surfaced later /// by [`Self::check_filter`]. pub(crate) fn add_filter(&mut self, new: QueryFilter) { + let new = match new { + QueryFilter::Sql(filter) => match crate::expr::canonicalize_sql_predicate(&filter) { + Ok(filter) => QueryFilter::Sql(filter), + Err(err) => { + self.filter_error = Some(err.to_string()); + return; + } + }, + other => other, + }; + self.filter = Some(match self.filter.take() { None => new, Some(existing) => match and_filters(existing, new) { @@ -1652,8 +1666,8 @@ mod tests { datatypes::{Int32Type, UInt8Type}, }; 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}; @@ -1882,6 +1896,157 @@ mod tests { query.execute().await.unwrap(); } + #[tokio::test] + 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("id", DataType::Int32, false), + ArrowField::new("PartyAbbrev", DataType::Utf8, false), + ArrowField::new("path", DataType::Utf8, false), + ])); + let batch = RecordBatch::try_new( + 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(); + + let conn = connect(uri).execute().await.unwrap(); + let table = conn.create_table("parties", batch).execute().await.unwrap(); + let batches = table + .query() + .only_if(r#""PartyAbbrev" = 'D'"#) + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + + assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::(), 2); + assert_eq!( + table + .count_rows(Some(r#""PartyAbbrev" = 'D'"#.to_string())) + .await + .unwrap(), + 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 + .query() + .only_if(predicate) + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::(), 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 = table + .base_table() + .merge_insert( + merge, + 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 + ); + + let update = table + .update() + .only_if(r#""PartyAbbrev" = 'R'"#) + .column("PartyAbbrev", "'X'"); + table.base_table().update(update).await.unwrap(); + assert_eq!( + table + .count_rows(Some(r#""PartyAbbrev" = 'X'"#.to_string())) + .await + .unwrap(), + 2 + ); + + 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); + } + #[tokio::test] async fn test_select_with_transform() { let batches = make_non_empty_batches(); diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index fad04a098..1afc2615a 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -1379,10 +1379,11 @@ impl RemoteTable { query: &AnyQuery, version: Option, ) -> Result> { + let query = query.canonicalized()?; 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)?; @@ -2491,7 +2492,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": read_snapshot.version }) @@ -2747,7 +2748,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 @@ -2794,7 +2796,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 }); @@ -2851,9 +2853,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; @@ -3864,13 +3867,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); } @@ -4353,7 +4360,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 { @@ -4369,14 +4376,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 }); @@ -4463,10 +4468,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")); @@ -4483,11 +4488,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 { @@ -4549,7 +4552,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() @@ -4567,7 +4570,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 }); } @@ -4659,6 +4666,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. @@ -4674,9 +4682,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 af8bcb5e2..8436657ca 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -1164,7 +1164,10 @@ impl Table { /// /// * `filter` if present, only count rows matching the filter pub async fn count_rows(&self, filter: Option) -> Result { - 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. @@ -1364,7 +1367,13 @@ impl Table { /// # }); /// ``` pub async fn delete(&self, predicate: impl Into>) -> Result { - 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). @@ -3239,7 +3248,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 6b87d1080..7a3cad071 100644 --- a/rust/lancedb/src/table/merge.rs +++ b/rust/lancedb/src/table/merge.rs @@ -220,9 +220,32 @@ 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) -> Result { + pub async fn execute( + mut self, + new_data: Box, + ) -> Result { + 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> { + 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 @@ -230,9 +253,10 @@ impl MergeInsertBuilder { /// 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, ) -> Result { + params.canonicalize_filters()?; super::computed_columns::ensure_no_function_bindings_for_mutation( table.schema().await?.as_ref(), "merge_insert", diff --git a/rust/lancedb/src/table/query.rs b/rust/lancedb/src/table/query.rs index 8658ad0b7..2737471bf 100644 --- a/rust/lancedb/src/table/query.rs +++ b/rust/lancedb/src/table/query.rs @@ -20,9 +20,9 @@ use arrow::array::{AsArray, FixedSizeListBuilder, Float32Builder}; use arrow::datatypes::{Float32Type, UInt8Type}; use arrow_array::Array; use arrow_schema::{DataType, Schema}; -use datafusion_common::ScalarValue; +use datafusion_common::{Column, DataFusionError, ScalarValue, SchemaError}; use datafusion_expr::Operator; -use datafusion_physical_expr::expressions::{BinaryExpr, Column, Literal}; +use datafusion_physical_expr::expressions::{BinaryExpr, Column as PhysicalColumn, Literal}; use datafusion_physical_plan::PhysicalExpr; use datafusion_physical_plan::projection::ProjectionExec; use datafusion_physical_plan::repartition::RepartitionExec; @@ -56,6 +56,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 @@ -64,15 +80,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 { @@ -147,9 +164,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()?; @@ -185,7 +203,7 @@ pub async fn create_plan( if query.query_vector.len() > 1 { if column.is_none() { // Infer a vector column with the same dimension of the query vector. - let arrow_schema = Schema::from(ds_ref.schema()); + let arrow_schema = Schema::from(schema); column = Some(default_vector_column( &arrow_schema, Some(query.query_vector[0].len() as i32), @@ -262,7 +280,7 @@ pub async fn create_plan( let column = if let Some(col) = column { col } else { - let arrow_schema = Schema::from(ds_ref.schema()); + let arrow_schema = Schema::from(schema); default_vector_column(&arrow_schema, Some(query_vector.len() as i32))? }; @@ -368,7 +386,10 @@ pub async fn create_plan( scanner.order_by(Some(order_by.clone()))?; } - let mut plan = scanner.create_plan().await?; + let mut plan = scanner + .create_plan() + .await + .map_err(|error| enrich_lance_field_not_found(error, schema))?; let normalized_l2_indices = normalized_l2_ann_indices(plan.as_ref()).await?; if !normalized_l2_indices.is_empty() { // Rebuild only the affected ANN nodes with internal normalized squared-L2 @@ -378,7 +399,10 @@ pub async fn create_plan( query.lower_bound.map(|bound| bound / COSINE_ANN_SCALE), query.upper_bound.map(|bound| bound / COSINE_ANN_SCALE), ); - scanner.create_plan().await? + scanner + .create_plan() + .await + .map_err(|error| enrich_lance_field_not_found(error, schema))? } else { plan.clone() }; @@ -388,6 +412,93 @@ pub async fn create_plan( Ok(plan) } +/// Replace DataFusion's top-level field candidates with qualified leaf paths. +/// +/// DataFusion resolves nested fields but its `FieldNotFound` error only lists the +/// top-level Arrow fields. This makes a missing leaf look unavailable even when it +/// exists below a struct. Keep every other Lance/DataFusion error unchanged and +/// enrich only this one schema error at the LanceDB query boundary. +fn enrich_lance_field_not_found( + error: lance::Error, + schema: &lance_core::datatypes::Schema, +) -> Error { + let Some(field) = find_missing_field(&error) else { + return error.into(); + }; + field_not_found_error(field, &Schema::from(schema)) +} + +fn field_not_found_diagnostic( + error: &(dyn std::error::Error + 'static), + schema: &Schema, +) -> Option { + let field = find_missing_field(error)?; + Some(field_not_found_error(field, schema)) +} + +fn field_not_found_error(field: &Column, schema: &Schema) -> Error { + let valid_fields = leaf_field_paths(schema); + let mut message = format!("Schema error: No field named {}", field.quoted_flat_name()); + if !valid_fields.is_empty() { + message.push_str(". Valid fields are "); + message.push_str(&valid_fields.join(", ")); + } + message.push('.'); + + Error::InvalidInput { message } +} + +fn find_missing_field<'a>(error: &'a (dyn std::error::Error + 'static)) -> Option<&'a Column> { + if let Some(DataFusionError::SchemaError(schema_error, _)) = + error.downcast_ref::() + && let SchemaError::FieldNotFound { field, .. } = schema_error.as_ref() + { + return Some(field); + } + + error.source().and_then(find_missing_field) +} + +fn leaf_field_paths(schema: &Schema) -> Vec { + fn format_segment(segment: &str) -> String { + // Quote every segment instead of maintaining a SQL keyword list. Bare + // lowercase names such as `true` can be parsed as expressions rather + // than identifiers, while backticks preserve all field names in both + // local SQL parsers. + format!("`{}`", segment.replace('`', "``")) + } + + fn visit(fields: &arrow_schema::Fields, path: &mut Vec, paths: &mut Vec) { + for field in fields { + // Neither local planner can address an empty field-path segment, + // even when it is backtick-quoted. Do not advertise leaves beneath + // such a segment as valid filter fields. + if field.name().is_empty() { + continue; + } + path.push(field.name().clone()); + match field.data_type() { + DataType::Struct(children) if !children.is_empty() => { + visit(children, path, paths); + } + _ => { + paths.push( + path.iter() + .map(|segment| format_segment(segment)) + .collect::>() + .join("."), + ); + } + } + path.pop(); + } + } + + let mut paths = Vec::new(); + visit(schema.fields(), &mut Vec::new(), &mut paths); + paths +} + //Helper functions below const COSINE_ANN_SCALE: f32 = 0.5; @@ -553,7 +664,7 @@ fn scale_distance_column( .iter() .enumerate() .map(|(index, field)| { - let column: Arc = Arc::new(Column::new(field.name(), index)); + let column: Arc = Arc::new(PhysicalColumn::new(field.name(), index)); let expression = if field.name() == DIST_COL { let scale: Arc = Arc::new(Literal::new(ScalarValue::Float32(Some(scale)))); @@ -923,7 +1034,10 @@ async fn parse_arrow_ipc_response(bytes: bytes::Bytes) -> Result Field { + let mut segments = path.iter().rev(); + let mut field = Field::new( + *segments.next().expect("path must have a leaf"), + DataType::Int32, + false, + ); + for segment in segments { + field = Field::new(*segment, DataType::Struct(vec![field].into()), false); + } + field + } + + let schema = Schema::new(vec![ + nested_field(&["a", "b", "c", "d", "e"]), + nested_field(&["metadata", "child.with.dot"]), + nested_field(&["metadata", "Title"]), + nested_field(&["metadata", "123child"]), + nested_field(&["metadata", "child`tick"]), + nested_field(&["metadata", ""]), + nested_field(&["", "child"]), + ]); + + assert_eq!( + leaf_field_paths(&schema), + vec![ + "`a`.`b`.`c`.`d`.`e`", + "`metadata`.`child.with.dot`", + "`metadata`.`Title`", + "`metadata`.`123child`", + "`metadata`.`child``tick`", + ] + ); + + let source = DataFusionError::SchemaError( + Box::new(SchemaError::FieldNotFound { + field: Box::new(Column::from_name("missing")), + valid_fields: Vec::new(), + }), + Box::new(None), + ); + let error = field_not_found_diagnostic(&source, &schema).unwrap(); + assert!( + error.to_string().contains( + "Valid fields are `a`.`b`.`c`.`d`.`e`, `metadata`.`child.with.dot`, `metadata`.`Title`, `metadata`.`123child`, `metadata`.`child``tick`" + ), + "unexpected error: {error}" + ); + } + #[derive(Debug, Default)] struct CountingNamespaceClient { query_table_calls: AtomicUsize, diff --git a/rust/lancedb/src/table/query/lsm.rs b/rust/lancedb/src/table/query/lsm.rs index 7e0aa6f7a..447a6e358 100644 --- a/rust/lancedb/src/table/query/lsm.rs +++ b/rust/lancedb/src/table/query/lsm.rs @@ -27,6 +27,8 @@ use std::sync::Arc; use arrow_array::Array; use arrow_schema::{DataType, Schema as ArrowSchema}; +use datafusion::common::{DataFusionError, ToDFSchema}; +use datafusion::prelude::SessionContext; use datafusion_physical_plan::expressions::Column; use datafusion_physical_plan::projection::ProjectionExec; use datafusion_physical_plan::{ExecutionPlan, PhysicalExpr}; @@ -395,7 +397,21 @@ fn base_scanner( } if let Some(filter) = &query.base.filter { scanner = match filter { - QueryFilter::Sql(sql) => scanner.filter(sql)?, + QueryFilter::Sql(sql) => { + // Parse here instead of inside `LsmScanner::filter` so the typed + // DataFusion `FieldNotFound` error is still available for the + // same nested-field enrichment used by the ordinary scanner. + let schema = ArrowSchema::from(dataset.schema()); + let df_schema = schema.clone().to_dfschema().map_err(|error| { + enrich_filter_error(error, &schema, "Failed to create DFSchema") + })?; + let expr = SessionContext::new() + .parse_sql_expr(sql, &df_schema) + .map_err(|error| { + enrich_filter_error(error, &schema, "Failed to parse filter expression") + })?; + scanner.filter_expr(expr) + } QueryFilter::Datafusion(expr) => scanner.filter_expr(expr.clone()), QueryFilter::Substrait(_) => { return Err(Error::NotSupported { @@ -407,6 +423,12 @@ fn base_scanner( Ok(scanner) } +fn enrich_filter_error(error: DataFusionError, schema: &ArrowSchema, context: &str) -> Error { + super::field_not_found_diagnostic(&error, schema).unwrap_or_else(|| Error::InvalidInput { + message: format!("{context}: {error}"), + }) +} + /// Plain scan: filter / projection / limit over base ∪ SSTables ∪ in-memory. /// The plain scan applies limit and offset inside the planner. async fn plain_plan( diff --git a/rust/lancedb/src/table/update.rs b/rust/lancedb/src/table/update.rs index 98050dfe8..f10594f23 100644 --- a/rust/lancedb/src/table/update.rs +++ b/rust/lancedb/src/table/update.rs @@ -62,22 +62,33 @@ impl UpdateBuilder { } /// Executes the update operation. - pub async fn execute(self) -> Result { + pub async fn execute(mut self) -> Result { if self.columns.is_empty() { Err(Error::InvalidInput { message: "at least one column must be specified in an update operation".to_string(), }) } else { + 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