diff --git a/Cargo.lock b/Cargo.lock index 33e43f9fa..e27f9f271 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3455,8 +3455,8 @@ checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" [[package]] name = "fsst" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "rand 0.9.5", @@ -4815,8 +4815,8 @@ checksum = "e037a2e1d8d5fdbd49b16a4ea09d5d6401c1f29eca5ff29d03d3824dba16256a" [[package]] name = "lance" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arc-swap", "arrow", @@ -4888,8 +4888,8 @@ dependencies = [ [[package]] name = "lance-arrow" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-buffer", @@ -4911,7 +4911,7 @@ dependencies = [ [[package]] name = "lance-arrow-scalar" version = "58.0.0" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-buffer", @@ -4925,7 +4925,7 @@ dependencies = [ [[package]] name = "lance-arrow-stats" version = "58.0.0" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-schema", @@ -4934,8 +4934,8 @@ dependencies = [ [[package]] name = "lance-bitpacking" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrayref", "crunchy", @@ -4945,8 +4945,8 @@ dependencies = [ [[package]] name = "lance-core" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-buffer", @@ -4983,8 +4983,8 @@ dependencies = [ [[package]] name = "lance-datafusion" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow", "arrow-array", @@ -5013,8 +5013,8 @@ dependencies = [ [[package]] name = "lance-datagen" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow", "arrow-array", @@ -5031,8 +5031,8 @@ dependencies = [ [[package]] name = "lance-derive" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "proc-macro2", "quote", @@ -5041,8 +5041,8 @@ dependencies = [ [[package]] name = "lance-encoding" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-arith", "arrow-array", @@ -5075,8 +5075,8 @@ dependencies = [ [[package]] name = "lance-file" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-arith", "arrow-array", @@ -5107,8 +5107,8 @@ dependencies = [ [[package]] name = "lance-index" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arc-swap", "arrow", @@ -5172,8 +5172,8 @@ dependencies = [ [[package]] name = "lance-index-core" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-schema", @@ -5195,8 +5195,8 @@ dependencies = [ [[package]] name = "lance-io" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow", "arrow-array", @@ -5236,8 +5236,8 @@ dependencies = [ [[package]] name = "lance-linalg" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-schema", @@ -5251,8 +5251,8 @@ dependencies = [ [[package]] name = "lance-namespace" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow", "async-trait", @@ -5264,8 +5264,8 @@ dependencies = [ [[package]] name = "lance-namespace-impls" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow", "arrow-ipc", @@ -5318,8 +5318,8 @@ dependencies = [ [[package]] name = "lance-select" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-buffer", @@ -5333,8 +5333,8 @@ dependencies = [ [[package]] name = "lance-table" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow", "arrow-array", @@ -5374,8 +5374,8 @@ dependencies = [ [[package]] name = "lance-testing" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-schema", @@ -5388,8 +5388,8 @@ dependencies = [ [[package]] name = "lance-tokenizer" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "frostem", "icu_segmenter", diff --git a/Cargo.toml b/Cargo.toml index 716b14d7a..a16f0412c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -13,20 +13,20 @@ categories = ["database-implementations"] rust-version = "1.91.0" [workspace.dependencies] -lance = { "version" = "=11.0.0-beta.22", default-features = false, "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-core = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-datagen = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-file = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-io = { "version" = "=11.0.0-beta.22", default-features = false, "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-index = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-linalg = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-namespace = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-namespace-impls = { "version" = "=11.0.0-beta.22", default-features = false, "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-table = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-testing = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-datafusion = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-encoding = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } -lance-arrow = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" } +lance = { "version" = "=12.0.0-beta.2", default-features = false, "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-core = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-datagen = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-file = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-io = { "version" = "=12.0.0-beta.2", default-features = false, "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-index = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-linalg = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-namespace = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-namespace-impls = { "version" = "=12.0.0-beta.2", default-features = false, "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-table = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-testing = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-datafusion = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-encoding = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-arrow = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } lancedb = { path = "rust/lancedb", default-features = false } ahash = "0.8" # Note that this one does not include pyarrow diff --git a/docs/src/js/classes/Table.md b/docs/src/js/classes/Table.md index 9a85d0d96..159348450 100644 --- a/docs/src/js/classes/Table.md +++ b/docs/src/js/classes/Table.md @@ -1292,6 +1292,18 @@ abstract updateFieldMetadata(updates): Promise Update per-field (column) metadata. +The following keys are treated specially, by convention, and should be +used when appropriate: + +- `lancedb:description`: for a human-readable description of a field. +- `lancedb:tag:`: for a user-defined key-value tag, where the suffix + names the tag category; e.g. `lancedb:tag:model: "clip"`. +- `lancedb:logical-column`: for a column grouping; e.g. `feature_v1` and + `feature_v2` might be in the same logical column. +- `lancedb:status`: for status options (`production`, `candidate`, + `deprecated`, `archived`) to designate the current life cycle state of + this column. + #### Parameters * **updates**: [`FieldMetadataUpdate`](../interfaces/FieldMetadataUpdate.md)[] diff --git a/docs/src/js/interfaces/FieldMetadataUpdate.md b/docs/src/js/interfaces/FieldMetadataUpdate.md index 38c675630..a85e3e7a0 100644 --- a/docs/src/js/interfaces/FieldMetadataUpdate.md +++ b/docs/src/js/interfaces/FieldMetadataUpdate.md @@ -17,7 +17,8 @@ metadata: Record; ``` Metadata key/value pairs. Merged into the field's existing metadata by -default; a value of `null` deletes that key. +default; a value of `null` deletes that key. See +[Table.updateFieldMetadata](../classes/Table.md#updatefieldmetadata) for the conventional `lancedb:*` keys. *** diff --git a/docs/src/python/python.md b/docs/src/python/python.md index 3cbeee6f0..3cb996a15 100644 --- a/docs/src/python/python.md +++ b/docs/src/python/python.md @@ -159,6 +159,8 @@ and combined with [BooleanQuery][lancedb.query.BooleanQuery]. ::: lancedb.query.FullTextOperator +::: lancedb.query.DocumentGranularity + ::: lancedb.query.Occur ## Embeddings diff --git a/java/pom.xml b/java/pom.xml index 4d83153ab..a2cec19c0 100644 --- a/java/pom.xml +++ b/java/pom.xml @@ -28,7 +28,7 @@ UTF-8 15.0.0 - 11.0.0-beta.22 + 12.0.0-beta.2 false 2.30.0 1.7 diff --git a/nodejs/__test__/table.test.ts b/nodejs/__test__/table.test.ts index 0aee5caf0..dae640850 100644 --- a/nodejs/__test__/table.test.ts +++ b/nodejs/__test__/table.test.ts @@ -685,6 +685,56 @@ describe.each([arrow15, arrow16, arrow17, arrow18])( }, ); +// https://github.com/lancedb/lancedb/issues/1963 +it("should query documents with LangChain PDF metadata", async () => { + const tmpDir = tmp.dirSync({ unsafeCleanup: true }); + try { + const db = await connect(tmpDir.name); + const documents = [ + { + text: "first page", + vector: [1, 0], + source: "first.pdf", + loc: { pageNumber: 1, lines: { from: 1, to: 12 } }, + pdf: { + version: "1.10.100", + info: { + format: "PDF 1.7", + producer: "pdf.js", + creator: "Writer", + }, + totalPages: 2, + }, + }, + { + text: "second page", + vector: [0, 1], + source: "second.pdf", + loc: { pageNumber: 2, lines: { from: 13, to: 24 } }, + pdf: { + version: "1.10.100", + info: { + format: "PDF 1.7", + producer: "pdf.js", + creator: "Writer", + }, + totalPages: 2, + }, + }, + ]; + const documentsTable = await db.createTable("documents", documents); + + const results = await documentsTable.query().toArray(); + + expect(results).toHaveLength(2); + expect(results[0].source).toBe("first.pdf"); + expect(results[0].pdf.info.producer).toBe("pdf.js"); + expect(results[1].loc.pageNumber).toBe(2); + } finally { + tmpDir.removeCallback(); + } +}); + describe("merge insert", () => { let tmpDir: tmp.DirResult; let table: Table; @@ -2535,7 +2585,24 @@ describe.each([arrow15, arrow16, arrow17, arrow18])( ); }); - test("full text search if no embedding function provided", async () => { + test("full text search if only an unrelated embedding function is registered", async () => { + register("unused")( + class extends EmbeddingFunction { + ndims() { + return 3; + } + embeddingDataType() { + return new Float32(); + } + async computeQueryEmbeddings(_data: string) { + return [1, 2, 3]; + } + async computeSourceEmbeddings(data: string[]) { + return data.map(() => [1, 2, 3]); + } + }, + ); + const db = await connect(tmpDir.name); const data = [ { text: "hello world", vector: [0.1, 0.2, 0.3] }, @@ -2557,6 +2624,306 @@ describe.each([arrow15, arrow16, arrow17, arrow18])( expect(results2[0].text).toBe(data[1].text); }); + test("auto search stays consistent with the active revision", async () => { + let initCalls = 0; + let queryCalls = 0; + let markStarted!: () => void; + const started = new Promise((resolve) => { + markStarted = resolve; + }); + let releaseEmbedding!: () => void; + const embeddingReleased = new Promise((resolve) => { + releaseEmbedding = resolve; + }); + + @register("refresh-test") + class TestEmbedding extends EmbeddingFunction { + async init() { + initCalls += 1; + } + ndims() { + return 1; + } + embeddingDataType() { + return new arrow.Float32(); + } + async computeQueryEmbeddings(value: string) { + queryCalls += 1; + if (value === "blocked") { + markStarted(); + await embeddingReleased; + } + return value === "greetings" ? [0.1] : [0.2]; + } + async computeSourceEmbeddings(values: string[]) { + return values.map((value) => + value === "hello world" ? [0.1] : [0.2], + ); + } + } + + const writer = await connect(tmpDir.name); + await writer.createTable("test", [{ text: "plain", vector: [0.0] }]); + const reader = await connect(tmpDir.name, { + readConsistencyInterval: 0, + }); + const tracked = await reader.openTable("test"); + type SnapshotCountingNative = { + querySnapshot: () => Promise; + }; + const native = (tracked as unknown as { inner: SnapshotCountingNative }) + .inner; + const querySnapshot = native.querySnapshot.bind(native); + let snapshotCalls = 0; + native.querySnapshot = async () => { + snapshotCalls += 1; + return await querySnapshot(); + }; + const autoQuery = tracked.search("greetings").select(["text"]).limit(1); + + const func = new TestEmbedding(); + const schema = LanceSchema({ + text: func.sourceField(new arrow.Utf8()), + vector: func.vectorField(), + }); + const data = [{ text: "hello world" }, { text: "goodbye world" }]; + await writer.createTable("test", data, { mode: "overwrite", schema }); + const baselineInitCalls = initCalls; + + expect( + (await tracked.schema()).metadata.get("embedding_functions"), + ).toBeDefined(); + const results = await autoQuery.toArray(); + expect(results[0].text).toBe(data[0].text); + expect(initCalls).toBe(baselineInitCalls + 1); + expect(queryCalls).toBe(1); + expect(snapshotCalls).toBe(1); + + const repeatedResults = await autoQuery.toArray(); + expect(repeatedResults[0].text).toBe(data[0].text); + expect(initCalls).toBe(baselineInitCalls + 1); + expect(queryCalls).toBe(1); + expect(snapshotCalls).toBe(2); + + const pending = tracked + .search("blocked") + .select(["text"]) + .limit(1) + .toArray(); + await started; + + const ftsData = [ + { text: "greetings from full text", vector: [0.0] }, + { text: "blocked from full text", vector: [0.0] }, + ]; + const ftsTable = await writer.createTable("test", ftsData, { + mode: "overwrite", + }); + await ftsTable.createIndex("text", { config: Index.fts() }); + releaseEmbedding(); + + const pendingResults = await pending; + expect(pendingResults[0].text).toBe(data[1].text); + + expect( + (await tracked.schema()).metadata.get("embedding_functions"), + ).toBeUndefined(); + const ftsResults = await autoQuery.toArray(); + expect(ftsResults[0].text).toBe(ftsData[0].text); + }); + + test("auto search keeps newer preparation during a revision race", async () => { + let aCalls = 0; + let bCalls = 0; + let markAStarted!: () => void; + const aStarted = new Promise((resolve) => { + markAStarted = resolve; + }); + let releaseA!: () => void; + const aReleased = new Promise((resolve) => { + releaseA = resolve; + }); + let markBStarted!: () => void; + const bStarted = new Promise((resolve) => { + markBStarted = resolve; + }); + let releaseB!: () => void; + const bReleased = new Promise((resolve) => { + releaseB = resolve; + }); + + @register("race-a") + class EmbeddingA extends EmbeddingFunction { + ndims() { + return 1; + } + embeddingDataType() { + return new arrow.Float32(); + } + async computeQueryEmbeddings() { + aCalls += 1; + markAStarted(); + await aReleased; + return [0.1]; + } + async computeSourceEmbeddings(values: string[]) { + return values.map(() => [0.1]); + } + } + + @register("race-b") + class EmbeddingB extends EmbeddingFunction { + ndims() { + return 1; + } + embeddingDataType() { + return new arrow.Float32(); + } + async computeQueryEmbeddings() { + bCalls += 1; + markBStarted(); + await bReleased; + return [0.2]; + } + async computeSourceEmbeddings(values: string[]) { + return values.map(() => [0.2]); + } + } + + const writer = await connect(tmpDir.name); + const embeddingA = new EmbeddingA(); + const schemaA = LanceSchema({ + text: embeddingA.sourceField(new arrow.Utf8()), + vector: embeddingA.vectorField(), + }); + await writer.createTable("race", [{ text: "revision a" }], { + schema: schemaA, + }); + const reader = await connect(tmpDir.name, { + readConsistencyInterval: 0, + }); + const tracked = await reader.openTable("race"); + const query = tracked.search("query"); + + const first = query.toArray(); + await aStarted; + + const embeddingB = new EmbeddingB(); + const schemaB = LanceSchema({ + text: embeddingB.sourceField(new arrow.Utf8()), + vector: embeddingB.vectorField(), + }); + await writer.createTable("race", [{ text: "revision b" }], { + mode: "overwrite", + schema: schemaB, + }); + const second = query.toArray(); + await bStarted; + + releaseA(); + releaseB(); + await Promise.all([first, second]); + expect(aCalls).toBe(1); + expect(bCalls).toBe(1); + }); + + test("stale FTS routing keeps newer vector preparation", async () => { + let vectorCalls = 0; + let markVectorStarted!: () => void; + const vectorStarted = new Promise((resolve) => { + markVectorStarted = resolve; + }); + let releaseVector!: () => void; + const vectorReleased = new Promise((resolve) => { + releaseVector = resolve; + }); + + @register("stale-fts-race") + class RaceEmbedding extends EmbeddingFunction { + ndims() { + return 1; + } + embeddingDataType() { + return new arrow.Float32(); + } + async computeQueryEmbeddings() { + vectorCalls += 1; + markVectorStarted(); + await vectorReleased; + return [0.1]; + } + async computeSourceEmbeddings(values: string[]) { + return values.map(() => [0.1]); + } + } + + const writer = await connect(tmpDir.name); + const ftsTable = await writer.createTable("stale_fts", [ + { text: "hello", vector: [0.0] }, + ]); + await ftsTable.createIndex("text", { config: Index.fts() }); + + const reader = await connect(tmpDir.name, { + readConsistencyInterval: 0, + }); + const tracked = await reader.openTable("stale_fts"); + type Snapshot = { + schema: () => Promise; + }; + type NativeWithSnapshot = { + querySnapshot: () => Promise; + }; + const native = (tracked as unknown as { inner: NativeWithSnapshot }) + .inner; + const querySnapshot = native.querySnapshot.bind(native); + let snapshotCalls = 0; + let markStaleSchemaStarted!: () => void; + const staleSchemaStarted = new Promise((resolve) => { + markStaleSchemaStarted = resolve; + }); + let releaseStaleSchema!: () => void; + const staleSchemaReleased = new Promise((resolve) => { + releaseStaleSchema = resolve; + }); + native.querySnapshot = async () => { + const snapshot = await querySnapshot(); + snapshotCalls += 1; + if (snapshotCalls === 1) { + const schema = snapshot.schema.bind(snapshot); + snapshot.schema = async () => { + markStaleSchemaStarted(); + await staleSchemaReleased; + return await schema(); + }; + } + return snapshot; + }; + + const query = tracked.search("hello"); + const staleFtsExecution = query.toArray(); + await staleSchemaStarted; + + const embedding = new RaceEmbedding(); + const vectorSchema = LanceSchema({ + text: embedding.sourceField(new arrow.Utf8()), + vector: embedding.vectorField(), + }); + await writer.createTable("stale_fts", [{ text: "hello" }], { + mode: "overwrite", + schema: vectorSchema, + }); + + const vectorExecution = query.toArray(); + await vectorStarted; + releaseStaleSchema(); + await staleFtsExecution; + releaseVector(); + await vectorExecution; + + await query.toArray(); + expect(vectorCalls).toBe(1); + }); + test("tokenizes FTS queries by column or index name", async () => { const db = await connect(tmpDir.name); const data = [ @@ -3107,6 +3474,30 @@ describe("column name options", () => { expect(results[1].query_index).toBe(1); }); + test("observes promised additional vectors while the query is pending", async () => { + const initialVector = new Promise(() => undefined); + const query = table.query().nearestTo(initialVector); + const unhandled: unknown[] = []; + const onUnhandled = (reason: unknown) => unhandled.push(reason); + process.on("unhandledRejection", onUnhandled); + + try { + query.addQueryVector(Promise.reject(new Error("extra vector failed"))); + await new Promise((resolve) => setImmediate(resolve)); + expect(unhandled).toEqual([]); + + const rejectedQuery = table + .query() + .nearestTo([0.1, 0.2]) + .addQueryVector(Promise.reject(new Error("consumed vector failed"))); + await expect(rejectedQuery.toArray()).rejects.toThrow( + "consumed vector failed", + ); + } finally { + process.off("unhandledRejection", onUnhandled); + } + }); + test("index and search multivectors", async () => { const db = await connect(tmpDir.name); const data = []; diff --git a/nodejs/lancedb/query.ts b/nodejs/lancedb/query.ts index 3b9b286a0..75a787fcd 100644 --- a/nodejs/lancedb/query.ts +++ b/nodejs/lancedb/query.ts @@ -100,6 +100,29 @@ export interface FullTextSearchOptions { columns?: string | string[]; } +function nearestToNative( + inner: NativeQuery, + vector: Awaited, +): NativeVectorQuery { + const raw = Array.isArray(vector) ? null : extractVectorBuffer(vector); + if (raw) { + return inner.nearestToRaw(raw.data, raw.dtype); + } + return inner.nearestTo(Float32Array.from(vector as number[])); +} + +function addQueryVectorToNative( + inner: NativeVectorQuery, + vector: Awaited, +) { + const raw = Array.isArray(vector) ? null : extractVectorBuffer(vector); + if (raw) { + inner.addQueryVectorRaw(raw.data, raw.dtype); + } else { + inner.addQueryVector(Float32Array.from(vector as number[])); + } +} + /** Common methods supported by all query types * * @see {@link Query} @@ -499,6 +522,13 @@ export class VectorQuery extends StandardQueryBase { super(inner); } + /** + * @hidden + */ + protected doVectorCall(fn: (inner: NativeVectorQuery) => void) { + super.doCall(fn); + } + /** * Set the number of partitions to search (probe) * @@ -526,7 +556,7 @@ export class VectorQuery extends StandardQueryBase { * the minimum and maximum to the same value. */ nprobes(nprobes: number): VectorQuery { - super.doCall((inner) => inner.nprobes(nprobes)); + this.doVectorCall((inner) => inner.nprobes(nprobes)); return this; } @@ -540,7 +570,7 @@ export class VectorQuery extends StandardQueryBase { * but will also increase latency. */ minimumNprobes(minimumNprobes: number): VectorQuery { - super.doCall((inner) => inner.minimumNprobes(minimumNprobes)); + this.doVectorCall((inner) => inner.minimumNprobes(minimumNprobes)); return this; } @@ -554,7 +584,7 @@ export class VectorQuery extends StandardQueryBase { * potential false negatives. */ maximumNprobes(maximumNprobes: number): VectorQuery { - super.doCall((inner) => inner.maximumNprobes(maximumNprobes)); + this.doVectorCall((inner) => inner.maximumNprobes(maximumNprobes)); return this; } @@ -567,7 +597,7 @@ export class VectorQuery extends StandardQueryBase { * `undefined` means no lower or upper bound. */ distanceRange(lowerBound?: number, upperBound?: number): VectorQuery { - super.doCall((inner) => inner.distanceRange(lowerBound, upperBound)); + this.doVectorCall((inner) => inner.distanceRange(lowerBound, upperBound)); return this; } @@ -581,7 +611,7 @@ export class VectorQuery extends StandardQueryBase { * also increase the latency of your query. The default value is 1.5*limit. */ ef(ef: number): VectorQuery { - super.doCall((inner) => inner.ef(ef)); + this.doVectorCall((inner) => inner.ef(ef)); return this; } @@ -595,7 +625,7 @@ export class VectorQuery extends StandardQueryBase { * whose data type is a fixed-size-list of floats. */ column(column: string): VectorQuery { - super.doCall((inner) => inner.column(column)); + this.doVectorCall((inner) => inner.column(column)); return this; } @@ -616,7 +646,7 @@ export class VectorQuery extends StandardQueryBase { distanceType( distanceType: Required["distanceType"], ): VectorQuery { - super.doCall((inner) => inner.distanceType(distanceType)); + this.doVectorCall((inner) => inner.distanceType(distanceType)); return this; } @@ -650,7 +680,7 @@ export class VectorQuery extends StandardQueryBase { * distance between the query vector and the actual uncompressed vector. */ refineFactor(refineFactor: number): VectorQuery { - super.doCall((inner) => inner.refineFactor(refineFactor)); + this.doVectorCall((inner) => inner.refineFactor(refineFactor)); return this; } @@ -675,7 +705,7 @@ export class VectorQuery extends StandardQueryBase { * factor can often help restore some of the results lost by post filtering. */ postfilter(): VectorQuery { - super.doCall((inner) => inner.postfilter()); + this.doVectorCall((inner) => inner.postfilter()); return this; } @@ -689,7 +719,7 @@ export class VectorQuery extends StandardQueryBase { * calculate your recall to select an appropriate value for nprobes. */ bypassVectorIndex(): VectorQuery { - super.doCall((inner) => inner.bypassVectorIndex()); + this.doVectorCall((inner) => inner.bypassVectorIndex()); return this; } @@ -697,43 +727,39 @@ export class VectorQuery extends StandardQueryBase { * Add a query vector to the search * * This method can be called multiple times to add multiple query vectors - * to the search. If multiple query vectors are added, then they will be searched - * in parallel, and the results will be concatenated. A column called `query_index` - * will be added to indicate the index of the query vector that produced the result. - * - * Performance wise, this is equivalent to running multiple queries concurrently. + * to the search. A column called `query_index` will be added to indicate the index + * of the query vector that produced the result. Flat searches share one table scan + * across the query vectors, avoiding the scan and memory amplification of running + * multiple queries concurrently. Indexed searches may still perform per-vector + * index work. */ addQueryVector(vector: IntoVector): VectorQuery { if (vector instanceof Promise) { + // Observe the promise as soon as it is accepted. The existing native + // query may still be pending, and delaying observation until it resolves + // can otherwise surface a fast rejection as unhandled. + const settledVector = vector.then( + (value) => ({ status: "fulfilled" as const, value }), + (reason) => ({ status: "rejected" as const, reason }), + ); const res = (async () => { - try { - const v = await vector; - // biome-ignore lint/suspicious/noExplicitAny: we need to get the `inner`, but js has no package scoping - const value: any = this.addQueryVector(v); - const inner = value.inner as - | NativeVectorQuery - | Promise; - return inner; - } catch (e) { - return Promise.reject(e); + const inner = await this.getInner(); + const outcome = await settledVector; + if (outcome.status === "rejected") { + throw outcome.reason; } + addQueryVectorToNative(inner, outcome.value); + return inner; })(); return new VectorQuery(res); } else { - super.doCall((inner) => { - const raw = Array.isArray(vector) ? null : extractVectorBuffer(vector); - if (raw) { - inner.addQueryVectorRaw(raw.data, raw.dtype); - } else { - inner.addQueryVector(Float32Array.from(vector as number[])); - } - }); + this.doVectorCall((inner) => addQueryVectorToNative(inner, vector)); return this; } } rerank(reranker: Reranker): VectorQuery { - super.doCall((inner) => + this.doVectorCall((inner) => inner.rerank(async (args) => { const vecResults = await fromBufferToRecordBatch(args.vecResults); const ftsResults = await fromBufferToRecordBatch(args.ftsResults); @@ -752,6 +778,71 @@ export class VectorQuery extends StandardQueryBase { } } +/** + * Create a string query whose vector/FTS routing is resolved against the active + * table schema when the query executes. + * + * @hidden + */ +export function createAutoQuery( + table: NativeTable, + query: string, + columns: string[] | null, + getVector: (metadata: string) => Promise>, +): AutoQuery { + type RouteSnapshot = { + table: NativeTable; + embeddingMetadata: string | undefined; + }; + type CachedPreparation = { + metadata: string; + vector: Promise>; + }; + + let cachedPreparation: CachedPreparation | undefined; + + const snapshotRoute = async (): Promise => { + const snapshot = await table.querySnapshot(); + const schema = tableFromIPC(await snapshot.schema()).schema; + return { + table: snapshot, + embeddingMetadata: schema.metadata.get("embedding_functions"), + }; + }; + + const createInner = async (): Promise => { + const route = await snapshotRoute(); + if (route.embeddingMetadata === undefined) { + const inner = route.table.query(); + inner.fullTextSearch({ query, columns }); + return inner; + } + + const metadata = route.embeddingMetadata; + if (cachedPreparation?.metadata !== metadata) { + cachedPreparation = { + metadata, + vector: Promise.resolve().then(() => getVector(metadata)), + }; + } + + const preparation = cachedPreparation; + let vector: Awaited; + try { + vector = await preparation.vector; + } catch (error) { + if (cachedPreparation === preparation) { + cachedPreparation = undefined; + } + throw error; + } + + return nearestToNative(route.table.query(), vector); + }; + + return new AutoQuery(createInner); +} + /** * A query that returns a subset of the rows in the table. * @@ -836,37 +927,6 @@ export class Query extends StandardQueryBase { super(tbl.query()); } - /** @hidden */ - static autoSearch( - tbl: () => Promise, - query: string, - vector: (tbl: NativeTable) => Promise | undefined>, - columns?: string[], - ): AutoQuery { - const nativeQuery = async () => { - const snapshot = await Promise.resolve(tbl()); - const resolved = await vector(snapshot); - const inner = snapshot.query(); - if (resolved === undefined) { - inner.fullTextSearch({ - query, - columns: columns ?? null, - }); - return inner; - } - - const raw = Array.isArray(resolved) - ? null - : extractVectorBuffer(resolved); - if (raw) { - return inner.nearestToRaw(raw.data, raw.dtype); - } - return inner.nearestTo(Float32Array.from(resolved as number[])); - }; - - return new AutoQuery(nativeQuery); - } - /** * Find the nearest vectors to the given query vector. * @@ -905,45 +965,19 @@ export class Query extends StandardQueryBase { * a default `limit` of 10 will be used. @see {@link Query#limit} */ nearestTo(vector: IntoVector): VectorQuery { - const callNearestTo = ( - inner: NativeQuery, - resolved: Float32Array | Float64Array | Uint8Array | number[], - ): NativeVectorQuery => { - const raw = Array.isArray(resolved) - ? null - : extractVectorBuffer(resolved); - if (raw) { - return inner.nearestToRaw(raw.data, raw.dtype); - } - return inner.nearestTo(Float32Array.from(resolved as number[])); - }; - - if (this.inner instanceof Promise) { - const nativeQuery = this.inner.then(async (inner) => { - const resolved = vector instanceof Promise ? await vector : vector; - return callNearestTo(inner, resolved); - }); + const inner = this.inner; + if (inner instanceof Promise) { + const nativeQuery = inner.then(async (resolvedInner) => + nearestToNative(resolvedInner, await vector), + ); return new VectorQuery(nativeQuery); } if (vector instanceof Promise) { - const res = (async () => { - try { - const v = await vector; - // biome-ignore lint/suspicious/noExplicitAny: we need to get the `inner`, but js has no package scoping - const value: any = this.nearestTo(v); - const inner = value.inner as - | NativeVectorQuery - | Promise; - return inner; - } catch (e) { - return Promise.reject(e); - } - })(); - return new VectorQuery(res); - } else { - const vectorQuery = callNearestTo(this.inner, vector); - return new VectorQuery(vectorQuery); + return new VectorQuery( + vector.then((resolvedVector) => nearestToNative(inner, resolvedVector)), + ); } + return new VectorQuery(nearestToNative(inner, vector)); } nearestToText(query: string | FullTextQuery, columns?: string[]): Query { diff --git a/nodejs/lancedb/table.ts b/nodejs/lancedb/table.ts index a4fc76ef1..dc062e337 100644 --- a/nodejs/lancedb/table.ts +++ b/nodejs/lancedb/table.ts @@ -48,6 +48,7 @@ import { Query, TakeQuery, VectorQuery, + createAutoQuery, instanceOfFullTextQuery, } from "./query"; import { sanitizeType } from "./sanitize"; @@ -629,6 +630,18 @@ export abstract class Table { /** * Update per-field (column) metadata. + * + * The following keys are treated specially, by convention, and should be + * used when appropriate: + * + * - `lancedb:description`: for a human-readable description of a field. + * - `lancedb:tag:`: for a user-defined key-value tag, where the suffix + * names the tag category; e.g. `lancedb:tag:model: "clip"`. + * - `lancedb:logical-column`: for a column grouping; e.g. `feature_v1` and + * `feature_v2` might be in the same logical column. + * - `lancedb:status`: for status options (`production`, `candidate`, + * `deprecated`, `archived`) to designate the current life cycle state of + * this column. * @param {FieldMetadataUpdate[]} updates One or more per-field updates. Each * update's metadata is merged into the field's existing metadata by default; * a value of `null` deletes that key, and `replace: true` swaps the whole map. @@ -1177,33 +1190,29 @@ export class LocalTable extends Table { }); } - if (queryType === "auto" && typeof query !== "string") { - return this.query().fullTextSearch(query, { - columns: ftsColumns, - }); - } + if (queryType === "auto") { + if (instanceOfFullTextQuery(query)) { + return this.query().fullTextSearch(query, { + columns: ftsColumns, + }); + } - if (queryType === "auto" && typeof query === "string") { - const vector = async (snapshot: _NativeTable) => { - const functions = await this.getEmbeddingFunctions(snapshot); + const columns = + typeof ftsColumns === "string" ? [ftsColumns] : (ftsColumns ?? null); + return createAutoQuery(this.inner, query, columns, async (metadata) => { + const functions = await getRegistry().parseFunctions( + new Map([["embedding_functions", metadata]]), + ); // TODO: Support multiple embedding functions const embeddingFunc: EmbeddingFunctionConfig | undefined = functions .values() .next().value; - if (embeddingFunc === undefined) { - return undefined; - } + // The route only calls this callback when embedding metadata exists. + // parseFunctions either yields a provider or reports malformed metadata. + if (!embeddingFunc) + throw new Error("Invalid embedding function metadata"); return await embeddingFunc.function.computeQueryEmbeddings(query); - }; - - const columns = - typeof ftsColumns === "string" ? [ftsColumns] : ftsColumns; - return Query.autoSearch( - () => this.inner.checkoutCurrent(), - query, - vector, - columns, - ); + }); } const queryPromise = this.getEmbeddingFunctions().then( @@ -1558,7 +1567,8 @@ export interface FieldMetadataUpdate { path: string; /** * Metadata key/value pairs. Merged into the field's existing metadata by - * default; a value of `null` deletes that key. + * default; a value of `null` deletes that key. See + * {@link Table.updateFieldMetadata} for the conventional `lancedb:*` keys. */ metadata: Record; /** If true, replace the field's entire metadata map instead of merging. */ diff --git a/nodejs/src/table.rs b/nodejs/src/table.rs index ff16ac042..db74d38fa 100644 --- a/nodejs/src/table.rs +++ b/nodejs/src/table.rs @@ -278,6 +278,13 @@ impl Table { Ok(Query::new(self.inner_ref()?.query())) } + /// Return a read-only table handle pinned to the current query revision. + #[napi(catch_unwind)] + pub async fn query_snapshot(&self) -> napi::Result { + let snapshot = self.inner_ref()?.query_snapshot().await.default_error()?; + Ok(Self::new(snapshot)) + } + #[napi(catch_unwind)] pub fn take_offsets(&self, offsets: Vec) -> napi::Result { Ok(TakeQuery::new( diff --git a/python/python/lancedb/functions.py b/python/python/lancedb/functions.py index 237ed73c2..8a19a9d37 100644 --- a/python/python/lancedb/functions.py +++ b/python/python/lancedb/functions.py @@ -222,6 +222,7 @@ class PythonEnvironmentSpec(_RemoteValue): kind: str packages: tuple[str, ...] = () + channels: tuple[str, ...] = () path: Optional[str] = None modules: tuple[str, ...] = () image: Optional[str] = None @@ -909,13 +910,25 @@ class UdfDefinition: pip: tuple[str, ...], env: Mapping[str, str], python_version: Optional[str], + conda: tuple[str, ...] = (), + conda_channels: tuple[str, ...] = (), ): function_name = name or function.__name__ if not _FUNCTION_NAME.fullmatch(function_name): raise ValueError(f"invalid Function name: {function_name!r}") - packages = tuple(sorted(set(pip))) + if pip and conda: + raise ValueError("a Function environment is pip or conda, not both") + if conda_channels and not conda: + raise ValueError("conda_channels requires conda packages") + packages = tuple(sorted(set(conda if conda else pip))) if any(not package or package != package.strip() for package in packages): - raise ValueError("pip requirements must be non-empty and trimmed") + raise ValueError("package requirements must be non-empty and trimmed") + if conda: + environment_spec = PythonEnvironmentSpec( + kind="conda", packages=packages, channels=tuple(conda_channels) + ) + else: + environment_spec = PythonEnvironmentSpec(kind="pip", packages=packages) environment = dict(env) if any( not isinstance(key, str) or not isinstance(value, str) @@ -929,7 +942,7 @@ class UdfDefinition: kind="python", python_version=python_version or f"{sys.version_info.major}.{sys.version_info.minor}", - environment=PythonEnvironmentSpec(kind="pip", packages=packages), + environment=environment_spec, env=environment, ) self._function = function @@ -976,6 +989,8 @@ def udf( pip: tuple[str, ...] | list[str] = (), env: Optional[Mapping[str, str]] = None, python_version: Optional[str] = None, + conda: tuple[str, ...] | list[str] = (), + conda_channels: tuple[str, ...] | list[str] = (), ) -> Callable[[Callable[..., Any]], UdfDefinition]: ... @@ -988,6 +1003,8 @@ def udf( pip: tuple[str, ...] | list[str] = (), env: Optional[Mapping[str, str]] = None, python_version: Optional[str] = None, + conda: tuple[str, ...] | list[str] = (), + conda_channels: tuple[str, ...] | list[str] = (), ): """Prepare a scalar Python callable for remote Function registration. @@ -1010,6 +1027,10 @@ def udf( provided together with ``input_schema``. pip : sequence of str, optional Pip requirements for the remote environment. + conda : sequence of str, optional + Conda packages for the remote environment, instead of ``pip``. + conda_channels : sequence of str, optional + Conda channels in priority order; requires ``conda``. env : mapping of str to str, optional Environment variables included in the Function definition. python_version : str, optional @@ -1049,6 +1070,8 @@ def udf( pip=tuple(pip), env={} if env is None else env, python_version=python_version, + conda=tuple(conda), + conda_channels=tuple(conda_channels), ) if function is None: diff --git a/python/python/lancedb/index.py b/python/python/lancedb/index.py index d2b63baf6..948342887 100644 --- a/python/python/lancedb/index.py +++ b/python/python/lancedb/index.py @@ -7,6 +7,7 @@ from typing import List, Literal, Optional from ._lancedb import ( IndexConfig, ) +from .query import DocumentGranularity from .types import BaseTokenizerType lang_mapping = { @@ -121,6 +122,11 @@ class FTS: >>> config = FTS(block_size=256) + Create an index that treats each deepest-list element as one document: + + >>> from lancedb.query import DocumentGranularity + >>> config = FTS(document_granularity=DocumentGranularity.LIST_ELEMENT) + Attributes ---------- with_position : bool, default False @@ -172,6 +178,11 @@ class FTS: roughly half of the available CPU cores. The effective value is limited by the available compute capacity. This build-only setting is not persisted with the index and does not apply to remote tables. + document_granularity : DocumentGranularity, default ROW + ``ROW`` treats the selected text in one table row as one document. + ``LIST_ELEMENT`` treats each element of the deepest list on the indexed + field path as one document and returns its physical coordinates in + ``_doc_index`` for matching queries. Notes ----- @@ -196,6 +207,7 @@ class FTS: custom_stop_words: Optional[List[str]] = None memory_limit: Optional[int] = None num_workers: Optional[int] = None + document_granularity: DocumentGranularity = DocumentGranularity.ROW @dataclass diff --git a/python/python/lancedb/query.py b/python/python/lancedb/query.py index bd5e2cea2..9301d7df8 100644 --- a/python/python/lancedb/query.py +++ b/python/python/lancedb/query.py @@ -375,6 +375,13 @@ class FullTextOperator(str, Enum): OR = "OR" +class DocumentGranularity(str, Enum): + """The unit treated as one full-text-search document.""" + + ROW = "row" + LIST_ELEMENT = "list_element" + + class Occur(str, Enum): SHOULD = "SHOULD" MUST = "MUST" @@ -478,6 +485,10 @@ class MatchQuery(FullTextQuery): prefix_length : int, optional The number of beginning characters being unchanged for fuzzy matching. This is useful to achieve prefix matching. + document_granularity : DocumentGranularity, optional + Explicitly select row or deepest-list-element documents. If omitted, + the indexed granularity is inferred. When both granularities are indexed + for the field, this must be specified. With no index, row granularity is used. """ query: str @@ -487,6 +498,9 @@ class MatchQuery(FullTextQuery): max_expansions: int = pydantic.Field(50, kw_only=True) operator: FullTextOperator = pydantic.Field(FullTextOperator.OR, kw_only=True) prefix_length: int = pydantic.Field(0, kw_only=True) + document_granularity: Optional[DocumentGranularity] = pydantic.Field( + None, kw_only=True + ) def query_type(self) -> FullTextQueryType: return FullTextQueryType.MATCH @@ -503,11 +517,20 @@ class PhraseQuery(FullTextQuery): The query string to match against. column : str The name of the column to match against. + slop : int, default 0 + The maximum number of intervening positions permitted in the phrase. + document_granularity : DocumentGranularity, optional + Explicitly select row or deepest-list-element documents. If omitted, + the indexed granularity is inferred. When both granularities are indexed + for the field, this must be specified. With no index, row granularity is used. """ query: str column: str slop: int = pydantic.Field(0, kw_only=True) + document_granularity: Optional[DocumentGranularity] = pydantic.Field( + None, kw_only=True + ) def query_type(self) -> FullTextQueryType: return FullTextQueryType.MATCH_PHRASE @@ -3378,9 +3401,10 @@ class AsyncQuery(AsyncStandardQuery): pass in multiple vectors. When multiple vectors are passed in, if the vector column is with multivector type, then the vectors will be treated as a single query. Or the vectors will be treated as multiple queries, this can be useful - if you want to find the nearest vectors to multiple query vectors. - This is not expected to be faster than making multiple queries concurrently; - it is just a convenience method. If multiple vectors are passed in then + if you want to find the nearest vectors to multiple query vectors. Flat + searches share one table scan across the query vectors, avoiding the scan + and memory amplification of making multiple queries concurrently. If + multiple vectors are passed in then an additional column `query_index` will be added to the results. This column will contain the index of the query vector that the result is nearest to. """ @@ -3509,8 +3533,8 @@ class AsyncFTSQuery(AsyncStandardQuery): Typically, a single vector is passed in as the query. However, you can also pass in multiple vectors. This can be useful if you want to find the nearest - vectors to multiple query vectors. This is not expected to be faster than - making multiple queries concurrently; it is just a convenience method. + vectors to multiple query vectors. Flat searches share one table scan across + the query vectors instead of issuing concurrent full scans. If multiple vectors are passed in then an additional column `query_index` will be added to the results. This column will contain the index of the query vector that the result is nearest to. diff --git a/python/python/lancedb/remote/table.py b/python/python/lancedb/remote/table.py index d0bf9f67a..d9139396b 100644 --- a/python/python/lancedb/remote/table.py +++ b/python/python/lancedb/remote/table.py @@ -61,6 +61,7 @@ from lancedb.table import _normalize_progress from ..query import ( AnalyzePlanDistributedMetrics, + DocumentGranularity, LanceQueryBuilder, LanceTakeQueryBuilder, LanceVectorQueryBuilder, @@ -349,6 +350,7 @@ class RemoteTable(Table): ngram_max_length: int = 3, prefix_only: bool = False, block_size: int = 128, + document_granularity: DocumentGranularity = DocumentGranularity.ROW, name: Optional[str] = None, ): """Create a full-text search index on a column. @@ -371,6 +373,7 @@ class RemoteTable(Table): ngram_max_length=ngram_max_length, prefix_only=prefix_only, block_size=block_size, + document_granularity=document_granularity, ) LOOP.run( self._table.create_index( diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index b3cab006e..c354a944e 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -85,6 +85,7 @@ from .query import ( AsyncQuery, AsyncTakeQuery, AsyncVectorQuery, + DocumentGranularity, FullTextQuery, LanceEmptyQueryBuilder, LanceFtsQueryBuilder, @@ -1168,6 +1169,7 @@ class Table(ABC): ngram_max_length: int = 3, prefix_only: bool = False, block_size: int = 128, + document_granularity: DocumentGranularity = DocumentGranularity.ROW, wait_timeout: Optional[timedelta] = None, name: Optional[str] = None, ): @@ -1246,6 +1248,11 @@ class Table(ABC): The number of documents per compressed posting block. Must be 128 or 256. A value of 256 uses the experimental FTS V3 format and may introduce breaking changes. + document_granularity: DocumentGranularity, default ROW + ``ROW`` treats the selected text in one table row as one document. + ``LIST_ELEMENT`` treats each element of the deepest list on the field + path as one document and returns its physical coordinates in + ``_doc_index`` for matching queries. wait_timeout: timedelta, optional The timeout to wait if indexing is asynchronous. name: str, optional @@ -2120,12 +2127,25 @@ class Table(ABC): ---------- updates : dict One or more dicts, each with: + - "path": str — dot-path to the field (e.g. "embedding" or "a.b.c"). - "metadata": dict[str, str | None] — keys to set; a value of ``None`` deletes that key. - "replace": bool, optional — replace the field's whole metadata map instead of merging (default False). + The following keys are treated specially, by convention, and should + be used when appropriate: + + - "lancedb:description": for a human-readable description of a field. + - ``"lancedb:tag:"`` for a user-defined key-value tag, where the + suffix names the tag category; e.g. "lancedb:tag:model": "clip". + - "lancedb:logical-column" for a column grouping; e.g. "feature_v1" + and "feature_v2" might be in the same logical column. + - "lancedb:status" for status options ("production", "candidate", + "deprecated", "archived") to designate the current life cycle + state of this column. + Returns ------- UpdateFieldMetadataResult @@ -3273,6 +3293,7 @@ class LanceTable(Table): ngram_max_length: int = 3, prefix_only: bool = False, block_size: int = 128, + document_granularity: DocumentGranularity = DocumentGranularity.ROW, name: Optional[str] = None, ): """Create a full-text search index on a column. @@ -3324,7 +3345,11 @@ class LanceTable(Table): tokenizer_configs = self.infer_tokenizer_configs(tokenizer_name) tokenizer_configs["custom_stop_words"] = custom_stop_words - config = FTS(block_size=block_size, **tokenizer_configs) + config = FTS( + block_size=block_size, + document_granularity=document_granularity, + **tokenizer_configs, + ) try: LOOP.run( diff --git a/python/python/tests/test_first_class_function_slice2.py b/python/python/tests/test_first_class_function_slice2.py index 55257d322..7ce6b6b91 100644 --- a/python/python/tests/test_first_class_function_slice2.py +++ b/python/python/tests/test_first_class_function_slice2.py @@ -69,6 +69,26 @@ def _run_packaged(definition, *args): return namespace[definition.registration_request.artifact.entrypoint](*args) +def test_udf_conda_environment(): + @udf(conda=["scipy", "numpy"], conda_channels=["conda-forge", "defaults"]) + def halve(value: float) -> float: + return value / 2 + + request = json.loads(halve.registration_request.to_canonical_json()) + assert request["runtime"]["environment"] == { + "kind": "conda", + "packages": ["numpy", "scipy"], + "channels": ["conda-forge", "defaults"], + } + pip_request = json.loads(normalize_score.registration_request.to_canonical_json()) + assert "channels" not in pip_request["runtime"]["environment"] + + with pytest.raises(ValueError, match="not both"): + udf(name="both", pip=["numpy"], conda=["numpy"])(lambda value: value) + with pytest.raises(ValueError, match="requires conda"): + udf(name="channels", conda_channels=["conda-forge"])(lambda value: value) + + def test_udf_packages_attribute_access_and_body_imports(): @udf def word_norm(body: str) -> float: diff --git a/python/python/tests/test_fts.py b/python/python/tests/test_fts.py index 625198d92..e5129dd9c 100644 --- a/python/python/tests/test_fts.py +++ b/python/python/tests/test_fts.py @@ -25,6 +25,7 @@ from lancedb.db import DBConnection from lancedb.index import FTS from lancedb.query import ( BoostQuery, + DocumentGranularity, MatchQuery, MultiMatchQuery, PhraseQuery, @@ -245,6 +246,55 @@ def test_create_inverted_index_rejects_invalid_block_size(table): table.create_index("text", config=FTS(block_size=129)) +def test_list_element_document_granularity(tmp_path): + docs_type = pa.list_(pa.struct([pa.field("content", pa.string())])) + docs = pa.array( + [ + [ + {"content": "alpha beta"}, + None, + {"content": ""}, + {"content": "the and"}, + {"content": "alpha beta"}, + ] + ], + type=docs_type, + ) + table = ldb.connect(tmp_path).create_table( + "list_element_docs", pa.table({"id": [0], "docs": docs}) + ) + row_table = ldb.connect(tmp_path).create_table( + "row_docs", pa.table({"id": [0], "docs": docs}) + ) + row_table.create_index("docs.content", config=FTS()) + row_result = row_table.search(MatchQuery("alpha", "docs.content")).to_arrow() + assert row_result.num_rows == 1 + assert "_doc_index" not in row_result.column_names + + granularity = DocumentGranularity.LIST_ELEMENT + table.create_index( + "docs.content", + config=FTS(with_position=True, document_granularity=granularity), + ) + assert table.list_indices()[0].columns == ["docs.content"] + + def coordinates(query): + result = table.search(query).limit(10).to_arrow() + doc_index_type = result.schema.field("_doc_index").type + assert pa.types.is_list(doc_index_type) + assert doc_index_type.value_type == pa.uint32() + return sorted(result["_doc_index"].to_pylist()) + + assert coordinates( + MatchQuery("alpha", "docs.content", document_granularity=granularity) + ) == [[0], [4]] + assert coordinates( + PhraseQuery("alpha beta", "docs.content", document_granularity=granularity) + ) == [[0], [4]] + assert coordinates(MatchQuery("alpha", "docs.content")) == [[0], [4]] + assert FTS().document_granularity is DocumentGranularity.ROW + + def test_create_inverted_index_respects_build_memory_limit(table): with pytest.raises(ValueError, match="exceeds worker memory limit"): table.create_index( @@ -1089,6 +1139,20 @@ def test_fts_query_to_json(): ) assert json_str == expected + # Test MatchQuery with list-element document granularity + match_query = MatchQuery( + "hello world", + "text", + document_granularity=DocumentGranularity.LIST_ELEMENT, + ) + json_str = match_query.to_json() + expected = ( + '{"match":{"column":"text","terms":"hello world","boost":1.0,' + '"fuzziness":0,"max_expansions":50,"operator":"Or","prefix_length":0,' + '"document_granularity":"list_element"}}' + ) + assert json_str == expected + # Test MatchQuery with options match_query = MatchQuery("puppy", "text", fuzziness=2, boost=1.5, prefix_length=3) json_str = match_query.to_json() @@ -1098,6 +1162,19 @@ def test_fts_query_to_json(): ) assert json_str == expected + # Test PhraseQuery with list-element document granularity + phrase_query = PhraseQuery( + "quick brown fox", + "title", + document_granularity=DocumentGranularity.LIST_ELEMENT, + ) + json_str = phrase_query.to_json() + expected = ( + '{"phrase":{"column":"title","terms":"quick brown fox","slop":0,' + '"document_granularity":"list_element"}}' + ) + assert json_str == expected + # Test PhraseQuery phrase_query = PhraseQuery("quick brown fox", "title") json_str = phrase_query.to_json() diff --git a/python/python/tests/test_query.py b/python/python/tests/test_query.py index d2629d1a8..4758f0e2d 100644 --- a/python/python/tests/test_query.py +++ b/python/python/tests/test_query.py @@ -897,6 +897,23 @@ def test_query_builder_batches(table): assert rs_list["id"][1] == 2 +def test_batch_vector_query_shares_filtered_flat_scan(table): + query = ( + table.search([[1.0, 2.0], [3.0, 4.0]]) + .where("id > 0", prefilter=True) + .limit(1) + .select(["id"]) + ) + + plan = query.explain_plan(verbose=True) + assert "KNNVectorDistance: queries=2" in plan + assert "UnionExec" not in plan + + results = query.to_arrow() + assert len(results) == 2 + assert results["query_index"].to_pylist() == [0, 1] + + def test_dynamic_projection(table): rs = ( LanceVectorQueryBuilder(table, [0, 0], "vector") diff --git a/python/python/tests/test_remote_db.py b/python/python/tests/test_remote_db.py index 2952f00a4..ab0df386d 100644 --- a/python/python/tests/test_remote_db.py +++ b/python/python/tests/test_remote_db.py @@ -1618,6 +1618,49 @@ def test_query_sync_fts(): ) +def test_query_sync_fts_document_granularity(): + from lancedb.query import DocumentGranularity, MatchQuery + + def handler(body): + assert body == { + "full_text_query": { + "query": { + "match": { + "column": "docs.content", + "terms": "alpha", + "boost": 1.0, + "fuzziness": 0, + "max_expansions": 50, + "operator": "Or", + "prefix_length": 0, + "document_granularity": "list_element", + } + } + }, + "k": 10, + "prefilter": True, + "vector": [], + "version": None, + } + return pa.table( + { + "id": [1, 1], + "_doc_index": pa.array([[0], [4]], type=pa.list_(pa.uint32())), + } + ) + + with query_test_table(handler, server_version=Version("0.6.0")) as table: + result = table.search( + MatchQuery( + "alpha", + "docs.content", + document_granularity=DocumentGranularity.LIST_ELEMENT, + ) + ).to_arrow() + + assert result["_doc_index"].to_pylist() == [[0], [4]] + + def test_query_sync_hybrid(): def handler(body): if "full_text_query" in body: diff --git a/python/python/tests/test_s3.py b/python/python/tests/test_s3.py index 256ccb1d4..70b423adb 100644 --- a/python/python/tests/test_s3.py +++ b/python/python/tests/test_s3.py @@ -4,6 +4,7 @@ import asyncio import copy +from concurrent.futures import ThreadPoolExecutor from datetime import timedelta import threading @@ -86,6 +87,25 @@ def test_s3_lifecycle(s3_bucket: str): asyncio.run(test()) +@pytest.mark.s3_test +def test_concurrent_open_table(s3_bucket: str): + uri = f"s3://{s3_bucket}/test_concurrent_open_table" + db = lancedb.connect(uri, storage_options=copy.copy(CONFIG)) + db.create_table("test", pa.table({"x": [1, 2, 3]})) + + num_workers = 32 + barrier = threading.Barrier(num_workers) + + def open_and_count(_): + barrier.wait() + return db.open_table("test").count_rows() + + with ThreadPoolExecutor(max_workers=num_workers) as pool: + row_counts = list(pool.map(open_and_count, range(num_workers))) + + assert row_counts == [3] * num_workers + + @pytest.fixture() def kms_key(): kms = get_boto3_client("kms", endpoint_url=CONFIG["aws_endpoint"]) diff --git a/python/src/index.rs b/python/src/index.rs index a5ca63c68..54b15f55e 100644 --- a/python/src/index.rs +++ b/python/src/index.rs @@ -8,7 +8,7 @@ use lancedb::index::vector::{ }; use lancedb::index::{ Index as LanceDbIndex, - scalar::{BTreeIndexBuilder, FmIndexBuilder, FtsIndexBuilder}, + scalar::{BTreeIndexBuilder, DocumentGranularity, FmIndexBuilder, FtsIndexBuilder}, }; use pyo3::IntoPyObject; use pyo3::types::PyStringMethods; @@ -60,7 +60,11 @@ pub fn extract_index_params(source: &Option>) -> PyResult, num_workers: Option, + document_granularity: String, } #[derive(FromPyObject)] @@ -481,6 +486,7 @@ mod tests { block_size = 128 memory_limit = 2048 num_workers = 7 + document_granularity = 'row' config = FTS()", None, diff --git a/python/src/query.rs b/python/src/query.rs index affbdf4fd..014e79e2d 100644 --- a/python/src/query.rs +++ b/python/src/query.rs @@ -16,8 +16,8 @@ use arrow::pyarrow::FromPyArrow; use arrow::pyarrow::IntoPyArrow; use arrow::pyarrow::ToPyArrow; use lancedb::index::scalar::{ - BooleanQuery, BoostQuery, FtsQuery, FullTextSearchQuery, MatchQuery, MultiMatchQuery, Occur, - Operator, PhraseQuery, + BooleanQuery, BoostQuery, DocumentGranularity, FtsQuery, FullTextSearchQuery, MatchQuery, + MultiMatchQuery, Occur, Operator, PhraseQuery, }; use lancedb::query::AnalyzePlanDistributedMetrics; use lancedb::query::QueryBase; @@ -76,8 +76,16 @@ impl<'a, 'py> FromPyObject<'a, 'py> for PyLanceDB { let max_expansions = ob.getattr("max_expansions")?.extract()?; let operator = ob.getattr("operator")?.extract::()?; let prefix_length = ob.getattr("prefix_length")?.extract()?; + let document_granularity = ob + .getattr("document_granularity")? + .extract::>()? + .map(|value| { + DocumentGranularity::try_from(value.as_str()) + .map_err(|err| PyValueError::new_err(err.to_string())) + }) + .transpose()?; - Ok(Self( + let mut query = MatchQuery::new(query) .with_column(Some(column)) .with_boost(boost) @@ -86,21 +94,32 @@ impl<'a, 'py> FromPyObject<'a, 'py> for PyLanceDB { .with_operator(Operator::try_from(operator.as_str()).map_err(|e| { PyValueError::new_err(format!("Invalid operator: {}", e)) })?) - .with_prefix_length(prefix_length) - .into(), - )) + .with_prefix_length(prefix_length); + if let Some(document_granularity) = document_granularity { + query = query.with_document_granularity(document_granularity); + } + Ok(Self(query.into())) } "PhraseQuery" => { let query = ob.getattr("query")?.extract()?; let column = ob.getattr("column")?.extract()?; let slop = ob.getattr("slop")?.extract()?; + let document_granularity = ob + .getattr("document_granularity")? + .extract::>()? + .map(|value| { + DocumentGranularity::try_from(value.as_str()) + .map_err(|err| PyValueError::new_err(err.to_string())) + }) + .transpose()?; - Ok(Self( - PhraseQuery::new(query) - .with_column(Some(column)) - .with_slop(slop) - .into(), - )) + let mut query = PhraseQuery::new(query) + .with_column(Some(column)) + .with_slop(slop); + if let Some(document_granularity) = document_granularity { + query = query.with_document_granularity(document_granularity); + } + Ok(Self(query.into())) } "BoostQuery" => { let positive: Self = ob.getattr("positive")?.extract()?; @@ -167,6 +186,13 @@ impl<'py> IntoPyObject<'py> for PyLanceDB { kwargs.set_item("max_expansions", query.max_expansions)?; kwargs.set_item::<_, &str>("operator", query.operator.into())?; kwargs.set_item("prefix_length", query.prefix_length)?; + if let Some(document_granularity) = query.document_granularity { + let value = match document_granularity { + DocumentGranularity::Row => "row", + DocumentGranularity::ListElement => "list_element", + }; + kwargs.set_item("document_granularity", value)?; + } namespace .getattr(intern!(py, "MatchQuery"))? .call((query.terms, query.column.unwrap()), Some(&kwargs)) @@ -174,6 +200,13 @@ impl<'py> IntoPyObject<'py> for PyLanceDB { FtsQuery::Phrase(query) => { let kwargs = PyDict::new(py); kwargs.set_item("slop", query.slop)?; + if let Some(document_granularity) = query.document_granularity { + let value = match document_granularity { + DocumentGranularity::Row => "row", + DocumentGranularity::ListElement => "list_element", + }; + kwargs.set_item("document_granularity", value)?; + } namespace .getattr(intern!(py, "PhraseQuery"))? .call((query.terms, query.column.unwrap()), Some(&kwargs)) diff --git a/rust/lancedb/src/database/listing.rs b/rust/lancedb/src/database/listing.rs index 17dc82756..064b5d28f 100644 --- a/rust/lancedb/src/database/listing.rs +++ b/rust/lancedb/src/database/listing.rs @@ -1476,7 +1476,7 @@ mod tests { use crate::table::{AnyQuery, WriteOptions}; use arrow_array::{Int32Array, RecordBatch, StringArray}; use arrow_schema::{DataType, Field, Schema, SchemaRef}; - use futures::{TryStreamExt, stream::once}; + use futures::{TryStreamExt, future::try_join_all, stream::once}; use std::path::PathBuf; use std::sync::Arc; use std::time::Duration; @@ -1614,6 +1614,59 @@ mod tests { ); } + #[tokio::test] + async fn test_concurrent_open_table_reuses_connection_object_store() { + let tempdir = tempdir().unwrap(); + let uri = tempdir.path().to_str().unwrap(); + let session = Arc::new(lance::session::Session::default()); + let request = ConnectRequest { + uri: uri.to_string(), + #[cfg(feature = "remote")] + client_config: Default::default(), + options: Default::default(), + namespace_client_properties: Default::default(), + manifest_enabled: false, + read_consistency_interval: None, + session: Some(session.clone()), + }; + let db = ListingDatabase::connect_with_options(&request) + .await + .unwrap(); + let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)])); + db.create_table(CreateTableRequest { + name: "test".to_string(), + namespace_path: vec![], + data: Box::new(RecordBatch::new_empty(schema)) as Box, + mode: CreateTableMode::Create, + write_options: Default::default(), + location: None, + namespace_client: None, + }) + .await + .unwrap(); + + let before = session.store_registry().stats(); + let opened_tables = try_join_all((0..32).map(|_| { + db.open_table(OpenTableRequest { + name: "test".to_string(), + namespace_path: vec![], + index_cache_size: None, + lance_read_params: None, + location: None, + namespace_client: None, + managed_versioning: None, + }) + })) + .await + .unwrap(); + let after = session.store_registry().stats(); + + assert_eq!(opened_tables.len(), 32); + assert_eq!(after.misses, before.misses); + assert_eq!(after.active_stores, before.active_stores); + assert!(after.hits >= before.hits + 32); + } + #[tokio::test] async fn test_listing_database_root_ops_do_not_create_manifest() { let tempdir = tempdir().unwrap(); diff --git a/rust/lancedb/src/function.rs b/rust/lancedb/src/function.rs index 52c70a4b1..5366d984e 100644 --- a/rust/lancedb/src/function.rs +++ b/rust/lancedb/src/function.rs @@ -186,6 +186,9 @@ pub struct PythonEnvironmentSpec { pub kind: String, #[serde(default, skip_serializing_if = "Vec::is_empty")] pub packages: Vec, + /// Conda channels in priority order; conda environments only. + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub channels: Vec, #[serde(default, skip_serializing_if = "Option::is_none")] pub path: Option, #[serde(default, skip_serializing_if = "Vec::is_empty")] @@ -583,3 +586,26 @@ impl RefreshColumnResult { } impl_json!(RefreshColumnResult); + +#[cfg(test)] +mod conda_environment_tests { + use super::PythonEnvironmentSpec; + + #[test] + fn conda_channels_round_trip_and_pip_stays_bare() { + let conda: PythonEnvironmentSpec = serde_json::from_str( + r#"{"kind":"conda","packages":["numpy"],"channels":["conda-forge"]}"#, + ) + .unwrap(); + assert_eq!(conda.channels, ["conda-forge"]); + assert!( + serde_json::to_string(&conda) + .unwrap() + .contains(r#""channels":["conda-forge"]"#) + ); + + let pip: PythonEnvironmentSpec = + serde_json::from_str(r#"{"kind":"pip","packages":["numpy"]}"#).unwrap(); + assert!(!serde_json::to_string(&pip).unwrap().contains("channels")); + } +} diff --git a/rust/lancedb/src/index/scalar.rs b/rust/lancedb/src/index/scalar.rs index 10d835bb1..dba05b776 100644 --- a/rust/lancedb/src/index/scalar.rs +++ b/rust/lancedb/src/index/scalar.rs @@ -63,4 +63,5 @@ pub struct FmIndexBuilder {} pub use lance_index::scalar::FullTextSearchQuery; pub use lance_index::scalar::InvertedIndexParams as FtsIndexBuilder; pub use lance_index::scalar::InvertedIndexParams; +pub use lance_index::scalar::inverted::DocumentGranularity; pub use lance_index::scalar::inverted::query::*; diff --git a/rust/lancedb/src/io/object_store.rs b/rust/lancedb/src/io/object_store.rs index d594bd857..c4a9a4f7e 100644 --- a/rust/lancedb/src/io/object_store.rs +++ b/rust/lancedb/src/io/object_store.rs @@ -10,7 +10,7 @@ use lance::io::WrappingObjectStore; use object_store::{ CopyOptions, Error, GetOptions, GetResult, ListResult, MultipartUpload, ObjectMeta, ObjectStore, ObjectStoreExt, PutMultipartOptions, PutOptions, PutPayload, PutResult, Result, - UploadPart, path::Path, + UploadPart, list::PaginatedListStore, path::Path, }; use async_trait::async_trait; @@ -187,6 +187,14 @@ impl WrappingObjectStore for MirroringObjectStoreWrapper { secondary: self.secondary.clone(), }) } + + fn wrap_paginated( + &self, + _store_prefix: &str, + original: Arc, + ) -> Option> { + Some(original) + } } // windows pathing can't be simply concatenated diff --git a/rust/lancedb/src/io/object_store/io_tracking.rs b/rust/lancedb/src/io/object_store/io_tracking.rs index bd4f8f54a..7f9750216 100644 --- a/rust/lancedb/src/io/object_store/io_tracking.rs +++ b/rust/lancedb/src/io/object_store/io_tracking.rs @@ -12,7 +12,7 @@ use lance::io::WrappingObjectStore; use object_store::{ CopyOptions, GetOptions, GetResult, ListResult, MultipartUpload, ObjectMeta, ObjectStore, PutMultipartOptions, PutOptions, PutPayload, PutResult, RenameOptions, Result as OSResult, - UploadPart, path::Path, + UploadPart, list::PaginatedListStore, path::Path, }; #[derive(Debug, Default)] @@ -57,6 +57,14 @@ impl WrappingObjectStore for IoStatsHolder { stats: self.0.clone(), }) } + + fn wrap_paginated( + &self, + _store_prefix: &str, + original: Arc, + ) -> Option> { + Some(original) + } } impl IoTrackingStore { diff --git a/rust/lancedb/src/job.rs b/rust/lancedb/src/job.rs index 94d1ba2b6..22f1a0450 100644 --- a/rust/lancedb/src/job.rs +++ b/rust/lancedb/src/job.rs @@ -47,6 +47,10 @@ impl TerminalResult { } } + pub(crate) fn value(&self) -> Option<&Value> { + self.value.as_ref() + } + fn decode(self) -> Result { let value = self.value.ok_or_else(|| match &self.request_id { Some(request_id) => Error::Http { diff --git a/rust/lancedb/src/query.rs b/rust/lancedb/src/query.rs index 654777adb..2a1283f22 100644 --- a/rust/lancedb/src/query.rs +++ b/rust/lancedb/src/query.rs @@ -1174,12 +1174,12 @@ impl VectorQuery { /// Add another query vector to the search. /// - /// Multiple searches will be dispatched as part of the query. - /// This is a convenience method for adding multiple query vectors - /// to the search. It is not expected to be faster than issuing - /// multiple queries concurrently. + /// Multiple searches will be dispatched as a batch. Flat searches share + /// one table scan across the query vectors, avoiding the scan and memory + /// amplification of issuing the searches concurrently. Indexed searches + /// may still perform per-vector index work. /// - /// The output data will contain an additional columns `query_index` which + /// The output data will contain an additional column `query_index` which /// will contain the index of the query vector that was used to generate the /// result. pub fn add_query_vector(mut self, vector: impl IntoQueryVector) -> Result { @@ -1646,7 +1646,11 @@ mod tests { use std::{collections::HashSet, sync::Arc}; use super::*; - use arrow::{array::downcast_array, compute::concat_batches, datatypes::Int32Type}; + use arrow::{ + array::downcast_array, + compute::concat_batches, + datatypes::{Int32Type, UInt8Type}, + }; use arrow_array::{ FixedSizeListArray, Float32Array, Int32Array, RecordBatch, StringArray, cast::AsArray, types::Float32Type, @@ -2334,7 +2338,8 @@ mod tests { .limit(1); let plan = query.explain_plan(true).await.unwrap(); - assert!(plan.contains("UnionExec")); + assert!(plan.contains("KNNVectorDistance: queries=2")); + assert!(!plan.contains("UnionExec")); let results = query .execute() @@ -2349,6 +2354,100 @@ mod tests { // We don't guarantee order. assert!(query_index.values().contains(&0)); assert!(query_index.values().contains(&1)); + + // Batch KNN does not support a per-query offset, so offset queries keep + // the legacy per-vector plan to preserve their result semantics. + let offset_query = table + .query() + .nearest_to(&[0.1, 0.2, 0.3, 0.4]) + .unwrap() + .add_query_vector(&[0.5, 0.6, 0.7, 0.8]) + .unwrap() + .limit(1) + .offset(1); + assert!( + offset_query + .explain_plan(true) + .await + .unwrap() + .contains("UnionExec") + ); + let offset_results = offset_query + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + assert_eq!( + offset_results + .iter() + .map(RecordBatch::num_rows) + .sum::(), + 2 + ); + } + + #[tokio::test] + async fn test_multiple_binary_query_vectors() { + let vectors = FixedSizeListArray::from_iter_primitive::( + vec![ + Some(vec![Some(0), Some(0)]), + Some(vec![Some(255), Some(255)]), + ], + 2, + ); + let schema = Arc::new(ArrowSchema::new(vec![ + ArrowField::new("id", DataType::Int32, false), + ArrowField::new("vector", vectors.data_type().clone(), false), + ])); + let batch = RecordBatch::try_new( + schema, + vec![Arc::new(Int32Array::from(vec![0, 1])), Arc::new(vectors)], + ) + .unwrap(); + + let conn = connect("memory://").execute().await.unwrap(); + let table = conn + .create_table("binary_batch", batch) + .execute() + .await + .unwrap(); + let query = table + .query() + .nearest_to(&[0.0, 0.0]) + .unwrap() + .add_query_vector(&[255.0, 255.0]) + .unwrap() + .distance_type(DistanceType::Hamming) + .limit(1); + + // Binary queries retain the per-vector plan because Lance's binary + // nearest path requires primitive UInt8 query arrays. + assert!( + query + .explain_plan(true) + .await + .unwrap() + .contains("UnionExec") + ); + + let results = query + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + let results = concat_batches(&results[0].schema(), &results).unwrap(); + assert_eq!(results.num_rows(), 2); + + let ids = results["id"].as_primitive::(); + assert!(ids.values().contains(&0)); + assert!(ids.values().contains(&1)); + let query_index = results["query_index"].as_primitive::(); + assert!(query_index.values().contains(&0)); + assert!(query_index.values().contains(&1)); } #[tokio::test] diff --git a/rust/lancedb/src/remote/db.rs b/rust/lancedb/src/remote/db.rs index 3f4216bc8..da9a4b09b 100644 --- a/rust/lancedb/src/remote/db.rs +++ b/rust/lancedb/src/remote/db.rs @@ -87,6 +87,10 @@ impl ServerVersion { pub fn support_blobs(&self) -> bool { self.0 >= semver::Version::new(0, 5, 0) } + + pub fn support_fts_document_granularity(&self) -> bool { + self.0 >= semver::Version::new(0, 6, 0) + } } pub const OPT_REMOTE_PREFIX: &str = "remote_database_"; diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index 77d5983eb..fad04a098 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -14,6 +14,7 @@ use crate::data::scannable::{PeekedScannable, Scannable, estimate_write_partitio use crate::expr::expr_to_sql_string; use crate::index::Index; use crate::index::IndexStatistics; +use crate::index::scalar::FtsQuery; use crate::index::waiter::wait_for_index; use crate::job::Job; use crate::query::{QueryFilter, QueryRequest, Select, VectorQueryRequest}; @@ -39,7 +40,8 @@ use crate::table::{ use crate::table::{AnyQuery, Filter, Predicate, PreprocessingOutput, TableStatistics}; use crate::utils::background_cache::BackgroundCache; use crate::utils::{ - resolve_arrow_field_path, supported_btree_data_type, supported_vector_data_type, + resolve_arrow_field_path, resolve_arrow_fts_field_path, supported_btree_data_type, + supported_vector_data_type, }; use crate::{DistanceType, Error}; use crate::{ @@ -86,22 +88,57 @@ const METRIC_TYPE_KEY: &str = "metric_type"; const INDEX_TYPE_KEY: &str = "index_type"; const SCHEMA_CACHE_TTL: Duration = Duration::from_secs(30); const SCHEMA_CACHE_REFRESH_WINDOW: Duration = Duration::from_secs(5); +const SCHEMA_SELECTOR_CHANGED: &str = "table selector changed while fetching schema"; + +fn fts_query_requires_document_granularity_support(query: &FtsQuery) -> bool { + match query { + FtsQuery::Match(query) => query + .document_granularity + .is_some_and(|granularity| granularity.is_list_element()), + FtsQuery::Phrase(query) => query + .document_granularity + .is_some_and(|granularity| granularity.is_list_element()), + FtsQuery::Boost(query) => { + fts_query_requires_document_granularity_support(&query.positive) + || fts_query_requires_document_granularity_support(&query.negative) + } + FtsQuery::MultiMatch(query) => query.match_queries.iter().any(|query| { + query + .document_granularity + .is_some_and(|granularity| granularity.is_list_element()) + }), + FtsQuery::Boolean(query) => query + .must + .iter() + .chain(&query.should) + .chain(&query.must_not) + .any(fts_query_requires_document_granularity_support), + } +} /// Per-table state driving the freshness headers (`x-lancedb-min-version`, -/// `x-lancedb-min-timestamp`, and `x-lancedb-min-read-version`) sent on read +/// `x-lancedb-min-timestamp`, and `x-lancedb-min-read-version`) sent on table /// requests. #[derive(Debug, Default, Clone, Copy)] struct FreshnessState { + /// Identifies the handle timeline that produced this state. Explicit + /// checkout operations advance the generation so responses from older + /// in-flight requests cannot repopulate the new timeline's constraints. + generation: u64, + /// Exact-version, tag, and snapshot handles must not carry latest-timeline + /// constraints. Their request body already selects the precise version. + pinned: bool, /// Provides read-your-write within a single handle: writes that return a /// version update this, and reads send it as `x-lancedb-min-version`. min_version: Option, - /// Highest dataset version observed in a *read* response on this handle. - /// Reads send it as `x-lancedb-min-read-version` so a load-balanced query + /// Highest committed dataset version advertised by a successful table + /// response on this handle. Later requests send it as + /// `x-lancedb-min-read-version` so a load-balanced query /// node whose cache is behind this version must refresh before serving, /// giving monotonic reads across nodes regardless of which one the load - /// balancer routes to. Sourced only from reads (always committed dataset - /// versions), never from writes (which may return WAL entry ids), so it is - /// unaffected by the WAL/version mismatch that retired `min_version`. + /// balancer routes to. Unlike write result bodies, this is sourced only + /// from the server's committed dataset-version response header or other + /// typed dataset-version fields, so WAL entry ids cannot enter it. min_read_version: Option, /// Wall-clock time captured at the last [`BaseTable::checkout_latest`] /// call. Subsequent reads send @@ -118,14 +155,22 @@ struct FreshnessState { checkout_baseline: Option, } -/// Snapshot of the headers that should be attached to a single read request. +/// Snapshot of the headers that should be attached to a single table request. #[derive(Debug, Default, Clone, Copy)] struct FreshnessHeaders { + generation: u64, min_version: Option, min_timestamp: Option, min_read_version: Option, } +#[derive(Debug, Default, Clone, Copy)] +struct ReadSnapshot { + version: Option, + freshness_state: FreshnessState, + freshness: FreshnessHeaders, +} + impl FreshnessHeaders { fn apply(self, mut request: RequestBuilder) -> RequestBuilder { if let Some(v) = self.min_version { @@ -140,6 +185,61 @@ impl FreshnessHeaders { } request } + + fn observe_version(self, freshness: &Mutex, version: u64) { + track_read_version_for_generation(freshness, self.generation, version); + } + + fn observe_headers( + self, + freshness: &Mutex, + headers: &reqwest::header::HeaderMap, + ) { + if let Some(version) = headers + .get(&VERSION_HEADER) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.parse::().ok()) + { + self.observe_version(freshness, version); + } + } + + fn update_if_current( + self, + freshness: &Mutex, + update: impl FnOnce(&mut FreshnessState), + ) { + let mut state = freshness.lock().unwrap(); + if state.generation == self.generation { + update(&mut state); + } + } + + fn is_current(self, freshness: &Mutex) -> bool { + freshness.lock().unwrap().generation == self.generation + } +} + +fn track_read_version(freshness: &Mutex, version: u64) { + if version == 0 { + return; + } + let mut state = freshness.lock().unwrap(); + state.min_read_version = Some(state.min_read_version.map_or(version, |v| v.max(version))); +} + +fn track_read_version_for_generation( + freshness: &Mutex, + generation: u64, + version: u64, +) { + if version == 0 { + return; + } + let mut state = freshness.lock().unwrap(); + if state.generation == generation { + state.min_read_version = Some(state.min_read_version.map_or(version, |v| v.max(version))); + } } /// A backfill job whose successful wait establishes a read-freshness @@ -150,6 +250,8 @@ struct FreshnessJob { inner: RemoteJob, freshness: Arc>, version: Arc>>, + track_refresh_result: bool, + freshness_request: FreshnessHeaders, } #[async_trait] @@ -166,7 +268,31 @@ impl crate::job::JobHandle for FreshnessJob { let result = crate::job::JobHandle::wait(&self.inner).await?; let version = self.version.read().await; if version.is_none() { - self.freshness.lock().unwrap().checkout_baseline = Some(SystemTime::now()); + let result_version = self + .track_refresh_result + .then(|| result.value()) + .flatten() + .and_then(|value| { + serde_json::from_value::(value.clone()) + .ok() + }) + .map(|result| { + result + .published_version + .map_or(result.source_version, |version| { + version.max(result.source_version) + }) + }) + .filter(|version| *version != 0); + if let Some(version) = result_version { + self.freshness_request + .observe_version(&self.freshness, version); + } else { + self.freshness_request + .update_if_current(&self.freshness, |state| { + state.checkout_baseline = Some(SystemTime::now()); + }); + } } Ok(result) } @@ -193,6 +319,38 @@ fn compute_min_timestamp( } } +fn freshness_headers_snapshot( + freshness: &Mutex, + interval: Option, +) -> FreshnessHeaders { + freshness_state_snapshot(freshness, interval).1 +} + +fn freshness_state_snapshot( + freshness: &Mutex, + interval: Option, +) -> (FreshnessState, FreshnessHeaders) { + let state = *freshness.lock().unwrap(); + if state.pinned { + return ( + state, + FreshnessHeaders { + generation: state.generation, + ..FreshnessHeaders::default() + }, + ); + } + ( + state, + FreshnessHeaders { + generation: state.generation, + min_version: state.min_version, + min_timestamp: compute_min_timestamp(&state, interval, SystemTime::now()), + min_read_version: state.min_read_version, + }, + ) +} + /// Normalize a branch selector: trim whitespace and treat `""` or `"main"` as /// the (absent) main branch, matching the server's convention. fn normalize_branch(branch: Option) -> Option { @@ -210,8 +368,9 @@ impl Tags for RemoteTags<'_, S> { async fn list(&self) -> Result> { let request = self .inner - .post_read(&format!("/v1/table/{}/tags/list/", self.inner.identifier)); - let (request_id, response) = self.inner.send(request, true).await?; + .client + .post(&format!("/v1/table/{}/tags/list/", self.inner.identifier)); + let (request_id, response) = self.inner.send_unfenced(request, true).await?; let response = self .inner .check_table_response(&request_id, response) @@ -241,12 +400,12 @@ impl Tags for RemoteTags<'_, S> { } async fn get_version(&self, tag: &str) -> Result { - let request = self.inner.post_read(&format!( + let request = self.inner.client.post(&format!( "/v1/table/{}/tags/version/", self.inner.identifier )); self.inner - .resolve_tag_version_with_request(tag, request) + .resolve_tag_version_with_request(tag, request, false) .await } @@ -276,7 +435,7 @@ impl Tags for RemoteTags<'_, S> { .post(&format!("/v1/table/{}/tags/delete/", self.inner.identifier)) .json(&serde_json::json!({ "tag": tag })); - let (request_id, response) = self.inner.send(request, true).await?; + let (request_id, response) = self.inner.send_unfenced(request, true).await?; self.inner .check_table_response(&request_id, response) .await?; @@ -349,8 +508,21 @@ impl RemoteTable { }); } }; + if matches!( + &index.index, + Index::FTS(params) if params.get_document_granularity().is_list_element() + ) && !self.server_version.support_fts_document_granularity() + { + return Err(Error::NotSupported { + message: "FTS document granularity requires remote server version 0.6.0 or later" + .into(), + }); + } let schema = self.schema().await?; - let (canonical_column, field) = resolve_arrow_field_path(&schema, &column)?; + let (canonical_column, field) = match &index.index { + Index::FTS(_) => resolve_arrow_fts_field_path(&schema, &column)?, + _ => resolve_arrow_field_path(&schema, &column)?, + }; let mut body = serde_json::json!({ "column": canonical_column }); @@ -386,7 +558,13 @@ impl RemoteTable { Index::Bitmap(p) => ("BITMAP", Some(to_json(p)?)), Index::LabelList(p) => ("LABEL_LIST", Some(to_json(p)?)), Index::Fm(p) => ("FM", Some(to_json(p)?)), - Index::FTS(p) => ("FTS", Some(to_json(p)?)), + Index::FTS(p) => { + let mut params = to_json(p)?; + if p.get_document_granularity().is_list_element() { + params["document_granularity"] = "list_element".into(); + } + ("FTS", Some(params)) + } Index::Auto => { if supported_vector_data_type(field.data_type()) { body[METRIC_TYPE_KEY] = @@ -468,6 +646,7 @@ impl RemoteTable { let Ok(description) = serde_json::from_str::(describe_body) else { return; }; + self.track_read_version(description.version); if let Ok(schema) = arrow_schema::Schema::try_from(description.schema) { self.schema_cache.seed(Arc::new(schema)); } @@ -510,23 +689,50 @@ impl RemoteTable { } async fn describe(&self) -> Result { - let version = self.current_version().await; - self.describe_version(version).await + self.describe_read_snapshot(self.snapshot_read_state().await) + .await } - async fn describe_version(&self, version: Option) -> Result { - let request = self.post_read(&format!("/v1/table/{}/describe/", self.identifier)); - self.describe_with_request(request, version).await + async fn describe_read_snapshot( + &self, + read_snapshot: ReadSnapshot, + ) -> Result { + let request = self + .client + .post(&format!("/v1/table/{}/describe/", self.identifier)); + self.describe_with_request( + request, + read_snapshot.version, + Some(read_snapshot.freshness), + ) + .await + } + + async fn schema_read_snapshot(&self, read_snapshot: ReadSnapshot) -> Result { + if read_snapshot.freshness.is_current(&self.freshness) + && let Some(schema) = self.schema_cache.try_get() + && read_snapshot.freshness.is_current(&self.freshness) + { + return Ok(schema); + } + + let description = self.describe_read_snapshot(read_snapshot).await?; + Ok(Arc::new(description.schema.try_into()?)) } async fn resolve_tag_version_with_request( &self, tag: &str, request: RequestBuilder, + fenced: bool, ) -> Result { let request = request.json(&serde_json::json!({ "tag": tag })); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = if fenced { + self.send(request, true).await? + } else { + self.send_unfenced(request, true).await? + }; let response = self.check_table_response(&request_id, response).await?; match response.text().await { @@ -565,7 +771,7 @@ impl RemoteTable { .client .post(&format!("/v1/table/{}/tags/version/", self.identifier)) .json(&serde_json::json!({ "tag": tag })); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = self.send_unfenced(request, true).await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; let value: serde_json::Value = serde_json::from_str(&body).map_err(|e| Error::Http { @@ -592,21 +798,34 @@ impl RemoteTable { &self, request: RequestBuilder, version: Option, + freshness_request: Option, ) -> Result { let mut body = serde_json::json!({ "version": version }); self.apply_branch_body(&mut body); let request = request.json(&body); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = if let Some(freshness_request) = freshness_request { + self.send_with_freshness(request, true, freshness_request) + .await? + } else { + self.send_unfenced(request, true).await? + }; let response = self.check_table_response(&request_id, response).await?; match response.text().await { - Ok(body) => serde_json::from_str(&body).map_err(|e| Error::Http { - source: format!("Failed to parse table description: {}", e).into(), - request_id, - status_code: None, - }), + Ok(body) => { + let description: TableDescription = + serde_json::from_str(&body).map_err(|e| Error::Http { + source: format!("Failed to parse table description: {}", e).into(), + request_id, + status_code: None, + })?; + if let Some(freshness_request) = freshness_request { + freshness_request.observe_version(&self.freshness, description.version); + } + Ok(description) + } Err(err) => { let status_code = err.status(); Err(Error::Http { @@ -619,14 +838,41 @@ impl RemoteTable { } async fn send(&self, req: RequestBuilder, with_retry: bool) -> Result<(String, Response)> { + let freshness_request = self.snapshot_freshness_headers(); + self.send_with_freshness(req, with_retry, freshness_request) + .await + } + + async fn send_with_freshness( + &self, + req: RequestBuilder, + with_retry: bool, + freshness_request: FreshnessHeaders, + ) -> Result<(String, Response)> { + let req = freshness_request.apply(req); let res = if with_retry { self.client.send_with_retry(req, None, true).await? } else { self.client.send(req).await? }; + if res.1.status().is_success() { + freshness_request.observe_headers(&self.freshness, res.1.headers()); + } Ok(res) } + async fn send_unfenced( + &self, + req: RequestBuilder, + with_retry: bool, + ) -> Result<(String, Response)> { + if with_retry { + self.client.send_with_retry(req, None, true).await + } else { + self.client.send(req).await + } + } + pub(super) async fn handle_table_not_found( table_name: &str, response: reqwest::Response, @@ -800,6 +1046,18 @@ impl RemoteTable { }); } + let requires_document_granularity_support = + fts_query_requires_document_granularity_support(&full_text_search.query); + if requires_document_granularity_support + && !self.server_version.support_fts_document_granularity() + { + return Err(Error::NotSupported { + message: + "FTS document granularity requires remote server version 0.6.0 or later" + .into(), + }); + } + if self.server_version.support_structural_fts() { body["full_text_query"] = serde_json::json!({ "query": full_text_search.query.clone(), @@ -1003,24 +1261,32 @@ impl RemoteTable { } } - async fn current_version(&self) -> Option { - let read_guard = self.version.read().await; - *read_guard + async fn snapshot_read_state(&self) -> ReadSnapshot { + let version = self.version.read().await; + let (freshness_state, freshness) = + freshness_state_snapshot(&self.freshness, self.client.read_consistency_interval); + ReadSnapshot { + version: *version, + freshness_state, + freshness, + } } - /// Snapshot the freshness headers to attach to a single read request. + /// Snapshot the freshness headers to attach to a single table request. /// Computed at call time so that retries reuse the same snapshot. fn snapshot_freshness_headers(&self) -> FreshnessHeaders { - let state = *self.freshness.lock().unwrap(); - FreshnessHeaders { - min_version: state.min_version, - min_timestamp: compute_min_timestamp( - &state, - self.client.read_consistency_interval, - SystemTime::now(), - ), - min_read_version: state.min_read_version, - } + freshness_headers_snapshot(&self.freshness, self.client.read_consistency_interval) + } + + fn reset_freshness(&self, checkout_baseline: Option, pinned: bool) { + let mut state = self.freshness.lock().unwrap(); + let generation = state.generation.wrapping_add(1); + *state = FreshnessState { + generation, + pinned, + checkout_baseline, + ..FreshnessState::default() + }; } /// Send an LSM operator request with the transport retry layer **off**. @@ -1035,46 +1301,25 @@ impl RemoteTable { Ok((request_id, response)) } - /// Build a POST request and attach the read-freshness headers - /// (`x-lancedb-min-version`, `x-lancedb-min-timestamp`). - fn post_read(&self, uri: &str) -> RequestBuilder { - self.snapshot_freshness_headers() - .apply(self.client.post(uri)) - } - /// Record a version returned by a write so subsequent reads can request at /// least that version via `x-lancedb-min-version`. A returned `0` from a /// backward-compatible old server is ignored. - fn track_write_version(&self, version: u64) { + fn track_write_version(&self, freshness_request: FreshnessHeaders, version: u64) { if version == 0 { return; } - let mut state = self.freshness.lock().unwrap(); - state.min_version = Some(state.min_version.map_or(version, |v| v.max(version))); + freshness_request.update_if_current(&self.freshness, |state| { + state.min_version = Some(state.min_version.map_or(version, |v| v.max(version))); + }); } - /// Record a dataset version observed in a *read* response so subsequent - /// reads request at least this version via `x-lancedb-min-read-version`, + /// Record a committed dataset version observed in a table response so + /// subsequent requests ask for at least this version via + /// `x-lancedb-min-read-version`, /// giving monotonic reads across load-balanced query nodes. A returned `0` /// (or absent header from an old server) is ignored. fn track_read_version(&self, version: u64) { - if version == 0 { - return; - } - let mut state = self.freshness.lock().unwrap(); - state.min_read_version = Some(state.min_read_version.map_or(version, |v| v.max(version))); - } - - /// Parse the `x-lancedb-version` response header (the dataset version a read - /// reflects) and fold it into the read-version watermark. - fn track_read_version_from_headers(&self, headers: &reqwest::header::HeaderMap) { - if let Some(version) = headers - .get(&VERSION_HEADER) - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.parse::().ok()) - { - self.track_read_version(version); - } + track_read_version(&self.freshness, version); } async fn execute_query( @@ -1082,7 +1327,9 @@ impl RemoteTable { query: &AnyQuery, options: &QueryExecutionOptions, ) -> Result>>> { - let mut request = self.post_read(&format!("/v1/table/{}/query/", self.identifier)); + let mut request = self + .client + .post(&format!("/v1/table/{}/query/", self.identifier)); if let Some(timeout) = options.timeout { // Also send to server, so it can abort the query if it takes too long. @@ -1092,15 +1339,17 @@ impl RemoteTable { } } - let query_bodies = self.prepare_query_bodies(query).await?; + let read_snapshot = self.snapshot_read_state().await; + let query_bodies = self.prepare_query_bodies(query, read_snapshot.version)?; let requests: Vec = query_bodies .into_iter() .map(|body| request.try_clone().unwrap().json(&body)) .collect(); let futures = requests.into_iter().map(|req| async move { - let (request_id, response) = self.send(req, true).await?; - self.track_read_version_from_headers(response.headers()); + let (request_id, response) = self + .send_with_freshness(req, true, read_snapshot.freshness) + .await?; self.read_arrow_response(&request_id, response).await }); let streams = futures::future::try_join_all(futures); @@ -1125,8 +1374,11 @@ impl RemoteTable { } } - async fn prepare_query_bodies(&self, query: &AnyQuery) -> Result> { - let version = self.current_version().await; + fn prepare_query_bodies( + &self, + query: &AnyQuery, + version: Option, + ) -> Result> { let mut base_body = serde_json::json!({ "version": version }); self.apply_branch_body(&mut base_body); @@ -1230,14 +1482,15 @@ async fn fetch_schema( client: &RestfulLanceDbClient, identifier: &str, table_name: &str, - version: Option, + read_snapshot: ReadSnapshot, branch: Option, - freshness_headers: FreshnessHeaders, + freshness: Arc>, ) -> Result { - let mut body = serde_json::json!({ "version": version }); + let mut body = serde_json::json!({ "version": read_snapshot.version }); if let Some(branch) = &branch { body["branch"] = serde_json::Value::String(branch.clone()); } + let freshness_headers = read_snapshot.freshness; let request = freshness_headers .apply(client.post(&format!("/v1/table/{}/describe/", identifier))) .json(&body); @@ -1257,6 +1510,7 @@ async fn fetch_schema( } let response = client.check_response(&request_id, response).await?; + freshness_headers.observe_headers(&freshness, response.headers()); let body = response.text().await.map_err(|e| { let status_code = e.status(); Error::Http { @@ -1271,6 +1525,12 @@ async fn fetch_schema( request_id, status_code: None, })?; + freshness_headers.observe_version(&freshness, description.version); + if !freshness_headers.is_current(&freshness) { + return Err(Error::Runtime { + message: SCHEMA_SELECTOR_CHANGED.to_string(), + }); + } let arrow_schema: arrow_schema::Schema = description.schema.try_into()?; Ok(Arc::new(arrow_schema)) @@ -1401,18 +1661,25 @@ impl RemoteTable { use crate::remote::retry::RetryCounter; let _guard = output.tracker.as_ref().map(|t| t.track_task()); + let freshness_request = self.snapshot_freshness_headers(); - let mut insert: Arc = Arc::new(RemoteWriteExec::new( - self.name.clone(), - self.identifier.clone(), - self.client.clone(), - output.plan, - WriteOp::Insert { - overwrite: output.overwrite, - }, - output.tracker.clone(), - self.branch.clone(), - )); + let mut insert: Arc = Arc::new( + RemoteWriteExec::new( + self.name.clone(), + self.identifier.clone(), + self.client.clone(), + output.plan, + WriteOp::Insert { + overwrite: output.overwrite, + }, + output.tracker.clone(), + self.branch.clone(), + ) + .with_freshness( + self.freshness.clone(), + self.client.read_consistency_interval, + ), + ); let mut retry_counter = RetryCounter::new(&self.client.retry_config, uuid::Uuid::new_v4().to_string()); @@ -1431,7 +1698,7 @@ impl RemoteTable { if output.overwrite { self.invalidate_schema_cache(); } - self.track_write_version(add_result.version); + self.track_write_version(freshness_request, add_result.version); return Ok(add_result); } @@ -1457,6 +1724,7 @@ impl RemoteTable { RetryCounter::new(&self.client.retry_config, uuid::Uuid::new_v4().to_string()); loop { + let freshness_request = self.snapshot_freshness_headers(); let upload_id = self.create_multipart_write().await?; let result = self @@ -1469,7 +1737,7 @@ impl RemoteTable { if output.overwrite { self.invalidate_schema_cache(); } - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); return Ok(result); } Err(e) => { @@ -1525,18 +1793,24 @@ impl RemoteTable { )?, ) as Arc; - let insert = Arc::new(RemoteWriteExec::new_multipart( - self.name.clone(), - self.identifier.clone(), - self.client.clone(), - plan, - output.overwrite, - upload_id.to_string(), - output.tracker.clone(), - self.branch.clone(), - self.client.max_bytes_per_request(), - self.client.max_request_duration(), - )); + let insert = Arc::new( + RemoteWriteExec::new_multipart( + self.name.clone(), + self.identifier.clone(), + self.client.clone(), + plan, + output.overwrite, + upload_id.to_string(), + output.tracker.clone(), + self.branch.clone(), + self.client.max_bytes_per_request(), + self.client.max_request_duration(), + ) + .with_freshness( + self.freshness.clone(), + self.client.read_consistency_interval, + ), + ); let task_ctx = Arc::new(datafusion_execution::TaskContext::default()); let tracker = output.tracker.clone(); @@ -1599,6 +1873,38 @@ where } impl RemoteTable { + async fn index_stats_read_snapshot( + &self, + index_name: &str, + read_snapshot: ReadSnapshot, + ) -> Result> { + let encoded_name = urlencoding::encode(index_name); + let mut body = serde_json::json!({ "version": read_snapshot.version }); + self.apply_branch_body(&mut body); + let request = self + .client + .post(&format!( + "/v1/table/{}/index/{encoded_name}/stats/", + self.identifier + )) + .json(&body); + + let (request_id, response) = self + .send_with_freshness(request, true, read_snapshot.freshness) + .await?; + if response.status() == StatusCode::NOT_FOUND { + return Ok(None); + } + let response = self.check_table_response(&request_id, response).await?; + let body = response.text().await.err_to_http(request_id.clone())?; + let stats = serde_json::from_str(&body).map_err(|e| Error::Http { + source: format!("Failed to parse index statistics: {}", e).into(), + request_id, + status_code: None, + })?; + Ok(Some(stats)) + } + /// Parse the response from `/index/list/` into `IndexConfig` entries. /// /// When the server returns `index_type` inline, all enriched fields are @@ -1610,6 +1916,7 @@ impl RemoteTable { body: &str, request_id: &str, schema: &SchemaRef, + read_snapshot: ReadSnapshot, ) -> Result> { use crate::index::IndexType; @@ -1678,7 +1985,10 @@ impl RemoteTable { })) } else { // Legacy response: fetch index type via stats endpoint. - match self.index_stats(&entry.index_name).await { + match self + .index_stats_read_snapshot(&entry.index_name, read_snapshot) + .await + { Ok(Some(stats)) => Ok(Some(IndexConfig { name: entry.index_name, index_type: stats.index_type, @@ -1722,6 +2032,20 @@ impl BaseTable for RemoteTable { fn id(&self) -> &str { &self.identifier } + async fn query_snapshot(&self) -> Result> { + let description = self.describe().await?; + let TableDescription { + version, + schema, + location, + } = description; + let schema = Arc::new(arrow_schema::Schema::try_from(schema)?); + let snapshot = self.with_branch(self.branch.clone()); + *snapshot.version.write().await = Some(version); + *snapshot.location.write().await = location; + snapshot.schema_cache.seed(schema); + Ok(Arc::new(snapshot)) + } async fn version(&self) -> Result { self.describe().await.map(|desc| desc.version) } @@ -1748,7 +2072,7 @@ impl BaseTable for RemoteTable { let request = self .client .post(&format!("/v1/table/{}/describe/", self.identifier)); - self.describe_with_request(request, Some(version)) + self.describe_with_request(request, Some(version), None) .await .map_err(|e| match e { // try to map the error to a more user-friendly error telling them @@ -1761,58 +2085,52 @@ impl BaseTable for RemoteTable { })?; let mut write_guard = self.version.write().await; + // Commit the selector and its freshness mode while holding the selector + // write lock, with no cancellation point between the two updates. + self.reset_freshness(None, true); *write_guard = Some(version); - drop(write_guard); - - // Explicit time-travel: drop any read-your-write / freshness - // constraints so the user sees exactly the requested version. - *self.freshness.lock().unwrap() = FreshnessState::default(); - - // Invalidate schema cache since we're switching versions self.invalidate_schema_cache(); + drop(write_guard); Ok(()) } async fn checkout_latest(&self) -> Result<()> { let mut write_guard = self.version.write().await; - *write_guard = None; - drop(write_guard); - // Drop any per-handle read/write tracking; subsequent reads use the // baseline timestamp captured now to guarantee freshness. - *self.freshness.lock().unwrap() = FreshnessState { - min_version: None, - checkout_baseline: Some(SystemTime::now()), - min_read_version: None, - }; - - // Invalidate schema cache since we're switching versions + self.reset_freshness(Some(SystemTime::now()), false); + *write_guard = None; self.invalidate_schema_cache(); + drop(write_guard); Ok(()) } async fn snapshot_at_current_version(&self) -> Result>> { // A checked-out handle already names its snapshot. Otherwise resolve // latest exactly once before creating the independent pinned handle. - let version = match self.current_version().await { + let read_snapshot = self.snapshot_read_state().await; + let version = match read_snapshot.version { Some(version) => version, - None => self.describe().await?.version, + None => self.describe_read_snapshot(read_snapshot).await?.version, }; let snapshot = self.with_branch(self.branch.clone()); *snapshot.version.write().await = Some(version); + snapshot.reset_freshness(None, true); Ok(Some(Arc::new(snapshot))) } async fn restore(&self) -> Result<()> { let mut request = self .client .post(&format!("/v1/table/{}/restore/", self.identifier)); - let version = self.current_version().await; - let mut body = serde_json::json!({ "version": version }); + let read_snapshot = self.snapshot_read_state().await; + let mut body = serde_json::json!({ "version": read_snapshot.version }); self.apply_branch_body(&mut body); request = request.json(&body); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = self + .send_with_freshness(request, true, read_snapshot.freshness) + .await?; self.check_table_response(&request_id, response).await?; self.checkout_latest().await?; Ok(()) @@ -1820,7 +2138,8 @@ impl BaseTable for RemoteTable { async fn list_versions(&self) -> Result> { let request = self.apply_branch_query( - self.post_read(&format!("/v1/table/{}/version/list/", self.identifier)), + self.client + .post(&format!("/v1/table/{}/version/list/", self.identifier)), ); let (request_id, response) = self.send(request, true).await?; let response = self.check_table_response(&request_id, response).await?; @@ -1884,31 +2203,42 @@ impl BaseTable for RemoteTable { } async fn schema(&self) -> Result { - if let Some(schema) = self.schema_cache.try_get() { - return Ok(schema); - } + loop { + let read_snapshot = self.snapshot_read_state().await; + if let Some(schema) = self.schema_cache.try_get() { + return Ok(schema); + } - let version = self.current_version().await; - let client = self.client.clone(); - let identifier = self.identifier.clone(); - let table_name = self.name.clone(); - let branch = self.branch.clone(); - let freshness_headers = self.snapshot_freshness_headers(); + let client = self.client.clone(); + let identifier = self.identifier.clone(); + let table_name = self.name.clone(); + let branch = self.branch.clone(); + let freshness = self.freshness.clone(); - self.schema_cache - .get(move || async move { - fetch_schema( - &client, - &identifier, - &table_name, - version, - branch, - freshness_headers, - ) + match self + .schema_cache + .get(move || async move { + fetch_schema( + &client, + &identifier, + &table_name, + read_snapshot, + branch, + freshness, + ) + .await + }) .await - }) - .await - .map_err(unwrap_shared_error) + { + Ok(schema) => return Ok(schema), + Err(error) + if matches!( + &*error, + Error::Runtime { message } if message == SCHEMA_SELECTOR_CHANGED + ) => {} + Err(error) => return Err(unwrap_shared_error(error)), + } + } } async fn create_branch( @@ -1950,7 +2280,7 @@ impl BaseTable for RemoteTable { // Send without retry so the expected 409 (branch already exists) is // surfaced as a response we can map, rather than being retried. - let (request_id, response) = self.send(request, false).await?; + let (request_id, response) = self.send_unfenced(request, false).await?; match response.status() { StatusCode::CONFLICT => { return Err(Error::TableAlreadyExists { @@ -2011,8 +2341,10 @@ impl BaseTable for RemoteTable { async fn list_branches(&self) -> Result> { use lance::dataset::refs::BranchContents; - let request = self.post_read(&format!("/v1/table/{}/branches/list/", self.identifier)); - let (request_id, response) = self.send(request, true).await?; + let request = self + .client + .post(&format!("/v1/table/{}/branches/list/", self.identifier)); + let (request_id, response) = self.send_unfenced(request, true).await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2066,7 +2398,7 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/branches/diff/", self.identifier)) .json(&serde_json::json!({ "from_branch": from_branch })); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = self.send_unfenced(request, true).await?; if response.status() == StatusCode::NOT_FOUND { return Err(Error::TableNotFound { name: format!("{} (branch: {})", self.name, from_branch), @@ -2092,6 +2424,12 @@ impl BaseTable for RemoteTable { message: "Branch name cannot be empty.".into(), }); } + let read_snapshot = self.snapshot_read_state().await; + let target_freshness = if self.branch.is_none() && read_snapshot.version.is_none() { + Some(read_snapshot.freshness) + } else { + None + }; let request = self .client .post(&format!( @@ -2103,7 +2441,7 @@ impl BaseTable for RemoteTable { "dry_run": dry_run, })); // No retry. HTTP 409 is CherryPickStatus::Failed with a body, not a transport error. - let (request_id, response) = self.send(request, false).await?; + let (request_id, response) = self.send_unfenced(request, false).await?; let status = response.status(); if status == StatusCode::NOT_FOUND { return Err(Error::TableNotFound { @@ -2121,7 +2459,7 @@ impl BaseTable for RemoteTable { }); } let body = response.text().await.err_to_http(request_id.clone())?; - serde_json::from_str(&body).map_err(|err| Error::Http { + let result: CherryPickResult = serde_json::from_str(&body).map_err(|err| Error::Http { source: format!( "Failed to parse cherry_pick response: {}, body: {}", err, body @@ -2129,7 +2467,15 @@ impl BaseTable for RemoteTable { .into(), request_id, status_code: Some(status), - }) + })?; + if !dry_run + && status == StatusCode::OK + && result.status == crate::table::CherryPickStatus::CherryPicked + && let (Some(freshness), Some(version)) = (target_freshness, result.main_version_after) + { + freshness.observe_version(&self.freshness, version); + } + Ok(result) } fn current_branch(&self) -> Option { @@ -2137,23 +2483,28 @@ impl BaseTable for RemoteTable { } async fn count_rows(&self, filter: Option) -> Result { - let mut request = self.post_read(&format!("/v1/table/{}/count_rows/", self.identifier)); + let mut request = self + .client + .post(&format!("/v1/table/{}/count_rows/", self.identifier)); - let version = self.current_version().await; + let read_snapshot = self.snapshot_read_state().await; let mut body = if let Some(filter) = filter { let filter_sql = match filter { Filter::Sql(sql) => sql.clone(), Filter::Datafusion(expr) => expr_to_sql_string(&expr)?, }; - serde_json::json!({ "predicate": filter_sql, "version": version }) + serde_json::json!({ "predicate": filter_sql, "version": read_snapshot.version }) } else { - serde_json::json!({ "version": version }) + serde_json::json!({ "version": read_snapshot.version }) }; self.apply_branch_body(&mut body); request = request.json(&body); - let (request_id, response) = match self.send(request, true).await { + let (request_id, response) = match self + .send_with_freshness(request, true, read_snapshot.freshness) + .await + { Ok((id, resp)) => { // check_table_response now handles error-based invalidation let response = self.check_table_response(&id, resp).await?; @@ -2165,7 +2516,6 @@ impl BaseTable for RemoteTable { } }; - self.track_read_version_from_headers(response.headers()); let body = response.text().await.err_to_http(request_id.clone())?; serde_json::from_str(&body).map_err(|e| Error::Http { @@ -2298,9 +2648,12 @@ impl BaseTable for RemoteTable { } async fn explain_plan(&self, query: &AnyQuery, verbose: bool) -> Result { - let base_request = self.post_read(&format!("/v1/table/{}/explain_plan/", self.identifier)); + let base_request = self + .client + .post(&format!("/v1/table/{}/explain_plan/", self.identifier)); - let query_bodies = self.prepare_query_bodies(query).await?; + let read_snapshot = self.snapshot_read_state().await; + let query_bodies = self.prepare_query_bodies(query, read_snapshot.version)?; let requests: Vec = query_bodies .into_iter() .map(|query_body| { @@ -2314,7 +2667,9 @@ impl BaseTable for RemoteTable { .collect::>(); let futures = requests.into_iter().map(|req| async move { - let (request_id, response) = self.send(req, true).await?; + let (request_id, response) = self + .send_with_freshness(req, true, read_snapshot.freshness) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2345,7 +2700,9 @@ impl BaseTable for RemoteTable { query: &AnyQuery, options: QueryExecutionOptions, ) -> Result { - let mut request = self.post_read(&format!("/v1/table/{}/analyze_plan/", self.identifier)); + let mut request = self + .client + .post(&format!("/v1/table/{}/analyze_plan/", self.identifier)); if options.analyze_plan_distributed_metrics != AnalyzePlanDistributedMetrics::Aggregate { request = request.query(&[( @@ -2354,14 +2711,17 @@ impl BaseTable for RemoteTable { )]); } - let query_bodies = self.prepare_query_bodies(query).await?; + let read_snapshot = self.snapshot_read_state().await; + let query_bodies = self.prepare_query_bodies(query, read_snapshot.version)?; let requests: Vec = query_bodies .into_iter() .map(|body| request.try_clone().unwrap().json(&body)) .collect(); let futures = requests.into_iter().map(|req| async move { - let (request_id, response) = self.send(req, true).await?; + let (request_id, response) = self + .send_with_freshness(req, true, read_snapshot.freshness) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2405,7 +2765,10 @@ impl BaseTable for RemoteTable { self.apply_branch_body(&mut body); let request = request.json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2424,7 +2787,7 @@ impl BaseTable for RemoteTable { status_code: None, })?; - self.track_write_version(update_response.version); + self.track_write_version(freshness_request, update_response.version); Ok(update_response) } @@ -2440,7 +2803,10 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/delete/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; if body.trim().is_empty() { @@ -2456,7 +2822,7 @@ impl BaseTable for RemoteTable { request_id, status_code: None, })?; - self.track_write_version(delete_response.version); + self.track_write_version(freshness_request, delete_response.version); Ok(delete_response) } @@ -2466,7 +2832,13 @@ impl BaseTable for RemoteTable { async fn create_index_async(&self, index: IndexBuilder) -> Result { Ok(match self.submit_create_index(index).await? { - Some(job_id) => Job::new(Box::new(RemoteJob::new(self.client.clone(), job_id))), + Some(job_id) => Job::new(Box::new(FreshnessJob { + inner: RemoteJob::new(self.client.clone(), job_id), + freshness: self.freshness.clone(), + version: self.version.clone(), + track_refresh_result: false, + freshness_request: self.snapshot_freshness_headers(), + })), None => Job::new_done(), }) } @@ -2506,16 +2878,23 @@ impl BaseTable for RemoteTable { let rescannable = source.rescannable(); let input: Arc = Arc::new(crate::table::datafusion::scannable_exec::ScannableExec::new(source, None)); + let freshness_request = self.snapshot_freshness_headers(); - let mut merge: Arc = Arc::new(RemoteWriteExec::new( - self.name.clone(), - self.identifier.clone(), - self.client.clone(), - input, - WriteOp::MergeInsert { query, timeout }, - None, - self.branch.clone(), - )); + let mut merge: Arc = Arc::new( + RemoteWriteExec::new( + self.name.clone(), + self.identifier.clone(), + self.client.clone(), + input, + WriteOp::MergeInsert { query, timeout }, + None, + self.branch.clone(), + ) + .with_freshness( + self.freshness.clone(), + self.client.read_consistency_interval, + ), + ); let mut retry_counter = crate::remote::retry::RetryCounter::new( &self.client.retry_config, @@ -2533,7 +2912,7 @@ impl BaseTable for RemoteTable { .and_then(|m| m.merge_result()) .unwrap_or_default(); - self.track_write_version(merge_result.version); + self.track_write_version(freshness_request, merge_result.version); return Ok(merge_result); } Err(err) if rescannable && self.is_retryable_write_error(&err) => { @@ -2572,7 +2951,8 @@ impl BaseTable for RemoteTable { async fn get_lsm_stats(&self, include_generation_rows: bool) -> Result> { // Read-semantics POST, like `get_lsm_write_spec`. let request = self - .post_read(&format!("/v1/table/{}/get_lsm_stats/", self.identifier)) + .client + .post(&format!("/v1/table/{}/get_lsm_stats/", self.identifier)) .json(&serde_json::json!({ "include_generation_rows": include_generation_rows, })); @@ -2646,7 +3026,7 @@ impl BaseTable for RemoteTable { // re-encodes it into the same sophon-owned shape the set endpoint // accepts — no lance/lancedb types cross the wire. `lsm_write_spec` is // null when the LSM write path is not enabled for the table. - let request = self.post_read(&format!( + let request = self.client.post(&format!( "/v1/table/{}/get_lsm_write_spec/", self.identifier )); @@ -2712,18 +3092,17 @@ impl BaseTable for RemoteTable { let request = self .client .post(&format!("/v1/table/{}/tags/version/", self.identifier)); - let version = self.resolve_tag_version_with_request(tag, request).await?; + let version = self + .resolve_tag_version_with_request(tag, request, false) + .await?; let mut write_guard = self.version.write().await; + // Commit the selector and its freshness mode while holding the selector + // write lock, with no cancellation point between the two updates. + self.reset_freshness(None, true); *write_guard = Some(version); - drop(write_guard); - - // Explicit time-travel: drop any read-your-write / freshness - // constraints so the user sees exactly the tagged version. - *self.freshness.lock().unwrap() = FreshnessState::default(); - - // Invalidate schema cache since we're switching versions self.invalidate_schema_cache(); + drop(write_guard); Ok(()) } @@ -2760,7 +3139,10 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/add_columns/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2777,7 +3159,7 @@ impl BaseTable for RemoteTable { })?; self.invalidate_schema_cache(); - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); Ok(result) } @@ -2813,7 +3195,10 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/add_columns/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2829,7 +3214,7 @@ impl BaseTable for RemoteTable { })?; self.invalidate_schema_cache(); - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); Ok(result) } @@ -2872,7 +3257,10 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/add_columns/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2887,7 +3275,7 @@ impl BaseTable for RemoteTable { })?; self.invalidate_schema_cache(); - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); Ok(result) } @@ -2909,7 +3297,8 @@ impl BaseTable for RemoteTable { let mut body = serde_json::json!({ "column": column }); self.apply_branch_body(&mut body); let request = self - .post_read(&format!("/v1/table/{}/backfill_column", self.identifier)) + .client + .post(&format!("/v1/table/{}/backfill_column", self.identifier)) .json(&body); let (request_id, response) = self.send(request, true).await?; let response = self.check_table_response(&request_id, response).await?; @@ -2929,6 +3318,8 @@ impl BaseTable for RemoteTable { inner: RemoteJob::new(self.client.clone(), response.job_id), freshness: self.freshness.clone(), version: self.version.clone(), + track_refresh_result: true, + freshness_request: self.snapshot_freshness_headers(), }))) } @@ -2960,7 +3351,10 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/alter_columns/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2976,7 +3370,7 @@ impl BaseTable for RemoteTable { })?; self.invalidate_schema_cache(); - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); Ok(result) } @@ -2995,7 +3389,10 @@ impl BaseTable for RemoteTable { self.identifier )) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -3007,7 +3404,7 @@ impl BaseTable for RemoteTable { })?; self.invalidate_schema_cache(); - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); Ok(result) } @@ -3019,7 +3416,10 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/drop_columns/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -3035,55 +3435,34 @@ impl BaseTable for RemoteTable { })?; self.invalidate_schema_cache(); - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); Ok(result) } async fn list_indices(&self) -> Result> { - let mut request = self.post_read(&format!("/v1/table/{}/index/list/", self.identifier)); - let version = self.current_version().await; - let mut body = serde_json::json!({ "version": version }); + let mut request = self + .client + .post(&format!("/v1/table/{}/index/list/", self.identifier)); + let read_snapshot = self.snapshot_read_state().await; + let mut body = serde_json::json!({ "version": read_snapshot.version }); self.apply_branch_body(&mut body); request = request.json(&body); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = self + .send_with_freshness(request, true, read_snapshot.freshness) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; - let schema = self.schema().await?; + let schema = self.schema_read_snapshot(read_snapshot).await?; - self.parse_index_list_response(&body, &request_id, &schema) + self.parse_index_list_response(&body, &request_id, &schema, read_snapshot) .await } async fn index_stats(&self, index_name: &str) -> Result> { - let encoded_name = urlencoding::encode(index_name); - let mut request = self.post_read(&format!( - "/v1/table/{}/index/{encoded_name}/stats/", - self.identifier - )); - let version = self.current_version().await; - let mut body = serde_json::json!({ "version": version }); - self.apply_branch_body(&mut body); - request = request.json(&body); - - let (request_id, response) = self.send(request, true).await?; - - if response.status() == StatusCode::NOT_FOUND { - return Ok(None); - } - - let response = self.check_table_response(&request_id, response).await?; - - let body = response.text().await.err_to_http(request_id.clone())?; - - let stats = serde_json::from_str(&body).map_err(|e| Error::Http { - source: format!("Failed to parse index statistics: {}", e).into(), - request_id, - status_code: None, - })?; - - Ok(Some(stats)) + self.index_stats_read_snapshot(index_name, self.snapshot_read_state().await) + .await } async fn drop_index(&self, index_name: &str) -> Result<()> { @@ -3173,7 +3552,9 @@ impl BaseTable for RemoteTable { } async fn stats(&self) -> Result { - let mut request = self.post_read(&format!("/v1/table/{}/stats/", self.identifier)); + let mut request = self + .client + .post(&format!("/v1/table/{}/stats/", self.identifier)); if let Some(branch) = &self.branch { request = request.json(&serde_json::json!({ "branch": branch })); } @@ -3195,15 +3576,21 @@ impl BaseTable for RemoteTable { write_params: lance::dataset::WriteParams, ) -> Result> { let overwrite = matches!(write_params.mode, lance::dataset::WriteMode::Overwrite); - Ok(Arc::new(insert::RemoteWriteExec::new( - self.name.clone(), - self.identifier.clone(), - self.client.clone(), - input, - WriteOp::Insert { overwrite }, - None, - self.branch.clone(), - ))) + Ok(Arc::new( + insert::RemoteWriteExec::new( + self.name.clone(), + self.identifier.clone(), + self.client.clone(), + input, + WriteOp::Insert { overwrite }, + None, + self.branch.clone(), + ) + .with_freshness( + self.freshness.clone(), + self.client.read_consistency_interval, + ), + )) } } @@ -3301,7 +3688,7 @@ mod tests { use arrow_schema::{DataType, Field, Schema}; use chrono::{DateTime, Utc}; use futures::{StreamExt, TryFutureExt, future::BoxFuture}; - use lance_index::scalar::inverted::query::MatchQuery; + use lance_index::scalar::inverted::{DocumentGranularity, query::MatchQuery}; use lance_index::scalar::{FullTextSearchQuery, InvertedIndexParams}; use reqwest::Body; use rstest::rstest; @@ -3578,6 +3965,15 @@ mod tests { DataType::Struct(vec![Field::new("text", DataType::Utf8, false)].into()), false, ), + Field::new( + "docs", + DataType::List(Arc::new(Field::new( + "item", + DataType::Struct(vec![Field::new("content", DataType::Utf8, true)].into()), + true, + ))), + true, + ), Field::new( "meta-data", DataType::Struct(vec![Field::new("user-id", DataType::Int32, false)].into()), @@ -5206,7 +5602,7 @@ mod tests { #[tokio::test] async fn test_query_structured_fts() { let table = - Table::new_with_handler_version("my_table", semver::Version::new(0, 3, 0), |request| { + Table::new_with_handler_version("my_table", semver::Version::new(0, 6, 0), |request| { assert_eq!(request.method(), "POST"); assert_eq!(request.url().path(), "/v1/table/my_table/query/"); assert_eq!( @@ -5227,6 +5623,7 @@ mod tests { "max_expansions": 50, "operator": "Or", "prefix_length": 0, + "document_granularity": "list_element", }, } }, @@ -5256,6 +5653,7 @@ mod tests { .full_text_search(FullTextSearchQuery::new_query( MatchQuery::new("hello world".to_owned()) .with_column(Some("payload.text".to_owned())) + .with_document_granularity(DocumentGranularity::ListElement) .into(), )) .with_row_id() @@ -5265,6 +5663,76 @@ mod tests { .unwrap(); } + #[tokio::test] + async fn test_query_row_document_granularity_uses_structured_fts() { + let table = + Table::new_with_handler_version("my_table", semver::Version::new(0, 3, 0), |request| { + let body = request.body().unwrap().as_bytes().unwrap(); + let body: serde_json::Value = serde_json::from_slice(body).unwrap(); + assert_eq!( + body["full_text_query"]["query"]["match"]["document_granularity"], + "row" + ); + + let data = RecordBatch::try_new( + Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])), + vec![Arc::new(Int32Array::from(vec![1]))], + ) + .unwrap(); + http::Response::builder() + .status(200) + .header(CONTENT_TYPE, ARROW_FILE_CONTENT_TYPE) + .body(write_ipc_file(&data)) + .unwrap() + }); + + table + .query() + .full_text_search(FullTextSearchQuery::new_query( + MatchQuery::new("hello world".to_owned()) + .with_column(Some("payload.text".to_owned())) + .with_document_granularity(DocumentGranularity::Row) + .into(), + )) + .execute() + .await + .unwrap(); + } + + #[rstest] + #[case(DEFAULT_SERVER_VERSION.clone())] + #[case(semver::Version::new(0, 3, 0))] + #[case(semver::Version::new(0, 5, 0))] + #[tokio::test] + async fn test_query_document_granularity_requires_server_support( + #[case] version: semver::Version, + ) { + let table = + Table::new_with_handler_version("my_table", version, |_| -> http::Response { + panic!("unsupported remote query must fail before sending a request") + }); + + let result = table + .query() + .full_text_search(FullTextSearchQuery::new_query( + MatchQuery::new("hello world".to_owned()) + .with_column(Some("payload.text".to_owned())) + .with_document_granularity(DocumentGranularity::ListElement) + .into(), + )) + .execute() + .await; + let err = match result { + Ok(_) => panic!("legacy remote query unexpectedly succeeded"), + Err(err) => err, + }; + + assert!( + err.to_string() + .contains("document granularity requires remote server version 0.6.0 or later") + ); + } + #[rstest] #[case(DEFAULT_SERVER_VERSION.clone())] #[case(semver::Version::new(0, 2, 0))] @@ -5543,40 +6011,56 @@ mod tests { "CAT".to_string(), ]))), ), + ( + "FTS", + { + let mut body = serde_json::to_value(InvertedIndexParams::default()).unwrap(); + body["document_granularity"] = "list_element".into(); + body + }, + Index::FTS( + InvertedIndexParams::default() + .document_granularity(DocumentGranularity::ListElement), + ), + ), ]; for (index_type, expected_body, index) in cases { - let table = Table::new_with_handler("my_table", move |request| { - assert_eq!(request.method(), "POST"); - match request.url().path() { - "/v1/table/my_table/describe/" => { - let schema = Schema::new(vec![Field::new("a", DataType::Int32, false)]); - http::Response::builder() - .status(200) - .body(describe_response(&schema)) - .unwrap() - } - "/v1/table/my_table/create_index/" => { - assert_eq!( - request.headers().get("Content-Type").unwrap(), - JSON_CONTENT_TYPE - ); - let body = request.body().unwrap().as_bytes().unwrap(); - let body: serde_json::Value = serde_json::from_slice(body).unwrap(); - let mut expected_body = expected_body.clone(); - expected_body["column"] = "a".into(); - expected_body[INDEX_TYPE_KEY] = index_type.into(); + let table = Table::new_with_handler_version( + "my_table", + semver::Version::new(0, 6, 0), + move |request| { + assert_eq!(request.method(), "POST"); + match request.url().path() { + "/v1/table/my_table/describe/" => { + let schema = Schema::new(vec![Field::new("a", DataType::Int32, false)]); + http::Response::builder() + .status(200) + .body(describe_response(&schema)) + .unwrap() + } + "/v1/table/my_table/create_index/" => { + assert_eq!( + request.headers().get("Content-Type").unwrap(), + JSON_CONTENT_TYPE + ); + let body = request.body().unwrap().as_bytes().unwrap(); + let body: serde_json::Value = serde_json::from_slice(body).unwrap(); + let mut expected_body = expected_body.clone(); + expected_body["column"] = "a".into(); + expected_body[INDEX_TYPE_KEY] = index_type.into(); - assert_eq!(body, expected_body); + assert_eq!(body, expected_body); - http::Response::builder() - .status(200) - .body("{}".to_string()) - .unwrap() + http::Response::builder() + .status(200) + .body("{}".to_string()) + .unwrap() + } + path => panic!("Unexpected path: {}", path), } - path => panic!("Unexpected path: {}", path), - } - }); + }, + ); table.create_index(&["a"], index).execute().await.unwrap(); } @@ -5835,6 +6319,19 @@ mod tests { body["index_type"] = "FTS".into(); body }, + { + let mut body = serde_json::to_value(InvertedIndexParams::default()).unwrap(); + body["column"] = "docs.content".into(); + body["index_type"] = "FTS".into(); + body + }, + { + let mut body = serde_json::to_value(InvertedIndexParams::default()).unwrap(); + body["column"] = "docs.content".into(); + body["index_type"] = "FTS".into(); + body["document_granularity"] = "list_element".into(); + body + }, json!({ "column": "`meta-data`.`user-id`", "index_type": "BTREE", @@ -5845,7 +6342,7 @@ mod tests { }), ]); let request_idx = Arc::new(AtomicUsize::new(0)); - let table = Table::new_with_handler("my_table", { + let table = Table::new_with_handler_version("my_table", semver::Version::new(0, 6, 0), { let schema = schema.clone(); let expected_requests = expected_requests.clone(); let request_idx = request_idx.clone(); @@ -5910,6 +6407,22 @@ mod tests { .execute() .await .unwrap(); + table + .create_index(&["Docs.Content"], Index::FTS(Default::default())) + .execute() + .await + .unwrap(); + table + .create_index( + &["Docs.Content"], + Index::FTS( + InvertedIndexParams::default() + .document_granularity(DocumentGranularity::ListElement), + ), + ) + .execute() + .await + .unwrap(); table .create_index(&["`META-DATA`.`USER-ID`"], Index::BTree(Default::default())) .execute() @@ -5924,6 +6437,35 @@ mod tests { assert_eq!(request_idx.load(Ordering::SeqCst), expected_requests.len()); } + #[tokio::test] + async fn test_create_list_element_fts_requires_server_support() { + let table = Table::new_with_handler_version( + "my_table", + semver::Version::new(0, 5, 0), + |_| -> http::Response { + panic!("unsupported index creation must fail before sending a request") + }, + ); + + let result = table + .create_index( + &["docs.content"], + Index::FTS( + InvertedIndexParams::default() + .document_granularity(DocumentGranularity::ListElement), + ), + ) + .execute() + .await; + + assert!( + result + .unwrap_err() + .to_string() + .contains("document granularity requires remote server version 0.6.0 or later") + ); + } + #[tokio::test] async fn test_list_indices() { let schema = Schema::new(vec![ @@ -7099,12 +7641,12 @@ mod tests { } /// The gate's reproducer: after a successful wait, a same-handle read - /// must carry a freshness baseline so a stale server cache cannot serve - /// the pre-backfill snapshot. + /// must carry the exact published version so a stale server cache cannot + /// serve the pre-backfill snapshot. #[tokio::test] async fn test_backfill_wait_establishes_read_freshness() { - let saw_min_timestamp = Arc::new(std::sync::atomic::AtomicBool::new(false)); - let saw = saw_min_timestamp.clone(); + let saw_published_version = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let saw = saw_published_version.clone(); let table = Table::new_with_handler("my_table", move |request| match request.url().path() { "/v1/table/my_table/backfill_column" => http::Response::builder() @@ -7117,7 +7659,11 @@ mod tests { .unwrap(), "/v1/table/my_table/count_rows/" => { saw.store( - request.headers().contains_key("x-lancedb-min-timestamp"), + request + .headers() + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()) + == Some("8"), std::sync::atomic::Ordering::SeqCst, ); http::Response::builder() @@ -7134,8 +7680,8 @@ mod tests { assert_eq!(result.published_version, Some(8)); table.count_rows(None).await.unwrap(); assert!( - saw_min_timestamp.load(std::sync::atomic::Ordering::SeqCst), - "read after wait carried no freshness baseline" + saw_published_version.load(std::sync::atomic::Ordering::SeqCst), + "read after wait did not carry the published version" ); } @@ -7310,12 +7856,11 @@ mod tests { } /// checkout_latest keeps the handle on latest, so a completed backfill - /// must still establish its post-fill baseline -- strictly later than the - /// checkout's own, or a pre-fill cache could still serve. + /// must retain the checkout timestamp and add its exact published version. #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn test_checkout_latest_during_submission_keeps_the_fence() { - let seen_min_timestamp = Arc::new(std::sync::Mutex::new(None::)); - let saw = seen_min_timestamp.clone(); + let seen_headers = Arc::new(std::sync::Mutex::new(None::)); + let saw = seen_headers.clone(); let (release_tx, release_rx) = std::sync::mpsc::channel::<()>(); let release_rx = Arc::new(std::sync::Mutex::new(release_rx)); let (arrived_tx, arrived_rx) = std::sync::mpsc::channel::<()>(); @@ -7339,10 +7884,7 @@ mod tests { .body(refresh_done("j-11")) .unwrap(), "/v1/table/my_table/count_rows/" => { - *saw.lock().unwrap() = request - .headers() - .get("x-lancedb-min-timestamp") - .map(|v| v.to_str().unwrap().to_string()); + *saw.lock().unwrap() = Some(request.headers().clone()); http::Response::builder() .status(200) .body("1".to_string()) @@ -7363,25 +7905,18 @@ mod tests { .await .unwrap(); table.checkout_latest().await.unwrap(); - let after_checkout = SystemTime::now(); - // Real separation between the checkout baseline and completion. - tokio::time::sleep(std::time::Duration::from_millis(50)).await; release_tx.send(()).unwrap(); let job = submit.await.unwrap().unwrap(); job.wait().await.unwrap(); table.count_rows(None).await.unwrap(); - let header = seen_min_timestamp - .lock() - .unwrap() - .clone() - .expect("no baseline"); - let sent: SystemTime = chrono::DateTime::parse_from_rfc3339(&header) - .unwrap() - .into(); - assert!( - sent > after_checkout, - "baseline {header} did not advance past the checkout" + let headers = seen_headers.lock().unwrap().clone().expect("no request"); + assert!(headers.contains_key("x-lancedb-min-timestamp")); + assert_eq!( + headers + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()), + Some("8") ); } @@ -8839,6 +9374,50 @@ mod tests { assert_ne!(Arc::as_ptr(&schema3), Arc::as_ptr(&schema1)); } + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn test_schema_fetch_does_not_cross_checkout_generation() { + let (release_tx, release_rx) = std::sync::mpsc::channel::<()>(); + let release_rx = Arc::new(std::sync::Mutex::new(release_rx)); + let (arrived_tx, arrived_rx) = std::sync::mpsc::channel::<()>(); + let arrived_tx = Arc::new(std::sync::Mutex::new(arrived_tx)); + let table = Table::new_with_handler("my_table", move |request| { + let body = request_body_json(&request); + let pinned = body["version"].as_u64() == Some(5); + if !pinned { + arrived_tx.lock().unwrap().send(()).unwrap(); + release_rx + .lock() + .unwrap() + .recv_timeout(Duration::from_secs(10)) + .unwrap(); + } + let field = if pinned { "pinned" } else { "latest" }; + http::Response::builder() + .status(200) + .body(format!( + r#"{{"version":5,"schema":{{"fields":[{{"name":"{field}","type":{{"type":"int32"}},"nullable":false}}]}}}}"# + )) + .unwrap() + }); + + let schema_fetch = tokio::spawn({ + let table = table.clone(); + async move { table.schema().await } + }); + tokio::task::spawn_blocking(move || { + arrived_rx.recv_timeout(Duration::from_secs(10)).unwrap() + }) + .await + .unwrap(); + table.checkout(5).await.unwrap(); + release_tx.send(()).unwrap(); + + let schema = schema_fetch.await.unwrap().unwrap(); + assert!(schema.field_with_name("pinned").is_ok()); + let cached = table.schema().await.unwrap(); + assert!(cached.field_with_name("pinned").is_ok()); + } + /// Test that schema cache is invalidated after checkout_latest #[tokio::test] async fn test_schema_cache_invalidation_on_checkout_latest() { @@ -10073,6 +10652,7 @@ mod tests { min_version: None, checkout_baseline: Some(baseline), min_read_version: None, + ..FreshnessState::default() }; assert_eq!(compute_min_timestamp(&state, None, now), Some(baseline)); @@ -10098,6 +10678,7 @@ mod tests { min_version: None, checkout_baseline: Some(baseline), min_read_version: None, + ..FreshnessState::default() }; assert_eq!( compute_min_timestamp(&state, Some(Duration::from_secs(10)), now), @@ -10110,6 +10691,7 @@ mod tests { min_version: None, checkout_baseline: Some(recent_baseline), min_read_version: None, + ..FreshnessState::default() }; assert_eq!( compute_min_timestamp(&state, Some(Duration::from_secs(60)), now), @@ -10186,6 +10768,110 @@ mod tests { assert!(!headers.contains_key("x-lancedb-min-version")); } + #[tokio::test] + async fn test_checkout_disables_read_consistency_interval() { + let (handler, captured) = capturing_handler(|path| match path { + "/v1/table/my_table/describe/" => r#"{"version":5,"schema":{"fields":[]}}"#.to_string(), + "/v1/table/my_table/count_rows/" => "42".to_string(), + _ => panic!("unexpected path: {}", path), + }); + let table = + Table::new_with_handler_and_interval("my_table", handler, Some(Duration::from_secs(0))); + + table.checkout(5).await.unwrap(); + table.count_rows(None).await.unwrap(); + + let headers = captured.lock().unwrap().clone().unwrap(); + assert!(!headers.contains_key("x-lancedb-min-timestamp")); + assert!(!headers.contains_key("x-lancedb-min-version")); + assert!(!headers.contains_key("x-lancedb-min-read-version")); + } + + #[tokio::test] + async fn test_read_snapshot_keeps_selector_and_freshness_generation_bound() { + let table = RemoteTable::new_mock_with_consistency_interval( + "my_table".to_string(), + |_| { + http::Response::builder() + .status(200) + .body(r#"{"version":5,"schema":{"fields":[]}}"#.to_string()) + .unwrap() + }, + Some(Duration::ZERO), + ); + + let latest = table.snapshot_read_state().await; + table.checkout(5).await.unwrap(); + + let latest_request = latest + .freshness + .apply( + table + .client + .post("/v1/table/my_table/count_rows/") + .json(&serde_json::json!({ "version": latest.version })), + ) + .build() + .unwrap(); + assert!(request_body_json(&latest_request)["version"].is_null()); + assert!(latest_request.headers().contains_key(MIN_TIMESTAMP_HEADER)); + + let pinned = table.snapshot_read_state().await; + let pinned_request = pinned + .freshness + .apply( + table + .client + .post("/v1/table/my_table/count_rows/") + .json(&serde_json::json!({ "version": pinned.version })), + ) + .build() + .unwrap(); + assert_eq!(request_body_json(&pinned_request)["version"], 5); + assert!(!pinned_request.headers().contains_key(MIN_TIMESTAMP_HEADER)); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn test_cancelled_checkout_keeps_latest_freshness_enabled() { + let (described_tx, described_rx) = std::sync::mpsc::channel::<()>(); + let table = Arc::new(RemoteTable::new_mock_with_consistency_interval( + "my_table".to_string(), + move |_| { + described_tx.send(()).unwrap(); + http::Response::builder() + .status(200) + .body(r#"{"version":5,"schema":{"fields":[]}}"#.to_string()) + .unwrap() + }, + Some(Duration::from_secs(0)), + )); + + let version_guard = table.version.write().await; + let checkout = tokio::spawn({ + let table = table.clone(); + async move { table.checkout(5).await } + }); + tokio::task::spawn_blocking(move || { + described_rx.recv_timeout(Duration::from_secs(10)).unwrap() + }) + .await + .unwrap(); + for _ in 0..100 { + if table.freshness.lock().unwrap().pinned { + break; + } + tokio::task::yield_now().await; + } + assert!(!checkout.is_finished()); + + checkout.abort(); + assert!(checkout.await.unwrap_err().is_cancelled()); + drop(version_guard); + + assert_eq!(*table.version.read().await, None); + assert!(table.snapshot_freshness_headers().min_timestamp.is_some()); + } + #[tokio::test] async fn test_freshness_positive_interval_sends_now_minus_interval() { let (handler, captured) = capturing_handler(|_| "42".to_string()); @@ -10254,6 +10940,62 @@ mod tests { ); } + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn test_inflight_write_result_cannot_cross_checkout_generation() { + let (release_tx, release_rx) = std::sync::mpsc::channel::<()>(); + let release_rx = Arc::new(std::sync::Mutex::new(release_rx)); + let (arrived_tx, arrived_rx) = std::sync::mpsc::channel::<()>(); + let arrived_tx = Arc::new(std::sync::Mutex::new(arrived_tx)); + let count_headers = Arc::new(std::sync::Mutex::new(None)); + let captured = count_headers.clone(); + let table = + Table::new_with_handler("my_table", move |request| match request.url().path() { + "/v1/table/my_table/update/" => { + arrived_tx.lock().unwrap().send(()).unwrap(); + release_rx + .lock() + .unwrap() + .recv_timeout(Duration::from_secs(10)) + .unwrap(); + http::Response::builder() + .status(200) + .body(r#"{"rows_updated":1,"version":100}"#.to_string()) + .unwrap() + } + "/v1/table/my_table/describe/" => http::Response::builder() + .status(200) + .body(r#"{"version":5,"schema":{"fields":[]}}"#.to_string()) + .unwrap(), + "/v1/table/my_table/count_rows/" => { + *captured.lock().unwrap() = Some(request.headers().clone()); + http::Response::builder() + .status(200) + .body("1".to_string()) + .unwrap() + } + path => panic!("unexpected path: {path}"), + }); + + let update = tokio::spawn({ + let table = table.clone(); + async move { table.update().column("a", "a + 1").execute().await } + }); + tokio::task::spawn_blocking(move || { + arrived_rx.recv_timeout(Duration::from_secs(10)).unwrap() + }) + .await + .unwrap(); + table.checkout(5).await.unwrap(); + release_tx.send(()).unwrap(); + update.await.unwrap().unwrap(); + table.count_rows(None).await.unwrap(); + + let headers = count_headers.lock().unwrap(); + let headers = headers.as_ref().unwrap(); + assert!(!headers.contains_key("x-lancedb-min-version")); + assert!(!headers.contains_key("x-lancedb-min-read-version")); + } + /// A handler that records every request's headers and answers each read with /// an `x-lancedb-version` response header taken from `versions` (by call /// index, saturating at the last entry). An empty string means "no header". @@ -10300,6 +11042,164 @@ mod tests { ); } + #[tokio::test] + async fn test_schema_response_advances_read_watermark() { + let requests = Arc::new(std::sync::Mutex::new(Vec::new())); + let captured = requests.clone(); + let table = Table::new_with_handler("my_table", move |request| { + captured + .lock() + .unwrap() + .push((request.url().path().to_string(), request.headers().clone())); + match request.url().path() { + "/v1/table/my_table/describe/" => http::Response::builder() + .status(200) + .header("x-lancedb-version", "100") + .body(r#"{"version":100,"schema":{"fields":[]}}"#.to_string()) + .unwrap(), + "/v1/table/my_table/count_rows/" => http::Response::builder() + .status(200) + .body("42".to_string()) + .unwrap(), + path => panic!("unexpected path: {path}"), + } + }); + + table.schema().await.unwrap(); + assert_eq!(table.count_rows(None).await.unwrap(), 42); + + let requests = requests.lock().unwrap(); + let count_headers = &requests[1].1; + assert_eq!( + count_headers + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()), + Some("100") + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn test_inflight_schema_response_cannot_cross_checkout_generation() { + let (release_tx, release_rx) = std::sync::mpsc::channel::<()>(); + let release_rx = Arc::new(std::sync::Mutex::new(release_rx)); + let (arrived_tx, arrived_rx) = std::sync::mpsc::channel::<()>(); + let arrived_tx = Arc::new(std::sync::Mutex::new(arrived_tx)); + let count_headers = Arc::new(std::sync::Mutex::new(None)); + let captured = count_headers.clone(); + let table = + Table::new_with_handler("my_table", move |request| match request.url().path() { + "/v1/table/my_table/describe/" => { + let body = request_body_json(&request); + if body["version"].is_null() { + arrived_tx.lock().unwrap().send(()).unwrap(); + release_rx + .lock() + .unwrap() + .recv_timeout(Duration::from_secs(10)) + .unwrap(); + http::Response::builder() + .status(200) + .header("x-lancedb-version", "100") + .body(r#"{"version":100,"schema":{"fields":[]}}"#.to_string()) + .unwrap() + } else { + http::Response::builder() + .status(200) + .body(r#"{"version":5,"schema":{"fields":[]}}"#.to_string()) + .unwrap() + } + } + "/v1/table/my_table/count_rows/" => { + *captured.lock().unwrap() = Some(request.headers().clone()); + http::Response::builder() + .status(200) + .body("1".to_string()) + .unwrap() + } + path => panic!("unexpected path: {path}"), + }); + + let schema = tokio::spawn({ + let table = table.clone(); + async move { table.schema().await } + }); + tokio::task::spawn_blocking(move || { + arrived_rx.recv_timeout(Duration::from_secs(10)).unwrap() + }) + .await + .unwrap(); + table.checkout(5).await.unwrap(); + release_tx.send(()).unwrap(); + schema.await.unwrap().unwrap(); + table.count_rows(None).await.unwrap(); + + assert!( + !count_headers + .lock() + .unwrap() + .as_ref() + .unwrap() + .contains_key("x-lancedb-min-read-version") + ); + } + + #[tokio::test] + async fn test_streaming_write_uses_and_advances_read_watermark() { + let data = RecordBatch::try_new( + Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])), + vec![Arc::new(Int32Array::from(vec![1, 2, 3]))], + ) + .unwrap(); + let describe_body = serde_json::to_string(&json!({ + "version": 7, + "schema": JsonSchema::try_from(data.schema().as_ref()).unwrap(), + })) + .unwrap(); + let requests = Arc::new(std::sync::Mutex::new(Vec::new())); + let captured = requests.clone(); + let table = Table::new_with_handler("my_table", move |request| { + captured + .lock() + .unwrap() + .push((request.url().path().to_string(), request.headers().clone())); + match request.url().path() { + "/v1/table/my_table/describe/" => http::Response::builder() + .status(200) + .body(describe_body.clone()) + .unwrap(), + "/v1/table/my_table/insert/" => http::Response::builder() + .status(200) + .header("x-lancedb-version", "8") + .body(r#"{"version":8}"#.to_string()) + .unwrap(), + "/v1/table/my_table/count_rows/" => http::Response::builder() + .status(200) + .body("3".to_string()) + .unwrap(), + path => panic!("unexpected path: {path}"), + } + }); + + assert_eq!(table.add(data).execute().await.unwrap().version, 8); + assert_eq!(table.count_rows(None).await.unwrap(), 3); + + let requests = requests.lock().unwrap(); + let insert_headers = &requests[1].1; + assert_eq!( + insert_headers + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()), + Some("7") + ); + let count_headers = &requests[2].1; + assert_eq!( + count_headers + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()), + Some("8") + ); + } + #[tokio::test] async fn test_read_version_watermark_keeps_max() { // Server reports 100 then a stale 50; the watermark must not regress. @@ -10715,6 +11615,86 @@ mod tests { ); } + #[tokio::test] + async fn test_main_only_metadata_is_unfenced_from_branch_timeline() { + let requests = Arc::new(std::sync::Mutex::new(HashMap::new())); + let captured = requests.clone(); + let saw_delete_response_floor = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let saw_high_floor = saw_delete_response_floor.clone(); + let table = RemoteTable::new_mock( + "my_table".to_string(), + move |request| { + let path = request.url().path().to_string(); + captured + .lock() + .unwrap() + .insert(path.clone(), request.headers().clone()); + match path.as_str() { + "/v1/table/my_table/count_rows/" => { + saw_high_floor.store( + request + .headers() + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()) + == Some("100"), + std::sync::atomic::Ordering::SeqCst, + ); + http::Response::builder() + .status(200) + .header("x-lancedb-version", "2") + .body("1".to_string()) + .unwrap() + } + "/v1/table/my_table/tags/list/" => http::Response::builder() + .status(200) + .body("{}".to_string()) + .unwrap(), + "/v1/table/my_table/tags/version/" => http::Response::builder() + .status(200) + .body(r#"{"version":1}"#.to_string()) + .unwrap(), + "/v1/table/my_table/tags/delete/" => http::Response::builder() + .status(200) + .header("x-lancedb-version", "100") + .body("{}".to_string()) + .unwrap(), + "/v1/table/my_table/branches/list/" => http::Response::builder() + .status(200) + .body(r#"{"branches":{}}"#.to_string()) + .unwrap(), + path => panic!("unexpected path: {path}"), + } + }, + None, + ); + let branch = table.with_branch(Some("exp".to_string())); + + branch.count_rows(None).await.unwrap(); + let mut tags = branch.tags().await.unwrap(); + tags.list().await.unwrap(); + tags.get_version("v1").await.unwrap(); + tags.delete("v1").await.unwrap(); + branch.list_branches().await.unwrap(); + branch.count_rows(None).await.unwrap(); + + let requests = requests.lock().unwrap(); + for path in [ + "/v1/table/my_table/tags/list/", + "/v1/table/my_table/tags/version/", + "/v1/table/my_table/tags/delete/", + "/v1/table/my_table/branches/list/", + ] { + assert!( + !requests[path].contains_key("x-lancedb-min-read-version"), + "{path} inherited the branch timeline" + ); + } + assert!( + !saw_delete_response_floor.load(std::sync::atomic::Ordering::SeqCst), + "tag deletion contaminated the branch timeline" + ); + } + #[tokio::test] async fn test_delete_branch() { let table = Table::new_with_handler("my_table", |request| { @@ -10806,6 +11786,50 @@ mod tests { assert!(result.main_version_after.is_none()); } + #[tokio::test] + async fn test_successful_cherry_pick_advances_main_read_watermark() { + let count_headers = Arc::new(std::sync::Mutex::new(None)); + let captured = count_headers.clone(); + let table = + Table::new_with_handler("my_table", move |request| match request.url().path() { + "/v1/table/my_table/branches/cherry_pick/" => { + let response = serde_json::json!({ + "status": "cherryPicked", + "diff": serde_json::from_str::(sample_branch_diff_json()) + .unwrap(), + "preview": { "promotedColumns": ["tag"] }, + "mainVersionAfter": 2 + }); + http::Response::builder() + .status(200) + .body(response.to_string()) + .unwrap() + } + "/v1/table/my_table/count_rows/" => { + *captured.lock().unwrap() = Some(request.headers().clone()); + http::Response::builder() + .status(200) + .body("1".to_string()) + .unwrap() + } + path => panic!("unexpected path: {path}"), + }); + + let result = table.cherry_pick("exp", false).await.unwrap(); + assert_eq!(result.status, crate::table::CherryPickStatus::CherryPicked); + table.count_rows(None).await.unwrap(); + assert_eq!( + count_headers + .lock() + .unwrap() + .as_ref() + .unwrap() + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()), + Some("2") + ); + } + #[tokio::test] async fn test_cherry_pick_failed_returns_ok_with_body() { let table = Table::new_with_handler("my_table", |request| { diff --git a/rust/lancedb/src/remote/table/blobs.rs b/rust/lancedb/src/remote/table/blobs.rs index 387d3c6dc..597b41782 100644 --- a/rust/lancedb/src/remote/table/blobs.rs +++ b/rust/lancedb/src/remote/table/blobs.rs @@ -6,6 +6,7 @@ use std::ops::Range; use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; +use std::time::Duration; use arrow_array::{Array, LargeBinaryArray}; use arrow_schema::DataType; @@ -20,7 +21,7 @@ use crate::error::Result; use crate::remote::client::{HttpSend, RequestResultExt, RestfulLanceDbClient}; use crate::table::BaseTable; -use super::{FreshnessHeaders, RemoteTable}; +use super::{FreshnessHeaders, FreshnessState, RemoteTable, freshness_headers_snapshot}; #[derive(Debug, Clone, Copy)] enum RangeRequestMode { @@ -43,7 +44,10 @@ struct TableBlobRangeRequester { path: String, version: Option, branch: Option, - freshness: FreshnessHeaders, + freshness: Arc>, + parent_freshness: Arc>, + parent_freshness_request: FreshnessHeaders, + read_consistency_interval: Option, } #[async_trait::async_trait] @@ -53,8 +57,9 @@ impl BlobRangeRequester for TableBlobRangeRequester { range_header: &str, mode: RangeRequestMode, ) -> Result<(String, Response)> { - let mut request = self - .freshness + let freshness_request = + freshness_headers_snapshot(&self.freshness, self.read_consistency_interval); + let mut request = freshness_request .apply(self.client.get(&self.path)) .header(header::RANGE, range_header); if let Some(version) = self.version { @@ -71,6 +76,9 @@ impl BlobRangeRequester for TableBlobRangeRequester { return Ok((request_id, response)); } let response = self.client.check_response(&request_id, response).await?; + freshness_request.observe_headers(&self.freshness, response.headers()); + self.parent_freshness_request + .observe_headers(&self.parent_freshness, response.headers()); Ok((request_id, response)) } } @@ -361,18 +369,21 @@ impl RemoteTable { message: "fetch_blobs is not supported on this LanceDB Cloud server".into(), }); } - let version = self.current_version().await; + let read_snapshot = self.snapshot_read_state().await; let mut body = serde_json::json!({ - "version": version, + "version": read_snapshot.version, "column": column, "row_ids": row_ids, }); self.apply_branch_body(&mut body); let request = self - .post_read(&format!("/v1/table/{}/fetch_blobs/", self.identifier)) + .client + .post(&format!("/v1/table/{}/fetch_blobs/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = self + .send_with_freshness(request, true, read_snapshot.freshness) + .await?; let mut stream = self.read_arrow_response(&request_id, response).await?; let mut blob_chunks: Vec> = Vec::new(); @@ -448,8 +459,7 @@ impl RemoteTable { }); } - let version = self.current_version().await; - let freshness = self.snapshot_freshness_headers(); + let read_snapshot = self.snapshot_read_state().await; let encoded_column = urlencoding::encode(column); let requesters = row_ids .iter() @@ -461,9 +471,12 @@ impl RemoteTable { let requester: Arc = Arc::new(TableBlobRangeRequester { client: self.client.clone(), path, - version, + version: read_snapshot.version, branch: self.branch.clone(), - freshness, + freshness: Arc::new(std::sync::Mutex::new(read_snapshot.freshness_state)), + parent_freshness: self.freshness.clone(), + parent_freshness_request: read_snapshot.freshness, + read_consistency_interval: self.client.read_consistency_interval, }); requester }) @@ -685,6 +698,46 @@ mod tests { assert!(requests.lock().unwrap().contains(&"bytes=5-11".to_string())); } + #[tokio::test] + async fn remote_blob_file_keeps_the_open_timeline_after_parent_checkout() { + let range_requests = Arc::new(StdMutex::new(Vec::new())); + let captured = range_requests.clone(); + let table = RemoteTable::new_mock( + "my_table".to_string(), + move |request| match request.url().path() { + "/v1/table/my_table/describe/" => http::Response::builder() + .status(200) + .body(r#"{"version":5,"schema":{"fields":[]}}"#.as_bytes().to_vec()) + .unwrap(), + "/v1/table/my_table/blob/image/10/bytes" => { + captured.lock().unwrap().push(( + request.url().query().unwrap_or_default().to_string(), + request.headers().clone(), + )); + range_response(&request, PAYLOAD) + } + path => panic!("unexpected path: {path}"), + }, + Some(Version::new(0, 5, 0)), + ); + + table.checkout(5).await.unwrap(); + let file = table + .fetch_blob_files_impl("image", &[10]) + .await + .unwrap() + .pop() + .flatten() + .unwrap(); + table.checkout_latest().await.unwrap(); + file.read_range(5..12).await.unwrap(); + + let requests = range_requests.lock().unwrap(); + let (query, headers) = requests.last().unwrap(); + assert!(query.contains("version=5")); + assert!(!headers.contains_key("x-lancedb-min-timestamp")); + } + #[tokio::test] async fn remote_blob_file_reuses_sequential_response_until_seek() { let requests = Arc::new(StdMutex::new(Vec::new())); diff --git a/rust/lancedb/src/remote/table/insert.rs b/rust/lancedb/src/remote/table/insert.rs index a4a28a9c6..4e0e0d666 100644 --- a/rust/lancedb/src/remote/table/insert.rs +++ b/rust/lancedb/src/remote/table/insert.rs @@ -24,7 +24,10 @@ use lance::io::exec::utils::InstrumentedRecordBatchStreamAdapter; use crate::Error; use crate::remote::ARROW_STREAM_CONTENT_TYPE; use crate::remote::client::{HttpSend, RestfulLanceDbClient, Sender}; -use crate::remote::table::{MergeInsertRequest, REQUEST_TIMEOUT_HEADER, RemoteTable}; +use crate::remote::table::{ + FreshnessHeaders, FreshnessState, MergeInsertRequest, REQUEST_TIMEOUT_HEADER, RemoteTable, + freshness_headers_snapshot, +}; use crate::table::datafusion::insert::COUNT_SCHEMA; use crate::table::write_progress::WriteProgressTracker; use crate::table::{AddResult, MergeResult}; @@ -54,6 +57,38 @@ pub enum WriteResult { Merge(MergeResult), } +#[derive(Debug, Clone, Default)] +struct WriteFreshness { + state: Option>>, + read_consistency_interval: Option, +} + +impl WriteFreshness { + fn prepare( + &self, + request: reqwest::RequestBuilder, + ) -> (reqwest::RequestBuilder, Option) { + match &self.state { + Some(state) => { + let freshness_request = + freshness_headers_snapshot(state, self.read_consistency_interval); + (freshness_request.apply(request), Some(freshness_request)) + } + None => (request, None), + } + } + + fn observe( + &self, + freshness_request: Option, + headers: &reqwest::header::HeaderMap, + ) { + if let (Some(state), Some(freshness_request)) = (&self.state, freshness_request) { + freshness_request.observe_headers(state, headers); + } + } +} + /// ExecutionPlan for streaming a write (add or merge_insert) to a remote /// LanceDB table. /// @@ -71,6 +106,7 @@ pub struct RemoteWriteExec { table_name: String, identifier: String, client: RestfulLanceDbClient, + freshness: WriteFreshness, input: Arc, op: WriteOp, properties: Arc, @@ -170,6 +206,7 @@ impl RemoteWriteExec { table_name, identifier, client, + freshness: WriteFreshness::default(), input, op, properties: Arc::new(properties), @@ -183,6 +220,18 @@ impl RemoteWriteExec { } } + pub(super) fn with_freshness( + mut self, + state: Arc>, + read_consistency_interval: Option, + ) -> Self { + self.freshness = WriteFreshness { + state: Some(state), + read_consistency_interval, + }; + self + } + /// Get the add result after execution, if this exec ran an insert. pub fn add_result(&self) -> Option { match self @@ -285,6 +334,7 @@ impl RemoteWriteExec { /// each threading the same handful of arguments. struct PartRequestCtx<'a, S: HttpSend> { client: &'a RestfulLanceDbClient, + freshness: &'a WriteFreshness, identifier: &'a str, table_name: &'a str, upload_id: &'a str, @@ -352,7 +402,11 @@ impl PartRequestCtx<'_, S> { } /// Build the `/insert` request for a single multipart part. - fn build_part_request(&self, part_id: &str, body: reqwest::Body) -> reqwest::RequestBuilder { + fn build_part_request( + &self, + part_id: &str, + body: reqwest::Body, + ) -> (reqwest::RequestBuilder, Option) { let mut request = self .client .post(&format!("/v1/table/{}/insert/", self.identifier)) @@ -368,12 +422,16 @@ impl PartRequestCtx<'_, S> { if let Some(b) = self.branch { request = request.query(&[("branch", b)]); } - request.body(body) + self.freshness.prepare(request.body(body)) } /// Send a single part's request and drain the response, mapping HTTP and /// table-not-found errors into `DataFusionError`. - async fn send_part_request(&self, request: reqwest::RequestBuilder) -> DataFusionResult<()> { + async fn send_part_request( + &self, + request: reqwest::RequestBuilder, + freshness_request: Option, + ) -> DataFusionResult<()> { let (request_id, response) = self .client .send(request) @@ -388,6 +446,8 @@ impl PartRequestCtx<'_, S> { .check_response(&request_id, response) .await .map_err(|e| DataFusionError::External(Box::new(e)))?; + self.freshness + .observe(freshness_request, response.headers()); response.bytes().await.map_err(|e| { DataFusionError::External(Box::new(Error::Http { source: Box::new(e), @@ -419,7 +479,7 @@ impl PartRequestCtx<'_, S> { let body = reqwest::Body::wrap_stream(chunk_rx); let part_id = uuid::Uuid::new_v4().to_string(); - let request = self.build_part_request(&part_id, body); + let (request, freshness_request) = self.build_part_request(&part_id, body); // Measured from just before the request is sent, matching the window the // client read timeout applies to the upload. @@ -495,7 +555,7 @@ impl PartRequestCtx<'_, S> { Ok::(input_ended) }; - let send = self.send_part_request(request); + let send = self.send_part_request(request, freshness_request); // `join!` rather than `tokio::spawn`: the producer borrows `input` (and // `schema`), so it cannot satisfy the `'static` bound a spawned task @@ -569,7 +629,7 @@ impl ExecutionPlan for RemoteWriteExec { // Building a fresh exec (with a new, empty `result`) is what makes the // outer rescannable retry loop work: `reset_state()` clears the captured // result so a re-execution starts clean. - Ok(Arc::new(Self::new_inner( + let mut exec = Self::new_inner( self.table_name.clone(), self.identifier.clone(), self.client.clone(), @@ -580,7 +640,9 @@ impl ExecutionPlan for RemoteWriteExec { self.branch.clone(), self.max_bytes_per_request, self.max_request_duration, - ))) + ); + exec.freshness = self.freshness.clone(); + Ok(Arc::new(exec)) } fn execute( @@ -613,6 +675,7 @@ impl ExecutionPlan for RemoteWriteExec { &self.metrics, )); let client = self.client.clone(); + let freshness = self.freshness.clone(); let identifier = self.identifier.clone(); let op = self.op.clone(); let result_slot = self.result.clone(); @@ -634,6 +697,7 @@ impl ExecutionPlan for RemoteWriteExec { let overwrite = matches!(op, WriteOp::Insert { overwrite: true }); let ctx = PartRequestCtx { client: &client, + freshness: &freshness, identifier: &identifier, table_name: &table_name, upload_id, @@ -688,7 +752,7 @@ impl ExecutionPlan for RemoteWriteExec { let (error_tx, mut error_rx) = tokio::sync::oneshot::channel(); let body = Self::stream_as_http_body(input_stream, error_tx, tracker)?; - let request = request.body(body); + let (request, freshness_request) = freshness.prepare(request.body(body)); let result: DataFusionResult<(String, _)> = async { let (request_id, response) = client @@ -708,6 +772,7 @@ impl ExecutionPlan for RemoteWriteExec { .check_response(&request_id, response) .await .map_err(|e| DataFusionError::External(Box::new(e)))?; + freshness.observe(freshness_request, response.headers()); Ok((request_id, response)) } diff --git a/rust/lancedb/src/table.rs b/rust/lancedb/src/table.rs index 5007a61d9..af8bcb5e2 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -59,7 +59,9 @@ use crate::index::{IndexConfig, IndexStatisticsImpl, IndexType}; use crate::job::Job; use crate::query::{IntoQueryVector, Query, QueryExecutionOptions, TakeQuery, VectorQuery}; use crate::table::datafusion::insert::InsertExec; -use crate::utils::{PatchReadParam, PatchWriteParam, resolve_arrow_field_path}; +use crate::utils::{ + PatchReadParam, PatchWriteParam, public_fts_field_path_by_id, resolve_arrow_field_path, +}; use self::dataset::DatasetConsistencyWrapper; use self::merge::MergeInsertBuilder; @@ -560,6 +562,13 @@ pub trait BaseTable: std::fmt::Display + std::fmt::Debug + Send + Sync { fn id(&self) -> &str; /// Get the arrow [Schema] of the table. async fn schema(&self) -> Result; + /// Create a read-only handle pinned to the table's current active revision. + /// + /// The returned handle is independent from later refreshes or checkouts on + /// this handle. This is used by bindings that must prepare client-side + /// query state from the same revision that the query will execute against. + #[doc(hidden)] + async fn query_snapshot(&self) -> Result>; /// Count the number of rows in this table. async fn count_rows(&self, filter: Option) -> Result; /// Create a physical plan for the query. @@ -1139,6 +1148,16 @@ impl Table { self.inner.schema().await } + /// Create a read-only handle pinned to the current active revision. + #[doc(hidden)] + pub async fn query_snapshot(&self) -> Result { + Ok(Self { + inner: self.inner.query_snapshot().await?, + database: self.database.clone(), + embedding_registry: self.embedding_registry.clone(), + }) + } + /// Count the number of rows in this dataset. /// /// # Arguments @@ -1758,7 +1777,23 @@ impl Table { self.inner.alter_columns(alterations).await } - /// Update per-field metadata (merges by default). + /// Update per-field (column) metadata. + /// + /// Each [`FieldMetadataUpdate`] is merged into the field's existing metadata + /// by default; use [`FieldMetadataUpdate::remove`] to delete a key, or + /// [`FieldMetadataUpdate::replace`] to swap the field's entire metadata map. + /// + /// The following keys are treated specially, by convention, and should be + /// used when appropriate: + /// + /// - `lancedb:description`: for a human-readable description of a field. + /// - `lancedb:tag:`: for a user-defined key-value tag, where the suffix + /// names the tag category; e.g. `lancedb:tag:model: "clip"`. + /// - `lancedb:logical-column`: for a column grouping; e.g. `feature_v1` and + /// `feature_v2` might be in the same logical column. + /// - `lancedb:status`: for status options (`production`, `candidate`, + /// `deprecated`, `archived`) to designate the current life cycle state of + /// this column. pub async fn update_field_metadata( &self, updates: &[FieldMetadataUpdate], @@ -3059,6 +3094,17 @@ impl BaseTable for NativeTable { &self.id } + async fn query_snapshot(&self) -> Result> { + let snapshot = self.dataset.new_query_snapshot().await?; + let mut table = self.with_dataset(snapshot); + // QueryTable requests do not carry a revision. A pinned snapshot must + // execute locally until the namespace API can accept that revision. + table + .pushdown_operations + .remove(&NamespaceClientPushdownOperation::QueryTable); + Ok(Arc::new(table)) + } + async fn version(&self) -> Result { Ok(self.dataset.get().await?.version().version) } @@ -3531,7 +3577,14 @@ impl BaseTable for NativeTable { let field_ids = idx_desc.field_ids(); let mut columns = Vec::with_capacity(field_ids.len()); for field_id in field_ids { - let field_path = match dataset.schema().field_path(*field_id as i32) { + let field_path = match if index_type == crate::index::IndexType::FTS { + public_fts_field_path_by_id(dataset.schema(), *field_id as i32) + } else { + dataset + .schema() + .field_path(*field_id as i32) + .map_err(Into::into) + } { Ok(field_path) => field_path, Err(e) => { log::warn!( @@ -4118,6 +4171,14 @@ mod tests { parent_list_calls: self.parent_list_calls.clone(), }) } + + fn wrap_paginated( + &self, + _store_prefix: &str, + _original: Arc, + ) -> Option> { + None + } } #[tokio::test] @@ -4221,6 +4282,14 @@ mod tests { self.called.store(true, Ordering::Relaxed); original } + + fn wrap_paginated( + &self, + _store_prefix: &str, + original: Arc, + ) -> Option> { + Some(original) + } } #[tokio::test] diff --git a/rust/lancedb/src/table/create_index.rs b/rust/lancedb/src/table/create_index.rs index e373522bc..c7d6b5675 100644 --- a/rust/lancedb/src/table/create_index.rs +++ b/rust/lancedb/src/table/create_index.rs @@ -28,8 +28,9 @@ pub(super) type PreparedIndex = (String, Box, Ind use crate::index::Index; use crate::index::vector::{VectorIndex, suggested_num_sub_vectors}; use crate::utils::{ - supported_bitmap_data_type, supported_btree_data_type, supported_fm_data_type, - supported_fts_data_type, supported_label_list_data_type, supported_vector_data_type, + resolve_lance_fts_field_path, supported_bitmap_data_type, supported_btree_data_type, + supported_fm_data_type, supported_fts_data_type, supported_label_list_data_type, + supported_vector_data_type, }; use super::NativeTable; @@ -122,7 +123,20 @@ impl NativeTable { } self.dataset.ensure_mutable()?; let dataset = self.dataset.get().await?; - let (column, field) = Self::resolve_index_field(dataset.schema(), &opts.columns[0])?; + let (column, field) = if let Index::FTS(params) = &opts.index { + let resolved = resolve_lance_fts_field_path(dataset.schema(), &opts.columns[0])?; + if params.get_document_granularity().is_list_element() && resolved.list_depth == 0 { + return Err(Error::InvalidInput { + message: format!( + "FTS field path '{}' has no List layer and cannot use ListElement document granularity", + resolved.canonical_path + ), + }); + } + (resolved.canonical_path, resolved.field) + } else { + Self::resolve_index_field(dataset.schema(), &opts.columns[0])? + }; let params = self.make_index_params(&field, opts.index.clone()).await?; let index_type = self.get_index_type_for_field(&field, &opts.index); Ok((column, params, index_type)) @@ -436,7 +450,7 @@ mod tests { use crate::connection::ConnectBuilder; use crate::index::Index; use crate::index::scalar::{ - BTreeIndexBuilder, BitmapIndexBuilder, FmIndexBuilder, FtsIndexBuilder, + BTreeIndexBuilder, BitmapIndexBuilder, DocumentGranularity, FmIndexBuilder, FtsIndexBuilder, }; use crate::index::vector::{ IvfHnswFlatIndexBuilder, IvfHnswPqIndexBuilder, IvfHnswSqIndexBuilder, @@ -553,6 +567,38 @@ mod tests { job.cancel().await.unwrap(); } + #[tokio::test] + async fn test_execute_async_validates_fts_input_before_starting_job() { + let conn = connect("memory://").execute().await.unwrap(); + let batch = + record_batch!(("id", Int32, [1, 2]), ("text", Utf8, ["alpha", "beta"])).unwrap(); + let table = conn.create_table("t", batch).execute().await.unwrap(); + + let missing = table + .create_index(&["missing"], Index::FTS(FtsIndexBuilder::default())) + .execute_async() + .await; + assert!(missing.is_err()); + + let invalid_type = table + .create_index(&["id"], Index::FTS(FtsIndexBuilder::default())) + .execute_async() + .await; + assert!(invalid_type.is_err()); + + let invalid_granularity = table + .create_index( + &["text"], + Index::FTS( + FtsIndexBuilder::default() + .document_granularity(DocumentGranularity::ListElement), + ), + ) + .execute_async() + .await; + assert!(invalid_granularity.is_err()); + } + /// Concurrent waiters, and a wait issued after the job settled, all /// succeed once the build does. #[tokio::test] diff --git a/rust/lancedb/src/table/dataset.rs b/rust/lancedb/src/table/dataset.rs index e37c3fc9f..5e3733b85 100644 --- a/rust/lancedb/src/table/dataset.rs +++ b/rust/lancedb/src/table/dataset.rs @@ -32,6 +32,10 @@ struct DatasetState { /// `Some(version)` = pinned to a specific version (time travel), /// `None` = tracking latest. pinned_version: Option, + /// Whether the pin is an internal query snapshot rather than user-visible + /// time travel. Query snapshots remain read-only but preserve MemWAL read + /// semantics. + query_snapshot: bool, } #[derive(Debug, Clone)] @@ -70,6 +74,7 @@ impl DatasetConsistencyWrapper { state: Arc::new(Mutex::new(DatasetState { dataset, pinned_version: None, + query_snapshot: false, })), consistency, shard_writer: Arc::new(ShardWriterCache::default()), @@ -93,6 +98,36 @@ impl DatasetConsistencyWrapper { wrapper } + /// Create an independent read-only wrapper pinned to the current dataset + /// while retaining this wrapper's live MemWAL read context. + pub async fn new_query_snapshot(&self) -> Result { + // Apply the configured consistency policy before taking the snapshot. + // The returned dataset is intentionally discarded: a checkout may race + // after this await, so the dataset and its pin provenance must instead + // be cloned together from one authoritative state sample below. + self.get().await?; + + let (dataset, query_snapshot) = { + let state = self.state.lock()?; + // Preserve user time travel so the MemWAL safety guard still sees + // it. Latest and already-internal snapshots remain internal pins. + ( + state.dataset.clone(), + state.query_snapshot || state.pinned_version.is_none(), + ) + }; + let version = dataset.version().version; + Ok(Self { + state: Arc::new(Mutex::new(DatasetState { + dataset, + pinned_version: Some(version), + query_snapshot, + })), + consistency: ConsistencyMode::Lazy, + shard_writer: self.shard_writer.clone(), + }) + } + /// The MemWAL `ShardWriter` cache co-located with this dataset. pub(crate) fn shard_writer(&self) -> &Arc { &self.shard_writer @@ -169,6 +204,7 @@ impl DatasetConsistencyWrapper { let mut state = self.state.lock()?; state.dataset = Arc::new(new_dataset); state.pinned_version = None; + state.query_snapshot = false; drop(state); if let ConsistencyMode::Eventual(bg_cache) = &self.consistency { bg_cache.invalidate(); @@ -202,10 +238,10 @@ impl DatasetConsistencyWrapper { /// Returns the version, if in time travel mode, or None otherwise. pub fn time_travel_version(&self) -> Option { - self.state - .lock() - .unwrap_or_else(|e| e.into_inner()) - .pinned_version + let state = self.state.lock().unwrap_or_else(|e| e.into_inner()); + (!state.query_snapshot) + .then_some(state.pinned_version) + .flatten() } /// Convert into a wrapper in latest version mode. @@ -225,6 +261,7 @@ impl DatasetConsistencyWrapper { if state.pinned_version.is_some() { state.dataset = Arc::new(new_dataset); state.pinned_version = None; + state.query_snapshot = false; } drop(state); if let ConsistencyMode::Eventual(bg_cache) = &self.consistency { @@ -260,6 +297,7 @@ impl DatasetConsistencyWrapper { let mut state = self.state.lock()?; state.dataset = Arc::new(new_dataset); state.pinned_version = Some(version_value); + state.query_snapshot = false; Ok(()) } @@ -461,6 +499,29 @@ mod tests { assert_eq!(wrapper.time_travel_version(), Some(1)); } + #[tokio::test] + async fn test_query_snapshot_samples_dataset_and_pin_together() { + let dir = tempfile::tempdir().unwrap(); + let uri = dir.path().to_str().unwrap(); + let ds = create_test_dataset(uri).await; + + let wrapper = DatasetConsistencyWrapper::new_latest(ds, None); + wrapper.as_time_travel(1u64).await.unwrap(); + let stale_time_travel_dataset = wrapper.get().await.unwrap(); + + append_to_dataset(uri).await; + wrapper.as_latest().await.unwrap(); + + let snapshot = wrapper.new_query_snapshot().await.unwrap(); + let snapshot_dataset = snapshot.get().await.unwrap(); + assert_eq!(snapshot_dataset.version().version, 2); + assert_ne!( + snapshot_dataset.version().version, + stale_time_travel_dataset.version().version + ); + assert_eq!(snapshot.time_travel_version(), None); + } + #[tokio::test] async fn test_as_latest_from_time_travel() { let dir = tempfile::tempdir().unwrap(); diff --git a/rust/lancedb/src/table/merge.rs b/rust/lancedb/src/table/merge.rs index 564a15e17..6b87d1080 100644 --- a/rust/lancedb/src/table/merge.rs +++ b/rust/lancedb/src/table/merge.rs @@ -1056,6 +1056,44 @@ mod lsm_tests { ); } + #[tokio::test] + async fn query_snapshot_preserves_lsm_read_semantics() { + let dir = tempdir().unwrap(); + let table = id_value_table(&dir).await; + table + .set_lsm_write_spec(LsmWriteSpec::unsharded()) + .await + .unwrap(); + lsm_upsert(&table, vec![4, 5]).await; + + let snapshot = table.query_snapshot().await.unwrap(); + let rows = collect_id_value(snapshot.query().execute().await.unwrap()).await; + assert_eq!( + rows.iter().map(|(id, _)| *id).collect::>(), + vec![1, 2, 3, 4, 5] + ); + } + + #[tokio::test] + async fn query_snapshot_preserves_time_travel_lsm_guard() { + let dir = tempdir().unwrap(); + let table = id_value_table(&dir).await; + table + .set_lsm_write_spec(LsmWriteSpec::unsharded()) + .await + .unwrap(); + lsm_upsert(&table, vec![4]).await; + + let version = table.version().await.unwrap(); + table.checkout(version).await.unwrap(); + let direct_error = table.query().execute().await.err().unwrap(); + assert!(matches!(direct_error, Error::NotSupported { .. })); + + let snapshot = table.query_snapshot().await.unwrap(); + let snapshot_error = snapshot.query().execute().await.err().unwrap(); + assert!(matches!(snapshot_error, Error::NotSupported { .. })); + } + #[tokio::test] async fn lsm_read_dedup_newest_wins() { let dir = tempdir().unwrap(); diff --git a/rust/lancedb/src/table/query.rs b/rust/lancedb/src/table/query.rs index 9b675ed45..8658ad0b7 100644 --- a/rust/lancedb/src/table/query.rs +++ b/rust/lancedb/src/table/query.rs @@ -28,7 +28,6 @@ use datafusion_physical_plan::projection::ProjectionExec; use datafusion_physical_plan::repartition::RepartitionExec; use datafusion_physical_plan::union::UnionExec; use datafusion_physical_plan::{ExecutionPlan, with_new_children_if_necessary}; -use futures::future::try_join_all; use lance::dataset::mem_wal::DatasetMemWalExt; use lance::dataset::scanner::DatasetRecordBatchStream; use lance::dataset::scanner::Scanner; @@ -182,6 +181,7 @@ pub async fn create_plan( let mut column = query.column.clone(); let mut query_vector = query.query_vector.first().cloned(); + let mut is_batch_query = false; if query.query_vector.len() > 1 { if column.is_none() { // Infer a vector column with the same dimension of the query vector. @@ -192,16 +192,37 @@ pub async fn create_plan( )?); } let vector_field = schema.field(column.as_ref().unwrap()).unwrap(); - if let DataType::List(_) = vector_field.data_type() { - // Multivector handling: concatenate into FixedSizeList> + let (_, element_type) = + lance::index::vector::utils::get_vector_type(schema, column.as_ref().unwrap())?; + let is_binary = matches!(element_type, DataType::UInt8); + if matches!(vector_field.data_type(), DataType::List(_)) + || (query.base.offset.unwrap_or(0) == 0 && !is_binary) + { + // Lance distinguishes these cases from the vector column type: a + // list-like query against a List column is one multivector query, + // while the same query against a FixedSizeList column is a batch of + // independent queries. The batch path shares a single flat scan and + // bounds retained candidate data instead of running one scan per + // query vector. let vectors = query .query_vector .iter() .map(|arr| arr.as_ref()) .collect::>(); let dim = vectors[0].len(); + if let Some((query_index, actual_dim)) = vectors + .iter() + .enumerate() + .find_map(|(index, vector)| (vector.len() != dim).then_some((index, vector.len()))) + { + return Err(Error::InvalidInput { + message: format!( + "query vector at index {query_index} has dimension {actual_dim}, expected {dim}" + ), + }); + } let mut fsl_builder = FixedSizeListBuilder::with_capacity( - Float32Builder::with_capacity(dim), + Float32Builder::with_capacity(dim * vectors.len()), dim as i32, vectors.len(), ); @@ -212,8 +233,12 @@ pub async fn create_plan( fsl_builder.append(true); } query_vector = Some(Arc::new(fsl_builder.finish())); + is_batch_query = !matches!(vector_field.data_type(), DataType::List(_)); } else { - // Multiple query vectors: create a plan for each and union them + // Lance's batch path has no per-query offset, and its binary path + // requires primitive UInt8 queries rather than a fixed-size list. + // Keep the prior plan shape for these cases so offsets are applied + // per query and binary query vectors retain their primitive shape. let query_vecs = query.query_vector.clone(); let plan_futures = query_vecs .into_iter() @@ -226,7 +251,7 @@ pub async fn create_plan( } }) .collect::>(); - let plans = try_join_all(plan_futures).await?; + let plans = futures::future::try_join_all(plan_futures).await?; return create_multi_vector_plan(plans); } } @@ -263,10 +288,14 @@ pub async fn create_plan( } } - scanner.limit( - query.base.limit.map(|limit| limit as i64), - query.base.offset.map(|offset| offset as i64), - )?; + // For a batch query, `nearest` already applies k to each query vector. + // Adding Scanner's global limit would truncate the combined result to k rows. + if !is_batch_query { + scanner.limit( + query.base.limit.map(|limit| limit as i64), + query.base.offset.map(|offset| offset as i64), + )?; + } if let Some(ef) = query.ef { scanner.ef(ef); @@ -1264,7 +1293,38 @@ mod tests { } #[tokio::test] - async fn test_create_plan_multivector_structure() { + async fn test_query_snapshot_disables_namespace_pushdown() { + use crate::connect; + use crate::table::BaseTable; + use arrow_array::{Int32Array, RecordBatch}; + use arrow_schema::{DataType, Field, Schema}; + + let conn = connect("memory://").execute().await.unwrap(); + let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)])); + let batch = + RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(vec![1, 2, 3]))]).unwrap(); + let table = conn + .create_table("test_snapshot_namespace_fallback", vec![batch]) + .execute() + .await + .unwrap(); + let mut native_table = table.as_native().unwrap().clone(); + native_table.namespace_client = Some(Arc::new(CountingNamespaceClient::default())); + native_table + .pushdown_operations + .insert(NamespaceClientPushdownOperation::QueryTable); + + let snapshot = BaseTable::query_snapshot(&native_table).await.unwrap(); + let snapshot = snapshot.as_any().downcast_ref::().unwrap(); + assert!( + !can_execute_namespace_query(snapshot, &AnyQuery::Query(QueryRequest::default()),) + .await + .unwrap() + ); + } + + #[tokio::test] + async fn test_create_plan_batch_vector_uses_shared_scan() { use arrow_array::{Float32Array, RecordBatch}; use arrow_schema::{DataType, Field, Schema}; use datafusion_physical_plan::display::DisplayableExecutionPlan; @@ -1291,11 +1351,18 @@ mod tests { .unwrap(); let native_table = table.as_native().unwrap(); - // This triggers the "create_multi_vector_plan" logic branch + // A batch of vectors against a fixed-size vector column should use + // Lance's native batch KNN path instead of independent scan plans. let q1 = Arc::new(Float32Array::from(vec![1.0, 2.0])); let q2 = Arc::new(Float32Array::from(vec![3.0, 4.0])); let req = VectorQueryRequest { + base: QueryRequest { + filter: Some(QueryFilter::Sql("id >= 0".to_string())), + limit: Some(1), + select: Select::Columns(vec!["id".to_string()]), + ..Default::default() + }, column: Some("vector".to_string()), query_vector: vec![q1, q2], ..Default::default() @@ -1312,19 +1379,17 @@ mod tests { .indent(true) .to_string(); - // We expect a RepartitionExec wrapping a UnionExec assert!( - display.contains("RepartitionExec"), - "Plan should include Repartitioning" + display.contains("KNNVectorDistance: queries=2"), + "plan should use native batch KNN, got:\n{display}" ); assert!( - display.contains("UnionExec"), - "Plan should include a Union of multiple searches" + !display.contains("UnionExec"), + "flat batch KNN should share one scan, got:\n{display}" ); - // We expect the projection to add the 'query_index' column (logic inside multi_vector_plan) assert!( display.contains("query_index"), - "Plan should add query_index column" + "plan should add query_index column, got:\n{display}" ); } diff --git a/rust/lancedb/src/table/query/lsm.rs b/rust/lancedb/src/table/query/lsm.rs index 4500bde42..7e0aa6f7a 100644 --- a/rust/lancedb/src/table/query/lsm.rs +++ b/rust/lancedb/src/table/query/lsm.rs @@ -302,7 +302,7 @@ async fn build_read_context( for shard_id in shard_ids { let manifest_store = ShardManifestStore::new(store.clone(), &base_path, shard_id, scan_batch_size); - if let Some(manifest) = manifest_store.read_latest().await? { + if let Some(manifest) = manifest_store.latest().await? { snapshots.push(snapshot_from_manifest(shard_id, &manifest, &exclude)); } } diff --git a/rust/lancedb/src/table/schema_evolution.rs b/rust/lancedb/src/table/schema_evolution.rs index d10a45eea..4f8dc811a 100644 --- a/rust/lancedb/src/table/schema_evolution.rs +++ b/rust/lancedb/src/table/schema_evolution.rs @@ -55,7 +55,9 @@ pub struct DropColumnsResult { pub struct FieldMetadataUpdate { /// Dot-separated path to the field (e.g. `"embedding"` or `"address.zip"`). pub path: String, - /// Keys to set (`Some`) or delete (`None`). + /// Keys to set (`Some`) or delete (`None`). See + /// [`Table::update_field_metadata`](crate::Table::update_field_metadata) for + /// the conventional `lancedb:*` keys. pub metadata: HashMap>, /// If `true`, replace the field's entire metadata map instead of merging. pub replace: bool, diff --git a/rust/lancedb/src/utils/mod.rs b/rust/lancedb/src/utils/mod.rs index 8bd306988..07d1836a1 100644 --- a/rust/lancedb/src/utils/mod.rs +++ b/rust/lancedb/src/utils/mod.rs @@ -225,6 +225,159 @@ pub(crate) fn resolve_arrow_field_path(schema: &Schema, column: &str) -> Result< Ok((canonical_path, Field::from(*field))) } +pub(crate) struct ResolvedFtsField { + pub canonical_path: String, + pub field: Field, + pub list_depth: usize, +} + +/// Canonicalize a public FTS field path while keeping Arrow list item names hidden. +pub(crate) fn resolve_lance_fts_field_path( + schema: &lance_core::datatypes::Schema, + column: &str, +) -> Result { + let names = + lance_core::datatypes::parse_field_path(column).map_err(|e| Error::InvalidInput { + message: format!("Invalid field path `{}`: {}", column, e), + })?; + let (root_name, remaining_names) = names.split_first().ok_or_else(|| Error::InvalidInput { + message: "FTS field path cannot be empty".to_string(), + })?; + let mut field = schema + .fields + .iter() + .find(|field| field.name == *root_name) + .or_else(|| { + schema + .fields + .iter() + .find(|field| field.name.eq_ignore_ascii_case(root_name)) + }) + .ok_or_else(|| fts_field_not_found(schema, column))?; + let mut canonical_names = vec![field.name.clone()]; + let mut list_depth = 0; + + for name in remaining_names { + while matches!( + field.data_type(), + DataType::List(_) | DataType::LargeList(_) + ) { + list_depth += 1; + field = field.children.first().ok_or_else(|| Error::Schema { + message: format!( + "FTS field path `{}` has a list without an item field", + column + ), + })?; + } + if !matches!(field.data_type(), DataType::Struct(_)) { + return Err(fts_field_not_found(schema, column)); + } + field = field + .children + .iter() + .find(|field| field.name == *name) + .or_else(|| { + field + .children + .iter() + .find(|field| field.name.eq_ignore_ascii_case(name)) + }) + .ok_or_else(|| fts_field_not_found(schema, column))?; + canonical_names.push(field.name.clone()); + } + + let mut terminal = field; + while matches!( + terminal.data_type(), + DataType::List(_) | DataType::LargeList(_) + ) { + list_depth += 1; + terminal = terminal.children.first().ok_or_else(|| Error::Schema { + message: format!( + "FTS field path `{}` has a list without an item field", + column + ), + })?; + } + + let canonical_path = lance_core::datatypes::format_field_path( + &canonical_names + .iter() + .map(String::as_str) + .collect::>(), + ); + Ok(ResolvedFtsField { + canonical_path, + field: Field::from(field), + list_depth, + }) +} + +fn fts_field_not_found(schema: &lance_core::datatypes::Schema, column: &str) -> Error { + Error::Schema { + message: format!( + "Field path `{}` not found in schema. Available field paths: {}", + column, + schema.field_paths().join(", ") + ), + } +} + +fn find_public_fts_field_path_by_id( + field: &lance_core::datatypes::Field, + field_id: i32, + path: &mut Vec, +) -> bool { + if field.id == field_id { + return true; + } + match field.data_type() { + DataType::List(_) | DataType::LargeList(_) => field + .children + .first() + .is_some_and(|child| find_public_fts_field_path_by_id(child, field_id, path)), + DataType::Struct(_) => field.children.iter().any(|child| { + path.push(child.name.clone()); + let found = find_public_fts_field_path_by_id(child, field_id, path); + if !found { + path.pop(); + } + found + }), + _ => false, + } +} + +pub(crate) fn public_fts_field_path_by_id( + schema: &lance_core::datatypes::Schema, + field_id: i32, +) -> Result { + for root in &schema.fields { + let mut path = vec![root.name.clone()]; + if find_public_fts_field_path_by_id(root, field_id, &mut path) { + return Ok(lance_core::datatypes::format_field_path( + &path.iter().map(String::as_str).collect::>(), + )); + } + } + Err(Error::Schema { + message: format!("Field id `{}` not found in schema", field_id), + }) +} + +pub(crate) fn resolve_arrow_fts_field_path( + schema: &Schema, + column: &str, +) -> Result<(String, Field)> { + let lance_schema = + lance_core::datatypes::Schema::try_from(schema).map_err(|e| Error::Schema { + message: format!("Invalid schema: {}", e), + })?; + let resolved = resolve_lance_fts_field_path(&lance_schema, column)?; + Ok((resolved.canonical_path, resolved.field)) +} + pub fn supported_btree_data_type(dtype: &DataType) -> bool { dtype.is_integer() || dtype.is_floating() @@ -480,6 +633,36 @@ mod tests { use super::*; + #[test] + fn test_public_fts_field_path_prefers_exact_case() { + let text_list = || { + DataType::List(Arc::new(Field::new( + "item", + DataType::Struct(vec![Field::new("content", DataType::Utf8, true)].into()), + true, + ))) + }; + let schema = Schema::new(vec![ + Field::new("Docs", text_list(), true), + Field::new("docs", text_list(), true), + ]); + + let (path, _) = resolve_arrow_fts_field_path(&schema, "docs.content").unwrap(); + assert_eq!(path, "docs.content"); + + let lance_schema = lance_core::datatypes::Schema::try_from(&schema).unwrap(); + let field_id = lance_schema + .resolve_case_insensitive("docs.item.content") + .unwrap() + .last() + .unwrap() + .id; + assert_eq!( + public_fts_field_path_by_id(&lance_schema, field_id).unwrap(), + "docs.content" + ); + } + #[test] fn test_guess_default_column() { let schema_no_vector = Schema::new(vec![