diff --git a/.bumpversion.toml b/.bumpversion.toml index f5f981b2c..1c4aea809 100644 --- a/.bumpversion.toml +++ b/.bumpversion.toml @@ -1,5 +1,5 @@ [tool.bumpversion] -current_version = "0.38.0-beta.6" +current_version = "0.38.0-beta.10" parse = """(?x) (?P0|[1-9]\\d*)\\. (?P0|[1-9]\\d*)\\. diff --git a/Cargo.lock b/Cargo.lock index 6c6dd4aea..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", @@ -5402,7 +5402,7 @@ dependencies = [ [[package]] name = "lancedb" -version = "0.38.0-beta.6" +version = "0.38.0-beta.10" dependencies = [ "ahash", "anyhow", @@ -5490,7 +5490,7 @@ dependencies = [ [[package]] name = "lancedb-nodejs" -version = "0.38.0-beta.6" +version = "0.38.0-beta.10" dependencies = [ "arrow-array", "arrow-buffer", @@ -5515,7 +5515,7 @@ dependencies = [ [[package]] name = "lancedb-python" -version = "0.38.0-beta.6" +version = "0.38.0-beta.10" dependencies = [ "arrow", "async-trait", 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/java/java.md b/docs/src/java/java.md index d6c285fc3..06dc267f3 100644 --- a/docs/src/java/java.md +++ b/docs/src/java/java.md @@ -14,7 +14,7 @@ Add the following dependency to your `pom.xml`: com.lancedb lancedb-core - 0.38.0-beta.6 + 0.38.0-beta.10 ``` diff --git a/docs/src/js/classes/AutoQuery.md b/docs/src/js/classes/AutoQuery.md new file mode 100644 index 000000000..1d7ea6952 --- /dev/null +++ b/docs/src/js/classes/AutoQuery.md @@ -0,0 +1,518 @@ +[**@lancedb/lancedb**](../README.md) • **Docs** + +*** + +[@lancedb/lancedb](../globals.md) / AutoQuery + +# Class: AutoQuery + +A builder for automatic string searches. + +Automatic search determines whether to use full-text or vector search from +the table revision selected for each execution. This builder exposes the +common operations supported by both query families. + +## Extends + +- `StandardQueryBase`<`NativeQuery` \| `NativeVectorQuery`> + +## Properties + +### inner + +```ts +protected inner: Query | VectorQuery | Promise; +``` + +#### Inherited from + +`StandardQueryBase.inner` + +## Methods + +### analyzePlan() + +```ts +analyzePlan(distributedMetrics?): Promise +``` + +Executes the query and returns the physical query plan annotated with runtime metrics. + +This is useful for debugging and performance analysis, as it shows how the query was executed +and includes metrics such as elapsed time, rows processed, and I/O statistics. + +#### Parameters + +* **distributedMetrics?**: [`AnalyzePlanDistributedMetrics`](../type-aliases/AnalyzePlanDistributedMetrics.md) + How distributed worker metrics are displayed for remote query plans. + Defaults to `"aggregate"`. + +#### Returns + +`Promise`<`string`> + +A query execution plan with runtime metrics for each step. + +#### Example + +```ts +import * as lancedb from "@lancedb/lancedb" + +const db = await lancedb.connect("./.lancedb"); +const table = await db.createTable("my_table", [ + { vector: [1.1, 0.9], id: "1" }, +]); + +const plan = await table.query().nearestTo([0.5, 0.2]).analyzePlan(); + +Example output (with runtime metrics inlined): +AnalyzeExec verbose=true, metrics=[] + ProjectionExec: expr=[id@3 as id, vector@0 as vector, _distance@2 as _distance], metrics=[output_rows=1, elapsed_compute=3.292µs] + Take: columns="vector, _rowid, _distance, (id)", metrics=[output_rows=1, elapsed_compute=66.001µs, batches_processed=1, bytes_read=8, iops=1, requests=1] + CoalesceBatchesExec: target_batch_size=1024, metrics=[output_rows=1, elapsed_compute=3.333µs] + GlobalLimitExec: skip=0, fetch=10, metrics=[output_rows=1, elapsed_compute=167ns] + FilterExec: _distance@2 IS NOT NULL, metrics=[output_rows=1, elapsed_compute=8.542µs] + SortExec: TopK(fetch=10), expr=[_distance@2 ASC NULLS LAST], metrics=[output_rows=1, elapsed_compute=63.25µs, row_replacements=1] + KNNVectorDistance: metric=l2, metrics=[output_rows=1, elapsed_compute=114.333µs, output_batches=1] + LanceScan: uri=/path/to/data, projection=[vector], row_id=true, row_addr=false, ordered=false, metrics=[output_rows=1, elapsed_compute=103.626µs, bytes_read=549, iops=2, requests=2] +``` + +#### Inherited from + +`StandardQueryBase.analyzePlan` + +*** + +### execute() + +```ts +protected execute(options?): AsyncGenerator, void, unknown> +``` + +Execute the query and return the results as an + +#### Parameters + +* **options?**: `Partial`<[`QueryExecutionOptions`](../interfaces/QueryExecutionOptions.md)> + +#### Returns + +`AsyncGenerator`<`RecordBatch`<`any`>, `void`, `unknown`> + +#### See + + - AsyncIterator +of + - RecordBatch. + +By default, LanceDb will use many threads to calculate results and, when +the result set is large, multiple batches will be processed at one time. +This readahead is limited however and backpressure will be applied if this +stream is consumed slowly (this constrains the maximum memory used by a +single query) + +#### Inherited from + +`StandardQueryBase.execute` + +*** + +### explainPlan() + +```ts +explainPlan(verbose): Promise +``` + +Generates an explanation of the query execution plan. + +#### Parameters + +* **verbose**: `boolean` = `false` + If true, provides a more detailed explanation. Defaults to false. + +#### Returns + +`Promise`<`string`> + +A Promise that resolves to a string containing the query execution plan explanation. + +#### Example + +```ts +import * as lancedb from "@lancedb/lancedb" +const db = await lancedb.connect("./.lancedb"); +const table = await db.createTable("my_table", [ + { vector: [1.1, 0.9], id: "1" }, +]); +const plan = await table.query().nearestTo([0.5, 0.2]).explainPlan(); +``` + +#### Inherited from + +`StandardQueryBase.explainPlan` + +*** + +### fastSearch() + +```ts +fastSearch(): this +``` + +Skip searching un-indexed data. This can make search faster, but will miss +any data that is not yet indexed. + +Use [Table#optimize](Table.md#optimize) to index all un-indexed data. + +#### Returns + +`this` + +#### Inherited from + +`StandardQueryBase.fastSearch` + +*** + +### ~~filter()~~ + +```ts +filter(predicate): this +``` + +A filter statement to be applied to this query. + +#### Parameters + +* **predicate**: `string` + +#### Returns + +`this` + +#### See + +where + +#### Deprecated + +Use `where` instead + +#### Inherited from + +`StandardQueryBase.filter` + +*** + +### fullTextSearch() + +```ts +fullTextSearch(query, options?): this +``` + +#### Parameters + +* **query**: `string` \| [`FullTextQuery`](../interfaces/FullTextQuery.md) + +* **options?**: `Partial`<[`FullTextSearchOptions`](../interfaces/FullTextSearchOptions.md)> + +#### Returns + +`this` + +#### Inherited from + +`StandardQueryBase.fullTextSearch` + +*** + +### limit() + +```ts +limit(limit): this +``` + +Set the maximum number of results to return. + +By default, a plain search has no limit. If this method is not +called then every valid row from the table will be returned. + +#### Parameters + +* **limit**: `number` + +#### Returns + +`this` + +#### Inherited from + +`StandardQueryBase.limit` + +*** + +### offset() + +```ts +offset(offset): this +``` + +Set the number of rows to skip before returning results. + +This is useful for pagination. + +#### Parameters + +* **offset**: `number` + +#### Returns + +`this` + +#### Inherited from + +`StandardQueryBase.offset` + +*** + +### orderBy() + +```ts +orderBy(ordering): this +``` + +Sort the results by the specified column(s). + +#### Parameters + +* **ordering**: [`ColumnOrdering`](../interfaces/ColumnOrdering.md) \| [`ColumnOrdering`](../interfaces/ColumnOrdering.md)[] + +#### Returns + +`this` + +This query builder. + +#### Inherited from + +`StandardQueryBase.orderBy` + +*** + +### outputSchema() + +```ts +outputSchema(): Promise> +``` + +Returns the schema of the output that will be returned by this query. + +This can be used to inspect the types and names of the columns that will be +returned by the query before executing it. + +#### Returns + +`Promise`<`Schema`<`any`>> + +An Arrow Schema describing the output columns. + +#### Inherited from + +`StandardQueryBase.outputSchema` + +*** + +### select() + +```ts +select(columns): this +``` + +Return only the specified columns. + +By default a query will return all columns from the table. However, this can have +a very significant impact on latency. LanceDb stores data in a columnar fashion. This +means we can finely tune our I/O to select exactly the columns we need. + +As a best practice you should always limit queries to the columns that you need. If you +pass in an array of column names then only those columns will be returned. + +You can also use this method to create new "dynamic" columns based on your existing columns. +For example, you may not care about "a" or "b" but instead simply want "a + b". This is often +seen in the SELECT clause of an SQL query (e.g. `SELECT a+b FROM my_table`). + +To create dynamic columns you can pass in a Map. A column will be returned +for each entry in the map. The key provides the name of the column. The value is +an SQL string used to specify how the column is calculated. + +For example, an SQL query might state `SELECT a + b AS combined, c`. The equivalent +input to this method would be: + +#### Parameters + +* **columns**: `string` \| `string`[] \| `Record`<`string`, `string`> \| `Map`<`string`, `string`> + +#### Returns + +`this` + +#### Example + +```ts +new Map([["combined", "a + b"], ["c", "c"]]) + +Columns will always be returned in the order given, even if that order is different than +the order used when adding the data. + +Note that you can pass in a `Record` (e.g. an object literal). This method +uses `Object.entries` which should preserve the insertion order of the object. However, +object insertion order is easy to get wrong and `Map` is more foolproof. +``` + +#### Inherited from + +`StandardQueryBase.select` + +*** + +### toArray() + +```ts +toArray(options?): Promise +``` + +Collect the results as an array of objects. + +#### Parameters + +* **options?**: `Partial`<[`QueryExecutionOptions`](../interfaces/QueryExecutionOptions.md)> + +#### Returns + +`Promise`<`any`[]> + +#### Inherited from + +`StandardQueryBase.toArray` + +*** + +### toArrow() + +```ts +toArrow(options?): Promise> +``` + +Collect the results as an Arrow + +#### Parameters + +* **options?**: `Partial`<[`QueryExecutionOptions`](../interfaces/QueryExecutionOptions.md)> + +#### Returns + +`Promise`<`Table`<`any`>> + +#### See + +ArrowTable. + +#### Inherited from + +`StandardQueryBase.toArrow` + +*** + +### useLsm() + +```ts +useLsm(enable): this +``` + +Control MemWAL read routing for this query. + +By default (unset), when the table carries a MemWAL write spec (see +[Table#setLsmWriteSpec](Table.md#setlsmwritespec)), reads are routed through the LSM scanner so +they also return data written via the `mergeInsert` LSM path that has not yet +been compacted into the base table (the active/frozen in-memory memtables and +the flushed generations), deduplicated by primary key; a table without a spec +reads the base table. + +#### Parameters + +* **enable**: `boolean` + `true` forces the LSM scanner and errors if the table has no + MemWAL write spec. `false` bypasses the MemWAL and reads the base table only, + even when a spec is present. + Note: the LSM scanner does not support every query shape (e.g. reranking, + hybrid search, `orderBy`). On a MemWAL table those shapes error unless + `useLsm(false)` is set, because a base-only read would silently exclude + un-compacted MemWAL data. + +#### Returns + +`this` + +#### Inherited from + +`StandardQueryBase.useLsm` + +*** + +### where() + +```ts +where(predicate): this +``` + +A filter statement to be applied to this query. + +The filter should be supplied as an SQL query string. For example: + +#### Parameters + +* **predicate**: `string` + +#### Returns + +`this` + +#### Example + +```ts +x > 10 +y > 0 AND y < 100 +x > 5 OR y = 'test' + +Filtering performance can often be improved by creating a scalar index +on the filter column(s). + +Calling this multiple times combines the filters with a logical AND rather +than replacing the previous filter. +``` + +#### Inherited from + +`StandardQueryBase.where` + +*** + +### withRowId() + +```ts +withRowId(): this +``` + +Whether to return the row id in the results. + +This column can be used to match results between different queries. For +example, to match results from a full text search and a vector search in +order to perform hybrid search. + +#### Returns + +`this` + +#### Inherited from + +`StandardQueryBase.withRowId` diff --git a/docs/src/js/classes/Connection.md b/docs/src/js/classes/Connection.md index 92cfd2568..1c5abd89f 100644 --- a/docs/src/js/classes/Connection.md +++ b/docs/src/js/classes/Connection.md @@ -584,6 +584,70 @@ Child namespace names and *** +### listTables() + +#### listTables(options) + +```ts +abstract listTables(options?): Promise +``` + +List a page of the tables in this database. + +To retrieve the tables after the page, pass the `pageToken` the response +carries back in. A page can be shorter than `limit` without being the last +one, so walk until a response carries no page token: + +```ts +const names = []; +let pageToken = undefined; +do { + const page = await conn.listTables({ pageToken, limit: 100 }); + names.push(...page.tables); + pageToken = page.pageToken; +} while (pageToken); +``` + +##### Parameters + +* **options?**: `Partial`<[`ListTablesOptions`](../interfaces/ListTablesOptions.md)> + Pagination options + (`pageToken`, `limit`). + +##### Returns + +`Promise`<[`ListTablesResponse`](../interfaces/ListTablesResponse.md)> + +A page of table names and an + optional token for the tables after it. + +#### listTables(namespacePath, options) + +```ts +abstract listTables(namespacePath?, options?): Promise +``` + +List a page of the tables in this database. + +##### Parameters + +* **namespacePath?**: `string`[] + The namespace path to list tables from + (defaults to root namespace) + +* **options?**: `Partial`<[`ListTablesOptions`](../interfaces/ListTablesOptions.md)> + Pagination options + (`pageToken`, `limit`). + +##### Returns + +`Promise`<[`ListTablesResponse`](../interfaces/ListTablesResponse.md)> + +A page of table names and an + optional token for the tables after it. + +*** + ### openMaterializedView() ```ts @@ -660,7 +724,7 @@ a "not supported" error. *** -### tableNames() +### ~~tableNames()~~ #### tableNames(options) @@ -682,6 +746,10 @@ Tables will be returned in lexicographical order. `Promise`<`string`[]> +##### Deprecated + +Use [Connection.listTables](Connection.md#listtables) instead. + #### tableNames(namespacePath, options) ```ts @@ -704,3 +772,7 @@ Tables will be returned in lexicographical order. ##### Returns `Promise`<`string`[]> + +##### Deprecated + +Use [Connection.listTables](Connection.md#listtables) instead. diff --git a/docs/src/js/classes/Table.md b/docs/src/js/classes/Table.md index 06dc8479e..9a85d0d96 100644 --- a/docs/src/js/classes/Table.md +++ b/docs/src/js/classes/Table.md @@ -942,7 +942,7 @@ Get the schema of the table. abstract search( query, queryType?, - ftsColumns?): Query | VectorQuery + ftsColumns?): Query | VectorQuery | AutoQuery ``` Create a search query to find the nearest neighbors @@ -964,7 +964,7 @@ of the given query #### Returns -[`Query`](Query.md) \| [`VectorQuery`](VectorQuery.md) +[`Query`](Query.md) \| [`VectorQuery`](VectorQuery.md) \| [`AutoQuery`](AutoQuery.md) *** diff --git a/docs/src/js/globals.md b/docs/src/js/globals.md index 6a8f644eb..beb9cbeff 100644 --- a/docs/src/js/globals.md +++ b/docs/src/js/globals.md @@ -18,6 +18,7 @@ ## Classes +- [AutoQuery](classes/AutoQuery.md) - [BooleanQuery](classes/BooleanQuery.md) - [BoostQuery](classes/BoostQuery.md) - [BranchContents](classes/BranchContents.md) @@ -100,6 +101,8 @@ - [JobInfo](interfaces/JobInfo.md) - [ListNamespacesOptions](interfaces/ListNamespacesOptions.md) - [ListNamespacesResponse](interfaces/ListNamespacesResponse.md) +- [ListTablesOptions](interfaces/ListTablesOptions.md) +- [ListTablesResponse](interfaces/ListTablesResponse.md) - [LsmStats](interfaces/LsmStats.md) - [LsmWriteSpec](interfaces/LsmWriteSpec.md) - [MaterializedViewDefinition](interfaces/MaterializedViewDefinition.md) diff --git a/docs/src/js/interfaces/ListTablesOptions.md b/docs/src/js/interfaces/ListTablesOptions.md new file mode 100644 index 000000000..52ace47e5 --- /dev/null +++ b/docs/src/js/interfaces/ListTablesOptions.md @@ -0,0 +1,34 @@ +[**@lancedb/lancedb**](../README.md) • **Docs** + +*** + +[@lancedb/lancedb](../globals.md) / ListTablesOptions + +# Interface: ListTablesOptions + +## Properties + +### limit? + +```ts +optional limit: number; +``` + +An upper bound on how many tables to return. + +A page may hold fewer than this and still not be the last one, so keep +going while the response carries a page token rather than while pages are +full. + +*** + +### pageToken? + +```ts +optional pageToken: string; +``` + +Token from a previous response, to resume listing where it left off. + +The token is opaque: it carries whatever the database needs to resume, and +callers should not construct or interpret one. diff --git a/docs/src/js/interfaces/ListTablesResponse.md b/docs/src/js/interfaces/ListTablesResponse.md new file mode 100644 index 000000000..76cac2b23 --- /dev/null +++ b/docs/src/js/interfaces/ListTablesResponse.md @@ -0,0 +1,23 @@ +[**@lancedb/lancedb**](../README.md) • **Docs** + +*** + +[@lancedb/lancedb](../globals.md) / ListTablesResponse + +# Interface: ListTablesResponse + +## Properties + +### pageToken? + +```ts +optional pageToken: string; +``` + +*** + +### tables + +```ts +tables: string[]; +``` diff --git a/docs/src/js/interfaces/TableNamesOptions.md b/docs/src/js/interfaces/TableNamesOptions.md index 45fa3d1b0..9254e9fa7 100644 --- a/docs/src/js/interfaces/TableNamesOptions.md +++ b/docs/src/js/interfaces/TableNamesOptions.md @@ -4,11 +4,16 @@ [@lancedb/lancedb](../globals.md) / TableNamesOptions -# Interface: TableNamesOptions +# Interface: ~~TableNamesOptions~~ + +## Deprecated + +Use [ListTablesOptions](ListTablesOptions.md) with [Connection.listTables](../classes/Connection.md#listtables) +instead. ## Properties -### limit? +### ~~limit?~~ ```ts optional limit: number; @@ -18,7 +23,7 @@ An optional limit to the number of results to return. *** -### startAfter? +### ~~startAfter?~~ ```ts optional startAfter: string; diff --git a/docs/src/js/namespaces/embedding/functions/getRegistry.md b/docs/src/js/namespaces/embedding/functions/getRegistry.md index 331dbd60b..b149a4a88 100644 --- a/docs/src/js/namespaces/embedding/functions/getRegistry.md +++ b/docs/src/js/namespaces/embedding/functions/getRegistry.md @@ -10,16 +10,12 @@ function getRegistry(): EmbeddingFunctionRegistry ``` -Utility function to get the global instance of the registry +Get the global embedding function registry. + +LanceDB built-in providers are initialized when this public API is first +used, so importing the root package does not change automatic search +selection for tables without embedding metadata. ## Returns [`EmbeddingFunctionRegistry`](../classes/EmbeddingFunctionRegistry.md) - -`EmbeddingFunctionRegistry` The global instance of the registry - -## Example - -```ts -const registry = getRegistry(); -const openai = registry.get("openai").create(); diff --git a/docs/src/python/python.md b/docs/src/python/python.md index 8a24ea199..3cbeee6f0 100644 --- a/docs/src/python/python.md +++ b/docs/src/python/python.md @@ -261,6 +261,8 @@ instead of being materialized with the rest of the row. ::: lancedb.streaming.StreamingDataset +::: lancedb.streaming.StreamingDataLoader + ::: lancedb.permutation.permutation_builder ::: lancedb.permutation.PermutationBuilder diff --git a/java/lancedb-core/pom.xml b/java/lancedb-core/pom.xml index 23d58124d..25e3b10e3 100644 --- a/java/lancedb-core/pom.xml +++ b/java/lancedb-core/pom.xml @@ -8,7 +8,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.6 + 0.38.0-beta.10 ../pom.xml diff --git a/java/pom.xml b/java/pom.xml index a07414ec0..a2cec19c0 100644 --- a/java/pom.xml +++ b/java/pom.xml @@ -6,7 +6,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.6 + 0.38.0-beta.10 pom ${project.artifactId} LanceDB Java SDK Parent POM @@ -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/Cargo.toml b/nodejs/Cargo.toml index ead873855..fd08e7a5e 100644 --- a/nodejs/Cargo.toml +++ b/nodejs/Cargo.toml @@ -1,7 +1,7 @@ [package] name = "lancedb-nodejs" edition.workspace = true -version = "0.38.0-beta.6" +version = "0.38.0-beta.10" publish = false license.workspace = true description.workspace = true diff --git a/nodejs/__test__/arrow.test.ts b/nodejs/__test__/arrow.test.ts index 160f4d5ef..c5bbbf169 100644 --- a/nodejs/__test__/arrow.test.ts +++ b/nodejs/__test__/arrow.test.ts @@ -515,6 +515,137 @@ describe.each([arrow15, arrow16, arrow17, arrow18])( ); }); + it("will allow matching inferred types across records", function () { + expect(() => + makeArrowTable([{ value: 1 }, { value: 2 }]), + ).not.toThrow(); + }); + + it("will reject mismatched inferred types across records", function () { + expect(() => makeArrowTable([{ value: 1 }, { value: "two" }])).toThrow( + "Failed to infer schema for data. Previously inferred type Float64 but found Utf8 for field value at row 1. Consider providing an explicit schema.", + ); + }); + + it("will ignore generated dictionary IDs when comparing inferred types", function () { + const table = makeArrowTable([{ str: "a" }, { str: "b" }], { + dictionaryEncodeStrings: true, + }); + + expect(table.getChild("str")?.toJSON()).toEqual(["a", "b"]); + }); + + it("will preserve null values without treating them as type mismatches", function () { + for (const records of [ + [{ vector: [1, 2, 3] }, { vector: null }], + [{ vector: null }, { vector: [1, 2, 3] }], + ]) { + const table = makeArrowTable(records); + + expect(table.numRows).toBe(2); + expect(table.getChild("vector")?.nullCount).toBe(1); + } + }); + + it("will preserve empty variable-size lists", function () { + for (const records of [ + [{ items: [1] }, { items: [] }], + [{ items: [] }, { items: [1] }], + ]) { + const table = makeArrowTable(records); + expect( + table + .getChild("items") + ?.toJSON() + .map((value) => value.toJSON()), + ).toEqual(records.map((record) => record.items)); + } + }); + + it("will propagate deferred evidence through nested lists", function () { + for (const records of [ + [{ items: [1] }, { items: [null] }], + [{ items: [null] }, { items: [1] }], + [{ items: [null, 1] }, { items: [2, null] }], + ]) { + const table = makeArrowTable(records); + expect( + table + .getChild("items") + ?.toJSON() + .map((value) => value.toJSON()), + ).toEqual(records.map((record) => record.items)); + } + + const nestedRecords = [{ items: [[1]] }, { items: [[null]] }]; + const nestedTable = makeArrowTable(nestedRecords); + expect( + nestedTable + .getChild("items") + ?.toJSON() + .map((value) => + value + .toJSON() + .map((nestedValue: { toJSON: () => unknown[] }) => + nestedValue.toJSON(), + ), + ), + ).toEqual(nestedRecords.map((record) => record.items)); + }); + + it("will reject incompatible deferred evidence within a list", function () { + for (const items of [ + [[], 1], + [1, []], + [[null], 1], + [1, [null]], + ]) { + expect(() => makeArrowTable([{ items }])).toThrow( + "Failed to infer data type for field items at row 0.", + ); + } + }); + + it("will reject empty fixed-size lists", function () { + expect(() => + makeArrowTable([{ vector: [1, 2, 3] }, { vector: [] }]), + ).toThrow( + "Failed to infer schema for data. Previously inferred type FixedSizeList[3] but found List[0] for field vector at row 1.", + ); + }); + + it("will reject inferred leaf and branch shape changes", function () { + expect(() => + makeArrowTable([{ value: 1 }, { value: { nested: 2 } }]), + ).toThrow( + "Failed to infer schema for data. Previously inferred type Float64 but found Struct for field value at row 1.", + ); + expect(() => + makeArrowTable([{ value: { nested: 1 } }, { value: 2 }]), + ).toThrow( + "Failed to infer schema for data. Previously inferred type Struct but found Float64 for field value at row 1.", + ); + }); + + it("will allow null values around inferred struct values", function () { + for (const { records, nullIndex } of [ + { + records: [{ value: null }, { value: { nested: 2 } }], + nullIndex: 0, + }, + { + records: [{ value: { nested: 1 } }, { value: null }], + nullIndex: 1, + }, + ]) { + const table = makeArrowTable(records); + const values = table.getChild("value"); + + expect(values?.nullCount).toBe(1); + expect(values?.get(nullIndex)).toBeNull(); + } + }); + it("will allow a schema to be provided", async function () { await checkTableCreation( async (records, _, schema) => diff --git a/nodejs/__test__/connection.test.ts b/nodejs/__test__/connection.test.ts index af471b478..816f8c3de 100644 --- a/nodejs/__test__/connection.test.ts +++ b/nodejs/__test__/connection.test.ts @@ -4,7 +4,13 @@ import { readdirSync } from "fs"; import { Field, Float64, Schema } from "apache-arrow"; import * as tmp from "tmp"; -import { Connection, Table, connect, connectNamespace } from "../lancedb"; +import { + Connection, + ListTablesResponse, + Table, + connect, + connectNamespace, +} from "../lancedb"; import { LocalTable } from "../lancedb/table"; describe("when connecting", () => { @@ -47,6 +53,7 @@ describe("given a connection", () => { await db.close(); expect(db.isOpen()).toBe(false); await expect(db.tableNames()).rejects.toThrow("Connection is closed"); + await expect(db.listTables()).rejects.toThrow("Connection is closed"); await expect(db.renameTable("a", "b")).rejects.toThrow( "Connection is closed", ); @@ -129,6 +136,66 @@ describe("given a connection", () => { expect(tables).toEqual(["b", "c"]); }); + it("should respect limit and page token when listing tables", async () => { + const db = await connect(tmpDir.name); + + await db.createTable("b", [{ id: 1 }]); + await db.createTable("a", [{ id: 1 }]); + await db.createTable("c", [{ id: 1 }]); + + const all = await db.listTables(); + expect(all.tables).toEqual(["a", "b", "c"]); + expect(all.pageToken).toBeUndefined(); + + const first = await db.listTables({ limit: 1 }); + expect(first.tables).toEqual(["a"]); + expect(first.pageToken).toBeDefined(); + + const second = await db.listTables({ + limit: 1, + pageToken: first.pageToken, + }); + expect(second.tables).toEqual(["b"]); + }); + + it("should visit every table exactly once when walking pages", async () => { + const db = await connect(tmpDir.name); + + const created = ["a", "b", "c", "d", "e"]; + for (const name of created) { + await db.createTable(name, [{ id: 1 }]); + } + + const seen: string[] = []; + let pageToken: string | undefined = undefined; + do { + const page: ListTablesResponse = await db.listTables({ + limit: 2, + pageToken, + }); + seen.push(...page.tables); + pageToken = page.pageToken; + } while (pageToken); + + expect(seen).toEqual(created); + }); + + it("should list tables in a namespace", async () => { + const db = await connect(tmpDir.name, { + // biome-ignore lint/style/useNamingConvention: opaque backend property key, must match Rust + namespaceClientProperties: { manifest_enabled: "true" }, + }); + await db.createNamespace(["child"]); + await db.createTable("nested", [{ id: 1 }], ["child"]); + + await expect(db.listTables(["child"])).resolves.toEqual( + expect.objectContaining({ tables: ["nested"] }), + ); + await expect(db.listTables()).resolves.toEqual( + expect.objectContaining({ tables: [] }), + ); + }); + it("should create tables in v2 mode", async () => { const db = await connect(tmpDir.name); const data = [...Array(10000).keys()].map((i) => ({ id: i })); diff --git a/nodejs/__test__/embedding_registry.test.ts b/nodejs/__test__/embedding_registry.test.ts new file mode 100644 index 000000000..83933399a --- /dev/null +++ b/nodejs/__test__/embedding_registry.test.ts @@ -0,0 +1,95 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The LanceDB Authors + +import { execFileSync } from "node:child_process"; +import { resolve } from "node:path"; + +import type { OpenAIEmbeddingFunction } from "../lancedb/embedding/openai"; +import type { EmbeddingFunctionRegistry } from "../lancedb/embedding/registry"; + +type EmbeddingModule = typeof import("../lancedb/embedding"); +type OpenAIModule = typeof import("../lancedb/embedding/openai"); +type RegistryModule = typeof import("../lancedb/embedding/registry"); + +describe("embedding function registry", () => { + const registries: EmbeddingFunctionRegistry[] = []; + + afterEach(() => { + for (const registry of registries) { + registry.reset(); + } + registries.length = 0; + }); + + it("defers built-in providers until the public registry API is used", () => { + jest.isolateModules(() => { + const embedding = require("../lancedb/embedding") as EmbeddingModule; + const { getRegistry: getInternalRegistry } = + require("../lancedb/embedding/registry") as RegistryModule; + const registry = getInternalRegistry(); + registries.push(registry); + + expect(registry.length()).toBe(0); + expect(embedding.getRegistry()).toBe(registry); + expect(registry.get("openai")).toBeDefined(); + expect(registry.get("huggingface")).toBeDefined(); + }); + }); + + it("preserves automatic FTS search in a fresh process", () => { + execFileSync( + process.execPath, + [resolve(__dirname, "fixtures", "auto_fts_search.cjs")], + { stdio: "pipe" }, + ); + }); + + it("shares registrations across duplicated provider module graphs", () => { + let registeringRegistry: EmbeddingFunctionRegistry | undefined; + let latestOpenAIConstructor: typeof OpenAIEmbeddingFunction | undefined; + + jest.isolateModules(() => { + require("../lancedb/embedding/openai"); + const { getRegistry } = + require("../lancedb/embedding/registry") as RegistryModule; + registeringRegistry = getRegistry(); + registries.push(registeringRegistry); + expect(registeringRegistry.get("openai")).toBeDefined(); + }); + + expect(() => { + jest.isolateModules(() => { + const { OpenAIEmbeddingFunction } = + require("../lancedb/embedding/openai") as OpenAIModule; + latestOpenAIConstructor = OpenAIEmbeddingFunction; + const { getRegistry } = + require("../lancedb/embedding/registry") as RegistryModule; + registries.push(getRegistry()); + }); + }).not.toThrow(); + + const previousApiKey = process.env.OPENAI_API_KEY; + process.env.OPENAI_API_KEY = "test"; + try { + const latestOpenAI = registeringRegistry! + .get("openai")! + .create(); + expect(latestOpenAI).toBeInstanceOf(latestOpenAIConstructor!); + } finally { + if (previousApiKey === undefined) { + delete process.env.OPENAI_API_KEY; + } else { + process.env.OPENAI_API_KEY = previousApiKey; + } + } + + jest.isolateModules(() => { + const { getRegistry } = + require("../lancedb/embedding") as EmbeddingModule; + const publicRegistry = getRegistry(); + registries.push(publicRegistry); + expect(publicRegistry).toBe(registeringRegistry); + expect(publicRegistry.get("openai")).toBeDefined(); + }); + }); +}); diff --git a/nodejs/__test__/fixtures/auto_fts_search.cjs b/nodejs/__test__/fixtures/auto_fts_search.cjs new file mode 100644 index 000000000..b5ab060b3 --- /dev/null +++ b/nodejs/__test__/fixtures/auto_fts_search.cjs @@ -0,0 +1,33 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The LanceDB Authors + +const assert = require("node:assert/strict"); +const tmp = require("tmp"); +const { connect, embedding, Index } = require("../../dist"); +const { getRegistry } = require("../../dist/embedding/registry"); + +async function main() { + assert.equal(typeof embedding.getRegistry, "function"); + assert.equal(getRegistry().length(), 0); + assert.equal(embedding.getRegistry(), getRegistry()); + assert.equal(getRegistry().length(), 2); + + const dir = tmp.dirSync({ unsafeCleanup: true }); + let db; + try { + db = await connect(dir.name); + const table = await db.createTable("docs", [{ text: "hello world" }]); + await table.createIndex("text", { config: Index.fts() }); + + const rows = await table.search("hello").toArray(); + assert.equal(rows[0].text, "hello world"); + } finally { + db?.close(); + dir.removeCallback(); + } +} + +main().catch((error) => { + console.error(error); + process.exitCode = 1; +}); diff --git a/nodejs/__test__/table.test.ts b/nodejs/__test__/table.test.ts index 0f6ed3615..dae640850 100644 --- a/nodejs/__test__/table.test.ts +++ b/nodejs/__test__/table.test.ts @@ -11,10 +11,13 @@ import * as arrow17 from "apache-arrow-17"; import * as arrow18 from "apache-arrow-18"; import { + AutoQuery, Connection, MatchQuery, PhraseQuery, + Query, Table, + VectorQuery, connect, tokenize, } from "../lancedb"; @@ -682,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; @@ -1777,6 +1830,194 @@ describe("Read consistency interval", () => { }); }); +describe("automatic search schema consistency", () => { + let tmpDir: tmp.DirResult; + + class SchemaRefreshEmbedding extends EmbeddingFunction { + ndims() { + return 2; + } + + embeddingDataType() { + return new Float32(); + } + + async computeSourceEmbeddings(data: string[]) { + return data.map((value) => [value.length, 1]); + } + + async computeQueryEmbeddings(value: string) { + return [value.length, 1]; + } + } + + function embeddingSchema() { + const func = new SchemaRefreshEmbedding(); + return LanceSchema({ + text: func.sourceField(new Utf8()), + vector: func.vectorField(), + }); + } + + beforeEach(() => { + getRegistry().reset(); + register("schema-refresh")(SchemaRefreshEmbedding); + tmpDir = tmp.dirSync({ unsafeCleanup: true }); + }); + + afterEach(() => { + getRegistry().reset(); + tmpDir.removeCallback(); + }); + + it("uses the schema refreshed from another connection", async () => { + const first = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + const second = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + + try { + const stale = await first.createTable("docs", [{ text: "before" }], { + schema: embeddingSchema(), + }); + const replacement = await second.createTable( + "docs", + [{ text: "after hello" }], + { mode: "overwrite" }, + ); + await replacement.createIndex("text", { config: Index.fts() }); + + const search = stale.search("hello"); + expect(search).toBeInstanceOf(AutoQuery); + expect(search).not.toBeInstanceOf(Query); + expect(search).not.toBeInstanceOf(VectorQuery); + expect("nprobes" in search).toBe(false); + + const rows = await search.toArray(); + expect(rows[0].text).toBe("after hello"); + expect((await stale.schema()).metadata.has("embedding_functions")).toBe( + false, + ); + } finally { + first.close(); + second.close(); + } + }); + + it("tracks embedding metadata across checkout and restore", async () => { + const first = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + const second = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + + try { + await first.createTable("docs", [{ text: "before" }], { + schema: embeddingSchema(), + }); + const table = await second.createTable( + "docs", + [{ text: "after hello" }], + { mode: "overwrite" }, + ); + await table.createIndex("text", { config: Index.fts() }); + + await table.checkout(1); + expect((await table.search("before").toArray())[0].text).toBe("before"); + + await table.checkoutLatest(); + expect((await table.search("hello").toArray())[0].text).toBe( + "after hello", + ); + + await table.checkout(1); + await table.restore(); + expect((await table.search("before").toArray())[0].text).toBe("before"); + } finally { + first.close(); + second.close(); + } + }); + + it("pins automatic search while computing an embedding", async () => { + let markStarted!: () => void; + let releaseEmbedding!: () => void; + const started = new Promise((resolve) => { + markStarted = resolve; + }); + const released = new Promise((resolve) => { + releaseEmbedding = resolve; + }); + + class BlockingEmbedding extends SchemaRefreshEmbedding { + async computeQueryEmbeddings(value: string) { + markStarted(); + await released; + return [value.length, 1]; + } + } + + register("schema-refresh-blocking")(BlockingEmbedding); + const func = new BlockingEmbedding(); + const schema = LanceSchema({ + text: func.sourceField(new Utf8()), + vector: func.vectorField(), + }); + const first = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + const second = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + + try { + const table = await first.createTable( + "docs", + [{ text: "hello before" }], + { schema }, + ); + const pending = table.search("hello").toArray(); + await started; + + const replacement = await second.createTable( + "docs", + [{ text: "hello after" }], + { mode: "overwrite" }, + ); + await replacement.createIndex("text", { config: Index.fts() }); + releaseEmbedding(); + + expect((await pending)[0].text).toBe("hello before"); + } finally { + releaseEmbedding(); + first.close(); + second.close(); + } + }); + + it("refreshes a reused automatic search for every execution", async () => { + const first = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + const second = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + + try { + const table = await first.createTable("docs", [ + { text: "hello before", marker: "before" }, + ]); + await table.createIndex("text", { config: Index.fts() }); + const search = table.search("hello").select(["text"]); + + const before = (await search.toArray())[0]; + expect(before.text).toBe("hello before"); + expect(before.marker).toBeUndefined(); + + const replacement = await second.createTable( + "docs", + [{ text: "hello after", marker: "after" }], + { mode: "overwrite" }, + ); + await replacement.createIndex("text", { config: Index.fts() }); + + const after = (await search.toArray())[0]; + expect(after.text).toBe("hello after"); + expect(after.marker).toBeUndefined(); + } finally { + first.close(); + second.close(); + } + }); +}); + describe("schema evolution", function () { let tmpDir: tmp.DirResult; beforeEach(() => { @@ -2344,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] }, @@ -2366,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 = [ @@ -2916,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/arrow.ts b/nodejs/lancedb/arrow.ts index 8b388d593..b52ab50ef 100644 --- a/nodejs/lancedb/arrow.ts +++ b/nodejs/lancedb/arrow.ts @@ -5,7 +5,6 @@ import { Data as ArrowData, Table as ArrowTable, Binary, - Bool, BufferType, DataType, DateUnit, @@ -18,12 +17,7 @@ import { FixedSizeList, Float, Float32, - Float64, Int, - Int8, - Int16, - Int32, - Int64, LargeBinary, List, Null, @@ -36,17 +30,16 @@ import { Struct, Timestamp, Type, - Uint8, - Uint16, - Uint32, Utf8, Vector, makeVector as arrowMakeVector, + util as arrowUtil, vectorFromArray as badVectorFromArray, makeBuilder, makeData, } from "apache-arrow"; import { Buffers } from "apache-arrow/data"; +import { typedArrayToArrowType } from "./arrow_type"; import { type EmbeddingFunction } from "./embedding/embedding_function"; import { EmbeddingFunctionConfig, @@ -59,14 +52,7 @@ import { sanitizeTable, sanitizeType, } from "./sanitize"; - -/** - * Check if a field name indicates a vector column. - */ -function nameSuggestsVectorColumn(fieldName: string): boolean { - const nameLower = fieldName.toLowerCase(); - return nameLower.includes("vector") || nameLower.includes("embedding"); -} +import { inferSchema } from "./schema"; export * from "apache-arrow"; export type SchemaLike = @@ -459,110 +445,6 @@ export function makeArrowTable( return new ArrowTable(inferredSchema, finalColumns); } -function inferSchema( - data: Array>, - schema: Schema | undefined, - opts: MakeArrowTableOptions, -): Schema { - // We will collect all fields we see in the data. - const pathTree = new PathTree(); - - for (const [rowI, row] of data.entries()) { - for (const [path, value] of rowPathsAndValues(row)) { - if (!pathTree.has(path)) { - // First time seeing this field. - if (schema !== undefined) { - const field = getFieldForPath(schema, path); - if (field === undefined) { - throw new Error( - `Found field not in schema: ${path.join(".")} at row ${rowI}`, - ); - } else { - pathTree.set(path, field.type); - } - } else { - const inferredType = inferType(value, path, opts); - if (inferredType === undefined) { - throw new Error(`Failed to infer data type for field ${path.join( - ".", - )} at row ${rowI}. \ - Consider providing an explicit schema.`); - } - pathTree.set(path, inferredType); - } - } else if (schema === undefined) { - const currentType = pathTree.get(path); - const newType = inferType(value, path, opts); - if (currentType !== newType) { - new Error(`Failed to infer schema for data. Previously inferred type \ - ${currentType} but found ${newType} at row ${rowI}. Consider \ - providing an explicit schema.`); - } - } - } - } - - if (schema === undefined) { - function fieldsFromPathTree(pathTree: PathTree): Field[] { - const fields = []; - for (const [name, value] of pathTree.map.entries()) { - if (value instanceof PathTree) { - const children = fieldsFromPathTree(value); - fields.push(new Field(name, new Struct(children), true)); - } else { - fields.push(new Field(name, value, true)); - } - } - return fields; - } - const fields = fieldsFromPathTree(pathTree); - return new Schema(fields); - } else { - function takeMatchingFields( - fields: Field[], - pathTree: PathTree, - ): Field[] { - const outFields = []; - for (const field of fields) { - if (pathTree.map.has(field.name)) { - const value = pathTree.get([field.name]); - if (value instanceof PathTree) { - const struct = field.type as Struct; - const children = takeMatchingFields(struct.children, value); - outFields.push( - new Field(field.name, new Struct(children), field.nullable), - ); - } else { - outFields.push( - new Field(field.name, value as DataType, field.nullable), - ); - } - } - } - return outFields; - } - const fields = takeMatchingFields(schema.fields, pathTree); - return new Schema(fields); - } -} - -function* rowPathsAndValues( - row: Record, - basePath: string[] = [], -): Generator<[string[], unknown]> { - for (const [key, value] of Object.entries(row)) { - if (isObject(value)) { - yield* rowPathsAndValues(value, [...basePath, key]); - } else { - // Skip undefined values - they should be treated the same as missing fields - // for embedding function purposes - if (value !== undefined) { - yield [[...basePath, key], value]; - } - } - } -} - function isObject(value: unknown): value is Record { return ( typeof value === "object" && @@ -577,146 +459,19 @@ function isObject(value: unknown): value is Record { ); } -function getFieldForPath(schema: Schema, path: string[]): Field | undefined { - let current: Field | Schema = schema; +function valueAtPath(datum: Record, path: string[]): unknown { + let current: unknown = datum; for (const key of path) { - if (current instanceof Schema) { - const field: Field | undefined = current.fields.find( - (f) => f.name === key, - ); - if (field === undefined) { - return undefined; - } - current = field; - } else if (current instanceof Field && DataType.isStruct(current.type)) { - const struct: Struct = current.type; - const field = struct.children.find((f) => f.name === key); - if (field === undefined) { - return undefined; - } - current = field; + if (current == null) { + return null; + } + if (isObject(current) && (Object.hasOwn(current, key) || key in current)) { + current = current[key]; } else { return undefined; } } - if (current instanceof Field) { - return current; - } else { - return undefined; - } -} - -/** - * Try to infer which Arrow type to use for a given value. - * - * May return undefined if the type cannot be inferred. - */ -function inferType( - value: unknown, - path: string[], - opts: MakeArrowTableOptions, -): DataType | undefined { - if (typeof value === "bigint") { - return new Int64(); - } else if (typeof value === "number") { - // Even if it's an integer, it's safer to assume Float64. Users can - // always provide an explicit schema or use BigInt if they mean integer. - return new Float64(); - } else if (typeof value === "string") { - if (opts.dictionaryEncodeStrings) { - return new Dictionary(new Utf8(), new Int32()); - } else { - return new Utf8(); - } - } else if (typeof value === "boolean") { - return new Bool(); - } else if (value instanceof Buffer) { - return new Binary(); - } else if (ArrayBuffer.isView(value) && !(value instanceof DataView)) { - const info = typedArrayToArrowType(value); - if (info !== undefined) { - const child = new Field("item", info.elementType, true); - return new FixedSizeList(info.length, child); - } - return undefined; - } else if (Array.isArray(value)) { - if (value.length === 0) { - return undefined; // Without any values we can't infer the type - } - if (path.length === 1 && Object.hasOwn(opts.vectorColumns, path[0])) { - const floatType = sanitizeType(opts.vectorColumns[path[0]].type); - return new FixedSizeList( - value.length, - new Field("item", floatType, true), - ); - } - const valueType = inferType(value[0], path, opts); - if (valueType === undefined) { - return undefined; - } - // Try to automatically detect embedding columns. - if (nameSuggestsVectorColumn(path[path.length - 1])) { - // Check if value is a Uint8Array for integer vector type determination - if (value instanceof Uint8Array) { - // For integer vectors, we default to Uint8 (matching Python implementation) - const child = new Field("item", new Uint8(), true); - return new FixedSizeList(value.length, child); - } else { - // For float vectors, we default to Float32 - const child = new Field("item", new Float32(), true); - return new FixedSizeList(value.length, child); - } - } else { - const child = new Field("item", valueType, true); - return new List(child); - } - } else { - // TODO: timestamp - return undefined; - } -} - -class PathTree { - map: Map>; - - constructor(entries?: [string[], V][]) { - this.map = new Map(); - if (entries !== undefined) { - for (const [path, value] of entries) { - this.set(path, value); - } - } - } - has(path: string[]): boolean { - let ref: PathTree = this; - for (const part of path) { - if (!(ref instanceof PathTree) || !ref.map.has(part)) { - return false; - } - ref = ref.map.get(part) as PathTree; - } - return true; - } - get(path: string[]): V | undefined { - let ref: PathTree = this; - for (const part of path) { - if (!(ref instanceof PathTree) || !ref.map.has(part)) { - return undefined; - } - ref = ref.map.get(part) as PathTree; - } - return ref as V; - } - set(path: string[], value: V): void { - let ref: PathTree = this; - for (const part of path.slice(0, path.length - 1)) { - if (!ref.map.has(part)) { - ref.map.set(part, new PathTree()); - } - ref = ref.map.get(part) as PathTree; - } - ref.map.set(path[path.length - 1], value); - } + return current; } function transposeData( @@ -724,37 +479,26 @@ function transposeData( field: Field, path: string[] = [], ): Vector { + const valuesPath = [...path, field.name]; + const values = data.map((datum) => valueAtPath(datum, valuesPath)); if (field.type instanceof Struct) { const childFields = field.type.children; - const fullPath = [...path, field.name]; const childVectors = childFields.map((child) => { - return transposeData(data, child, fullPath); + return transposeData(data, child, valuesPath); }); + const nullCount = values.filter((value) => value === null).length; const structData = makeData({ type: field.type, + length: values.length, + nullCount, + nullBitmap: + nullCount > 0 + ? arrowUtil.packBools(values.map((value) => value !== null)) + : undefined, children: childVectors as unknown as ArrowData[], }); return arrowMakeVector(structData); } else { - const valuesPath = [...path, field.name]; - const values = data.map((datum) => { - let current: unknown = datum; - for (const key of valuesPath) { - if (current == null) { - return null; - } - - if ( - isObject(current) && - (Object.hasOwn(current, key) || key in current) - ) { - current = current[key]; - } else { - return null; - } - } - return current; - }); return makeVector(values, field.type, undefined, field.nullable); } } @@ -797,32 +541,6 @@ function makeListVector(lists: unknown[][]): Vector { return listBuilder.finish().toVector(); } -/** - * Map a JS TypedArray instance to the corresponding Arrow element DataType - * and its length. Returns undefined if the value is not a recognized TypedArray. - */ -function typedArrayToArrowType( - value: ArrayBufferView, -): { elementType: DataType; length: number } | undefined { - if (value instanceof Float32Array) - return { elementType: new Float32(), length: value.length }; - if (value instanceof Float64Array) - return { elementType: new Float64(), length: value.length }; - if (value instanceof Uint8Array) - return { elementType: new Uint8(), length: value.length }; - if (value instanceof Uint16Array) - return { elementType: new Uint16(), length: value.length }; - if (value instanceof Uint32Array) - return { elementType: new Uint32(), length: value.length }; - if (value instanceof Int8Array) - return { elementType: new Int8(), length: value.length }; - if (value instanceof Int16Array) - return { elementType: new Int16(), length: value.length }; - if (value instanceof Int32Array) - return { elementType: new Int32(), length: value.length }; - return undefined; -} - /** Helper function to convert an Array of JS values to an Arrow Vector */ function makeVector( values: unknown[], @@ -1462,8 +1180,12 @@ export function ensureNestedFieldsExist( completeRow[field.name] = row[field.name]; } } else { - // Field is missing from the data - set to null - completeRow[field.name] = null; + // Keep a missing struct valid while filling each of its children with + // null. This is distinct from an explicitly null struct value. + completeRow[field.name] = + field.type.constructor.name === "Struct" + ? ensureStructFieldsExist({}, field.type as Struct) + : null; } } @@ -1498,8 +1220,12 @@ function ensureStructFieldsExist( completeStruct[childField.name] = data[childField.name]; } } else { - // Field is missing - set to null - completeStruct[childField.name] = null; + // Keep a missing struct valid while filling each of its children with + // null. This is distinct from an explicitly null struct value. + completeStruct[childField.name] = + childField.type.constructor.name === "Struct" + ? ensureStructFieldsExist({}, childField.type as Struct) + : null; } } diff --git a/nodejs/lancedb/arrow_type.ts b/nodejs/lancedb/arrow_type.ts new file mode 100644 index 000000000..35da346ef --- /dev/null +++ b/nodejs/lancedb/arrow_type.ts @@ -0,0 +1,40 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The LanceDB Authors + +import { + type DataType, + Float32, + Float64, + Int8, + Int16, + Int32, + Uint8, + Uint16, + Uint32, +} from "apache-arrow"; + +/** + * Map a JS TypedArray instance to the corresponding Arrow element type and + * length. Returns undefined when the view is not a supported TypedArray. + */ +export function typedArrayToArrowType( + value: ArrayBufferView, +): { elementType: DataType; length: number } | undefined { + if (value instanceof Float32Array) + return { elementType: new Float32(), length: value.length }; + if (value instanceof Float64Array) + return { elementType: new Float64(), length: value.length }; + if (value instanceof Uint8Array) + return { elementType: new Uint8(), length: value.length }; + if (value instanceof Uint16Array) + return { elementType: new Uint16(), length: value.length }; + if (value instanceof Uint32Array) + return { elementType: new Uint32(), length: value.length }; + if (value instanceof Int8Array) + return { elementType: new Int8(), length: value.length }; + if (value instanceof Int16Array) + return { elementType: new Int16(), length: value.length }; + if (value instanceof Int32Array) + return { elementType: new Int32(), length: value.length }; + return undefined; +} diff --git a/nodejs/lancedb/connection.ts b/nodejs/lancedb/connection.ts index e5528f8c6..263a338ab 100644 --- a/nodejs/lancedb/connection.ts +++ b/nodejs/lancedb/connection.ts @@ -31,12 +31,14 @@ import type { JobDescription, JobInfo, ListNamespacesResponse, + ListTablesResponse, } from "./native"; export type { CreateNamespaceResponse, DescribeNamespaceResponse, DropNamespaceResponse, ListNamespacesResponse, + ListTablesResponse, }; import { sanitizeTable } from "./sanitize"; import { LocalTable, Table } from "./table"; @@ -134,6 +136,10 @@ export interface OpenTableOptions { indexCacheSize?: number; } +/** + * @deprecated Use {@link ListTablesOptions} with {@link Connection.listTables} + * instead. + */ export interface TableNamesOptions { /** * If present, only return names that come lexicographically after the @@ -147,6 +153,24 @@ export interface TableNamesOptions { limit?: number; } +export interface ListTablesOptions { + /** + * Token from a previous response, to resume listing where it left off. + * + * The token is opaque: it carries whatever the database needs to resume, and + * callers should not construct or interpret one. + */ + pageToken?: string; + /** + * An upper bound on how many tables to return. + * + * A page may hold fewer than this and still not be the last one, so keep + * going while the response carries a page token rather than while pages are + * full. + */ + limit?: number; +} + export interface ListNamespacesOptions { /** Token from a previous response for pagination. */ pageToken?: string; @@ -231,6 +255,7 @@ export abstract class Connection { * @param {Partial} options - options to control the * paging / start point (backwards compatibility) * + * @deprecated Use {@link Connection.listTables} instead. */ abstract tableNames(options?: Partial): Promise; /** @@ -241,12 +266,53 @@ export abstract class Connection { * @param {Partial} options - options to control the * paging / start point * + * @deprecated Use {@link Connection.listTables} instead. */ abstract tableNames( namespacePath?: string[], options?: Partial, ): Promise; + /** + * List a page of the tables in this database. + * + * To retrieve the tables after the page, pass the `pageToken` the response + * carries back in. A page can be shorter than `limit` without being the last + * one, so walk until a response carries no page token: + * + * ```ts + * const names = []; + * let pageToken = undefined; + * do { + * const page = await conn.listTables({ pageToken, limit: 100 }); + * names.push(...page.tables); + * pageToken = page.pageToken; + * } while (pageToken); + * ``` + * + * @param {Partial} options - Pagination options + * (`pageToken`, `limit`). + * @returns {Promise} A page of table names and an + * optional token for the tables after it. + */ + abstract listTables( + options?: Partial, + ): Promise; + /** + * List a page of the tables in this database. + * + * @param {string[]} namespacePath - The namespace path to list tables from + * (defaults to root namespace) + * @param {Partial} options - Pagination options + * (`pageToken`, `limit`). + * @returns {Promise} A page of table names and an + * optional token for the tables after it. + */ + abstract listTables( + namespacePath?: string[], + options?: Partial, + ): Promise; + /** * Open a table in the database. * @param {string} name - The name of the table @@ -601,6 +667,25 @@ export class LocalConnection extends Connection { return await this.inner.listMaterializedViews(); } + async listTables( + namespacePathOrOptions?: string[] | Partial, + options?: Partial, + ): Promise { + // Detect if first argument is namespacePath array or options object + const namespacePath = Array.isArray(namespacePathOrOptions) + ? namespacePathOrOptions + : undefined; + const listTablesOptions = Array.isArray(namespacePathOrOptions) + ? options + : namespacePathOrOptions; + + return this.inner.listTables( + namespacePath ?? [], + listTablesOptions?.pageToken, + listTablesOptions?.limit, + ); + } + async openTable( name: string, namespacePath?: string[], diff --git a/nodejs/lancedb/embedding/index.ts b/nodejs/lancedb/embedding/index.ts index d0ffaec0d..d748e245a 100644 --- a/nodejs/lancedb/embedding/index.ts +++ b/nodejs/lancedb/embedding/index.ts @@ -4,7 +4,15 @@ import { Field, Schema } from "../arrow"; import { sanitizeType } from "../sanitize"; import { EmbeddingFunction } from "./embedding_function"; -import { EmbeddingFunctionConfig, getRegistry } from "./registry"; +import { + EmbeddingFunctionConfig, + EmbeddingFunctionRegistry, + getRegistry as getGlobalRegistry, + registerBuiltIn, +} from "./registry"; + +type OpenAIModule = typeof import("./openai"); +type TransformersModule = typeof import("./transformers"); export { FieldOptions, @@ -14,7 +22,39 @@ export { EmbeddingFunctionConstructor, } from "./embedding_function"; -export * from "./registry"; +export { + EmbeddingFunctionRegistry, + parseEmbeddingMetadata, + register, +} from "./registry"; +export type { + CreateReturnType, + EmbeddingFunctionConfig, + EmbeddingFunctionCreate, + EmbeddingMetadataEntry, + ResolvedEmbeddingFunctionConfig, +} from "./registry"; + +function initializeBuiltInProviders() { + const { OpenAIEmbeddingFunction } = require("./openai") as OpenAIModule; + const { TransformersEmbeddingFunction } = + require("./transformers") as TransformersModule; + + registerBuiltIn("openai", OpenAIEmbeddingFunction); + registerBuiltIn("huggingface", TransformersEmbeddingFunction); +} + +/** + * Get the global embedding function registry. + * + * LanceDB built-in providers are initialized when this public API is first + * used, so importing the root package does not change automatic search + * selection for tables without embedding metadata. + */ +export function getRegistry(): EmbeddingFunctionRegistry { + initializeBuiltInProviders(); + return getGlobalRegistry(); +} /** * Create a schema with embedding functions. diff --git a/nodejs/lancedb/embedding/openai.ts b/nodejs/lancedb/embedding/openai.ts index 5771cfeb5..2218d44bc 100644 --- a/nodejs/lancedb/embedding/openai.ts +++ b/nodejs/lancedb/embedding/openai.ts @@ -5,14 +5,13 @@ import type OpenAI from "openai"; import type { EmbeddingCreateParams } from "openai/resources/index"; import { Float, Float32 } from "../arrow"; import { EmbeddingFunction } from "./embedding_function"; -import { register } from "./registry"; +import { registerBuiltIn } from "./registry"; export type OpenAIOptions = { apiKey: string; model: EmbeddingCreateParams["model"]; }; -@register("openai") export class OpenAIEmbeddingFunction extends EmbeddingFunction< string, Partial @@ -100,3 +99,5 @@ export class OpenAIEmbeddingFunction extends EmbeddingFunction< return response.data[0].embedding; } } + +registerBuiltIn("openai", OpenAIEmbeddingFunction); diff --git a/nodejs/lancedb/embedding/registry.ts b/nodejs/lancedb/embedding/registry.ts index 5f32f683c..c9ee9135e 100644 --- a/nodejs/lancedb/embedding/registry.ts +++ b/nodejs/lancedb/embedding/registry.ts @@ -7,6 +7,10 @@ import { } from "./embedding_function"; import "reflect-metadata"; +const builtInFunctionsKey = Symbol.for( + "@lancedb/lancedb::embedding-built-in-functions::v1", +); + export type CreateReturnType = T extends { init: () => Promise } ? Promise : T; @@ -59,6 +63,15 @@ export class EmbeddingFunctionRegistry { }; } + /** @ignore */ + setBuiltIn< + T extends EmbeddingFunctionConstructor = EmbeddingFunctionConstructor, + >(name: string, ctor: T): T { + this.#functions.set(name, ctor); + Reflect.defineMetadata("lancedb::embedding::name", name, ctor); + return ctor; + } + get>( name: string, ): EmbeddingFunctionCreate | undefined; @@ -96,6 +109,7 @@ export class EmbeddingFunctionRegistry { */ reset(this: EmbeddingFunctionRegistry) { this.#functions.clear(); + getBuiltInFunctions(this).clear(); } /** @@ -183,12 +197,56 @@ export class EmbeddingFunctionRegistry { } } -const _REGISTRY = new EmbeddingFunctionRegistry(); +function getBuiltInFunctions(registry: EmbeddingFunctionRegistry): Set { + const registryWithBuiltIns = registry as EmbeddingFunctionRegistry & { + [key: symbol]: Set | undefined; + }; + let builtInFunctions = registryWithBuiltIns[builtInFunctionsKey]; + if (builtInFunctions === undefined) { + builtInFunctions = new Set(); + registryWithBuiltIns[builtInFunctionsKey] = builtInFunctions; + } + return builtInFunctions; +} + +// Server bundlers can load the side-effect embedding entry points and the public +// embedding API from separate module graphs. Keep their registry shared. +const registryKey = Symbol.for( + "@lancedb/lancedb::embedding-function-registry::v1", +); +const registryGlobal = globalThis as typeof globalThis & { + [key: symbol]: EmbeddingFunctionRegistry | undefined; +}; + +function getGlobalRegistry(): EmbeddingFunctionRegistry { + const existingRegistry = registryGlobal[registryKey]; + if (existingRegistry !== undefined) { + return existingRegistry; + } + const registry = new EmbeddingFunctionRegistry(); + registryGlobal[registryKey] = registry; + return registry; +} + +const _REGISTRY = getGlobalRegistry(); export function register(name?: string) { return _REGISTRY.register(name); } +/** @ignore */ +export function registerBuiltIn< + T extends EmbeddingFunctionConstructor = EmbeddingFunctionConstructor, +>(name: string, ctor: T): T { + const builtInFunctions = getBuiltInFunctions(_REGISTRY); + if (builtInFunctions.has(name)) { + return _REGISTRY.setBuiltIn(name, ctor); + } + _REGISTRY.register(name)(ctor); + builtInFunctions.add(name); + return ctor; +} + /** * Utility function to get the global instance of the registry * @returns `EmbeddingFunctionRegistry` The global instance of the registry diff --git a/nodejs/lancedb/embedding/transformers.ts b/nodejs/lancedb/embedding/transformers.ts index 161575285..06043ea7c 100644 --- a/nodejs/lancedb/embedding/transformers.ts +++ b/nodejs/lancedb/embedding/transformers.ts @@ -3,7 +3,7 @@ import { Float, Float32 } from "../arrow"; import { EmbeddingFunction } from "./embedding_function"; -import { register } from "./registry"; +import { registerBuiltIn } from "./registry"; export type XenovaTransformerOptions = { /** The wasm compatible model to use */ @@ -31,7 +31,6 @@ export type XenovaTransformerOptions = { }; }; -@register("huggingface") export class TransformersEmbeddingFunction extends EmbeddingFunction< string, Partial @@ -158,6 +157,8 @@ export class TransformersEmbeddingFunction extends EmbeddingFunction< } } +registerBuiltIn("huggingface", TransformersEmbeddingFunction); + const tensorDiv = ( src: import("@huggingface/transformers").Tensor, divBy: number, diff --git a/nodejs/lancedb/index.ts b/nodejs/lancedb/index.ts index da13b434f..34d7ce4d9 100644 --- a/nodejs/lancedb/index.ts +++ b/nodejs/lancedb/index.ts @@ -81,11 +81,13 @@ export { Connection, CreateTableOptions, TableNamesOptions, + ListTablesOptions, OpenTableOptions, ListNamespacesOptions, CreateNamespaceOptions, DropNamespaceOptions, ListNamespacesResponse, + ListTablesResponse, CreateNamespaceResponse, DropNamespaceResponse, DescribeNamespaceResponse, @@ -101,6 +103,7 @@ export { } from "./native.js"; export { + AutoQuery, ExecutableQuery, Query, QueryBase, diff --git a/nodejs/lancedb/query.ts b/nodejs/lancedb/query.ts index 843a1276f..f1d31eae1 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} @@ -111,13 +134,15 @@ export class QueryBase< NativeQueryType extends NativeQuery | NativeVectorQuery | NativeTakeQuery, > implements AsyncIterable { + protected inner!: NativeQueryType | Promise; + /** * @hidden */ - protected constructor( - protected inner: NativeQueryType | Promise, - ) { - // intentionally empty + protected constructor(inner?: NativeQueryType | Promise) { + if (inner !== undefined) { + this.inner = inner; + } } // call a function on the inner (either a promise or the actual object) @@ -135,6 +160,15 @@ export class QueryBase< } } + /** + * Return the native query used by the next terminal operation. + * + * @hidden + */ + protected async getInner(): Promise { + return this.inner; + } + /** * Return only the specified columns. * @@ -207,16 +241,11 @@ export class QueryBase< /** * @hidden */ - protected nativeExecute( + protected async nativeExecute( options?: Partial, ): Promise { - if (this.inner instanceof Promise) { - return this.inner.then((inner) => - inner.execute(options?.maxBatchLength, options?.timeoutMs), - ); - } else { - return this.inner.execute(options?.maxBatchLength, options?.timeoutMs); - } + const inner = await this.getInner(); + return inner.execute(options?.maxBatchLength, options?.timeoutMs); } /** @@ -245,12 +274,7 @@ export class QueryBase< /** Collect the results as an Arrow @see {@link ArrowTable}. */ async toArrow(options?: Partial): Promise { const batches = []; - let inner; - if (this.inner instanceof Promise) { - inner = await this.inner; - } else { - inner = this.inner; - } + const inner = await this.getInner(); for await (const batch of new RecordBatchIterable(inner, options)) { batches.push(batch); } @@ -279,11 +303,8 @@ export class QueryBase< * @returns A Promise that resolves to a string containing the query execution plan explanation. */ async explainPlan(verbose = false): Promise { - if (this.inner instanceof Promise) { - return this.inner.then((inner) => inner.explainPlan(verbose)); - } else { - return this.inner.explainPlan(verbose); - } + const inner = await this.getInner(); + return inner.explainPlan(verbose); } /** @@ -321,13 +342,8 @@ export class QueryBase< distributedMetrics?: AnalyzePlanDistributedMetrics, ): Promise { const distributedMetricsMode = distributedMetrics ?? "aggregate"; - if (this.inner instanceof Promise) { - return this.inner.then((inner) => - inner.analyzePlan(distributedMetricsMode), - ); - } else { - return this.inner.analyzePlan(distributedMetricsMode); - } + const inner = await this.getInner(); + return inner.analyzePlan(distributedMetricsMode); } /** @@ -339,12 +355,8 @@ export class QueryBase< * @returns An Arrow Schema describing the output columns. */ async outputSchema(): Promise { - let schemaBuffer: Buffer; - if (this.inner instanceof Promise) { - schemaBuffer = await this.inner.then((inner) => inner.outputSchema()); - } else { - schemaBuffer = await this.inner.outputSchema(); - } + const inner = await this.getInner(); + const schemaBuffer = await inner.outputSchema(); const schema = tableFromIPC(schemaBuffer).schema; return schema; } @@ -356,7 +368,7 @@ export class StandardQueryBase< extends QueryBase implements ExecutableQuery { - constructor(inner: NativeQueryType | Promise) { + constructor(inner?: NativeQueryType | Promise) { super(inner); } @@ -510,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) * @@ -537,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; } @@ -551,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; } @@ -565,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; } @@ -578,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; } @@ -592,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; } @@ -606,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; } @@ -627,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; } @@ -661,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; } @@ -686,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; } @@ -700,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; } @@ -716,35 +735,31 @@ export class VectorQuery extends StandardQueryBase { */ 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); @@ -763,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. * @@ -788,6 +868,51 @@ export class TakeQuery extends QueryBase { } } +/** + * A builder for automatic string searches. + * + * Automatic search determines whether to use full-text or vector search from + * the table revision selected for each execution. This builder exposes the + * common operations supported by both query families. + * + * @hideconstructor + */ +export class AutoQuery extends StandardQueryBase< + NativeQuery | NativeVectorQuery +> { + private readonly calls: Array< + (inner: NativeQuery | NativeVectorQuery) => void + > = []; + + /** @hidden */ + constructor( + private readonly createInner: () => Promise< + NativeQuery | NativeVectorQuery + >, + ) { + super(); + } + + /** @hidden */ + protected override doCall( + fn: (inner: NativeQuery | NativeVectorQuery) => void, + ) { + this.calls.push(fn); + } + + /** @hidden */ + protected override async getInner(): Promise< + NativeQuery | NativeVectorQuery + > { + const calls = [...this.calls]; + const inner = await this.createInner(); + for (const call of calls) { + call(inner); + } + return inner; + } +} + /** A builder for LanceDB queries. * * @see {@link Table#query}, {@link Table#search} @@ -840,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/schema.ts b/nodejs/lancedb/schema.ts new file mode 100644 index 000000000..e4749ef37 --- /dev/null +++ b/nodejs/lancedb/schema.ts @@ -0,0 +1,566 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The LanceDB Authors + +import { + Binary, + Bool, + DataType, + Dictionary, + Field, + FixedSizeList, + Float32, + Float64, + Int32, + Int64, + List, + Schema, + Struct, + Utf8, + util as arrowUtil, +} from "apache-arrow"; +import { typedArrayToArrowType } from "./arrow_type"; +import { sanitizeType } from "./sanitize"; + +type InferenceOptions = { + dictionaryEncodeStrings: boolean; + vectorColumns: Record; +}; + +/** + * Infer the Arrow schema represented by a set of records. + * + * This is the intentionally small interface to schema inference. The stateful + * details of combining partial type evidence are encapsulated below so callers + * only need to provide records, an optional schema, and inference options. + */ +export function inferSchema( + data: Array>, + schema: Schema | undefined, + options: InferenceOptions, +): Schema { + return new SchemaInferrer(schema, options).infer(data); +} + +class SchemaInferrer { + private readonly fields = new FieldTree(); + + constructor( + private readonly providedSchema: Schema | undefined, + private readonly options: InferenceOptions, + ) {} + + infer(data: Array>): Schema { + for (const [row, record] of data.entries()) { + for (const [path, value] of recordPathsAndValues(record)) { + this.observe(path, value, row); + } + } + + return this.providedSchema === undefined + ? new Schema(fieldsFromTree(this.fields)) + : new Schema(matchingFields(this.providedSchema.fields, this.fields)); + } + + private observe(path: string[], value: unknown, row: number): void { + const current = this.fields.get(path); + if (current === undefined) { + this.addField(path, value, row); + } else if (this.providedSchema === undefined) { + this.updateInferredField(path, value, row, current); + } + } + + private addField(path: string[], value: unknown, row: number): void { + if (this.providedSchema !== undefined) { + this.addSchemaField(this.providedSchema, path, row); + return; + } + + const evidence = + this.inferType(value, path) ?? DeferredTypeEvidence.from(value, row); + if (evidence === undefined) { + throw typeInferenceError(path, row); + } + + const conflict = this.fields.set( + path, + evidence, + (existing) => + existing instanceof DeferredTypeEvidence && existing.isOnlyNulls(), + ); + if (conflict !== undefined) { + throw branchConflictError(conflict, row, "Struct"); + } + } + + private addSchemaField(schema: Schema, path: string[], row: number): void { + const field = fieldAtPath(schema, path); + if (field === undefined) { + throw new Error( + `Found field not in schema: ${path.join(".")} at row ${row}`, + ); + } + + const conflict = this.fields.set(path, field.type); + if (conflict !== undefined) { + throw branchConflictError(conflict, row, "Struct"); + } + } + + private updateInferredField( + path: string[], + value: unknown, + row: number, + current: FieldNode, + ): void { + const newType = this.inferType(value, path); + const deferred = DeferredTypeEvidence.from(value, row); + + if (current instanceof FieldTree) { + if (deferred?.isOnlyNulls()) { + return; + } + throw schemaInferenceError( + path, + row, + "Struct", + describeEvidence(newType ?? deferred), + ); + } + + if (current instanceof DeferredTypeEvidence) { + this.resolveDeferredField(path, row, current, newType, deferred); + return; + } + + if (newType !== undefined) { + if (!inferredTypesEqual(current, newType)) { + throw schemaInferenceError( + path, + row, + describeEvidence(current), + describeEvidence(newType), + ); + } + return; + } + + if (deferred === undefined || !deferred.matches(current)) { + throw schemaInferenceError( + path, + row, + describeEvidence(current), + describeEvidence(deferred), + ); + } + } + + private resolveDeferredField( + path: string[], + row: number, + current: DeferredTypeEvidence, + newType: DataType | undefined, + deferred: DeferredTypeEvidence | undefined, + ): void { + if (newType !== undefined) { + if (!current.matches(newType)) { + throw schemaInferenceError( + path, + row, + current.describe(), + describeEvidence(newType), + ); + } + this.fields.set(path, newType); + return; + } + + if (deferred !== undefined) { + this.fields.set(path, current.merge(deferred)); + return; + } + + throw schemaInferenceError( + path, + row, + current.describe(), + describeEvidence(newType), + ); + } + + private inferType(value: unknown, path: string[]): DataType | undefined { + if (typeof value === "bigint") { + return new Int64(); + } + if (typeof value === "number") { + return new Float64(); + } + if (typeof value === "string") { + return this.options.dictionaryEncodeStrings + ? new Dictionary(new Utf8(), new Int32()) + : new Utf8(); + } + if (typeof value === "boolean") { + return new Bool(); + } + if (value instanceof Buffer) { + return new Binary(); + } + if (ArrayBuffer.isView(value) && !(value instanceof DataView)) { + const typedArray = typedArrayToArrowType(value); + return typedArray === undefined + ? undefined + : new FixedSizeList( + typedArray.length, + new Field("item", typedArray.elementType, true), + ); + } + if (!Array.isArray(value) || value.length === 0) { + return undefined; + } + + const configuredVector = + path.length === 1 ? this.options.vectorColumns[path[0]] : undefined; + if (configuredVector !== undefined) { + return new FixedSizeList( + value.length, + new Field("item", sanitizeType(configuredVector.type), true), + ); + } + + const itemType = this.inferArrayItemType(value, path); + if (itemType === undefined) { + return undefined; + } + + return nameSuggestsVectorColumn(path[path.length - 1]) + ? new FixedSizeList(value.length, new Field("item", new Float32(), true)) + : new List(new Field("item", itemType, true)); + } + + private inferArrayItemType( + values: unknown[], + path: string[], + ): DataType | undefined { + let itemType: DataType | undefined; + const deferredItems: unknown[] = []; + + for (const value of values) { + const candidate = this.inferType(value, path); + if (candidate === undefined) { + if (!isDeferredValue(value)) { + return undefined; + } + deferredItems.push(value); + } else if (itemType === undefined) { + itemType = candidate; + } else if (!inferredTypesEqual(itemType, candidate)) { + return undefined; + } + } + + if (itemType === undefined) { + return undefined; + } + return deferredItems.every((value) => + deferredValueMatchesType(value, itemType), + ) + ? itemType + : undefined; + } +} + +/** Nulls and empty/all-null lists that do not determine a type by themselves. */ +class DeferredTypeEvidence { + private constructor( + private readonly values: Array<{ value: unknown; row: number }>, + ) {} + + static from(value: unknown, row: number): DeferredTypeEvidence | undefined { + return isDeferredValue(value) + ? new DeferredTypeEvidence([{ value, row }]) + : undefined; + } + + isOnlyNulls(): boolean { + return this.values.every(({ value }) => value == null); + } + + matches(type: DataType): boolean { + return this.values.every(({ value }) => + deferredValueMatchesType(value, type), + ); + } + + merge(other: DeferredTypeEvidence): DeferredTypeEvidence { + return new DeferredTypeEvidence([...this.values, ...other.values]); + } + + describe(): string { + const list = this.values.find(({ value }) => Array.isArray(value)); + return list === undefined + ? "null" + : `List[${(list.value as unknown[]).length}]`; + } + + firstRow(): number { + return this.values[0].row; + } +} + +type FieldNode = DataType | DeferredTypeEvidence | FieldTree; +type LeafNode = Exclude; +type FieldConflict = { path: string[]; value: FieldNode }; + +/** Nested field state, kept separate from Arrow's eventual Struct types. */ +class FieldTree { + private readonly children = new Map(); + + get(path: string[]): FieldNode | undefined { + let current: FieldNode = this; + for (const part of path) { + if (!(current instanceof FieldTree)) { + return undefined; + } + const child = current.children.get(part); + if (child === undefined) { + return undefined; + } + current = child; + } + return current; + } + + set( + path: string[], + value: LeafNode, + canReplaceLeaf: (value: LeafNode) => boolean = () => false, + ): FieldConflict | undefined { + let branch: FieldTree = this; + for (const [index, part] of path.slice(0, -1).entries()) { + const child = branch.children.get(part); + if (child === undefined || (isLeaf(child) && canReplaceLeaf(child))) { + const nextBranch = new FieldTree(); + branch.children.set(part, nextBranch); + branch = nextBranch; + } else if (child instanceof FieldTree) { + branch = child; + } else { + return { path: path.slice(0, index + 1), value: child }; + } + } + + const name = path[path.length - 1]; + const current = branch.children.get(name); + if (current instanceof FieldTree) { + return { path, value: current }; + } + branch.children.set(name, value); + return undefined; + } + + entries(): IterableIterator<[string, FieldNode]> { + return this.children.entries(); + } + + has(name: string): boolean { + return this.children.has(name); + } +} + +function isLeaf(value: FieldNode): value is LeafNode { + return !(value instanceof FieldTree); +} + +function fieldsFromTree(tree: FieldTree, path: string[] = []): Field[] { + const fields: Field[] = []; + for (const [name, value] of tree.entries()) { + if (value instanceof FieldTree) { + fields.push( + new Field( + name, + new Struct(fieldsFromTree(value, [...path, name])), + true, + ), + ); + } else if (value instanceof DeferredTypeEvidence) { + throw typeInferenceError([...path, name], value.firstRow()); + } else { + fields.push(new Field(name, value, true)); + } + } + return fields; +} + +function matchingFields(fields: Field[], tree: FieldTree): Field[] { + const matches: Field[] = []; + for (const field of fields) { + if (!tree.has(field.name)) { + continue; + } + const value = tree.get([field.name]); + if (value instanceof FieldTree) { + const struct = field.type as Struct; + matches.push( + new Field( + field.name, + new Struct(matchingFields(struct.children, value)), + field.nullable, + ), + ); + } else { + matches.push(new Field(field.name, value as DataType, field.nullable)); + } + } + return matches; +} + +function* recordPathsAndValues( + record: Record, + path: string[] = [], +): Generator<[string[], unknown]> { + for (const [name, value] of Object.entries(record)) { + if (isRecord(value)) { + yield* recordPathsAndValues(value, [...path, name]); + } else if (value !== undefined) { + yield [[...path, name], value]; + } + } +} + +function isRecord(value: unknown): value is Record { + return ( + typeof value === "object" && + value !== null && + !Array.isArray(value) && + !(value instanceof RegExp) && + !(value instanceof Date) && + !(value instanceof Set) && + !(value instanceof Map) && + !(value instanceof Buffer) && + !ArrayBuffer.isView(value) + ); +} + +function fieldAtPath(schema: Schema, path: string[]): Field | undefined { + let fields = schema.fields; + let field: Field | undefined; + for (const [index, name] of path.entries()) { + field = fields.find((candidate) => candidate.name === name); + if (field === undefined || index === path.length - 1) { + return field; + } + if (!DataType.isStruct(field.type)) { + return undefined; + } + fields = field.type.children; + } + return field; +} + +function isDeferredValue(value: unknown): boolean { + return ( + value == null || (Array.isArray(value) && value.every(isDeferredValue)) + ); +} + +function deferredValueMatchesType(value: unknown, type: DataType): boolean { + if (value == null) { + return true; + } + if (!Array.isArray(value)) { + return false; + } + if (DataType.isList(type)) { + return value.every((item) => + deferredValueMatchesType(item, type.valueType), + ); + } + if (DataType.isFixedSizeList(type)) { + return ( + value.length === type.listSize && + value.every((item) => deferredValueMatchesType(item, type.valueType)) + ); + } + return false; +} + +function inferredTypesEqual(current: DataType, candidate: DataType): boolean { + if (DataType.isDictionary(current)) { + return ( + DataType.isDictionary(candidate) && + current.isOrdered === candidate.isOrdered && + inferredTypesEqual(current.indices, candidate.indices) && + inferredTypesEqual(current.dictionary, candidate.dictionary) + ); + } + if (DataType.isList(current)) { + return ( + DataType.isList(candidate) && + current.valueField.name === candidate.valueField.name && + current.valueField.nullable === candidate.valueField.nullable && + inferredTypesEqual(current.valueType, candidate.valueType) + ); + } + if (DataType.isFixedSizeList(current)) { + return ( + DataType.isFixedSizeList(candidate) && + current.listSize === candidate.listSize && + current.valueField.name === candidate.valueField.name && + current.valueField.nullable === candidate.valueField.nullable && + inferredTypesEqual(current.valueType, candidate.valueType) + ); + } + return arrowUtil.compareTypes(current, candidate); +} + +function describeEvidence( + evidence: DataType | DeferredTypeEvidence | undefined, +): string { + if (evidence === undefined) { + return "an unsupported value"; + } + return evidence instanceof DeferredTypeEvidence + ? evidence.describe() + : evidence.toString(); +} + +function branchConflictError( + conflict: FieldConflict, + row: number, + candidate: string, +): Error { + return schemaInferenceError( + conflict.path, + row, + conflict.value instanceof FieldTree + ? "Struct" + : describeEvidence(conflict.value), + candidate, + ); +} + +function schemaInferenceError( + path: string[], + row: number, + currentType: string, + newType: string, +): Error { + return new Error( + `Failed to infer schema for data. Previously inferred type ${currentType} ` + + `but found ${newType} for field ${path.join(".")} at row ${row}. ` + + "Consider providing an explicit schema.", + ); +} + +function typeInferenceError(path: string[], row: number): Error { + return new Error( + `Failed to infer data type for field ${path.join(".")} at row ${row}. ` + + "Consider providing an explicit schema.", + ); +} + +function nameSuggestsVectorColumn(name: string): boolean { + const normalized = name.toLowerCase(); + return normalized.includes("vector") || normalized.includes("embedding"); +} diff --git a/nodejs/lancedb/table.ts b/nodejs/lancedb/table.ts index 28603cca9..82f23e2f2 100644 --- a/nodejs/lancedb/table.ts +++ b/nodejs/lancedb/table.ts @@ -43,10 +43,12 @@ import { Table as _NativeTable, } from "./native"; import { + AutoQuery, FullTextQuery, Query, TakeQuery, VectorQuery, + createAutoQuery, instanceOfFullTextQuery, } from "./query"; import { sanitizeType } from "./sanitize"; @@ -523,7 +525,7 @@ export abstract class Table { query: string | IntoVector | MultiVector | FullTextQuery, queryType?: string, ftsColumns?: string | string[], - ): VectorQuery | Query; + ): VectorQuery | Query | AutoQuery; /** * Search the table with a given query vector. * @@ -975,10 +977,11 @@ export class LocalTable extends Table { return this.inner.display(); } - private async getEmbeddingFunctions(): Promise< - Map - > { - const schema = await this.schema(); + private async getEmbeddingFunctions( + inner: _NativeTable = this.inner, + ): Promise> { + const schemaBuf = await inner.schema(); + const schema = tableFromIPC(schemaBuf).schema; const registry = getRegistry(); return registry.parseFunctions(schema.metadata); } @@ -1160,7 +1163,7 @@ export class LocalTable extends Table { query: string | IntoVector | MultiVector | FullTextQuery, queryType: string = "auto", ftsColumns?: string | string[], - ): VectorQuery | Query { + ): VectorQuery | Query | AutoQuery { if (typeof query !== "string" && !instanceOfFullTextQuery(query)) { if (queryType === "fts") { throw new Error("Cannot perform full text search on a vector query"); @@ -1175,14 +1178,28 @@ export class LocalTable extends Table { }); } - // The query type is auto or vector - // fall back to full text search if no embedding functions are defined and the query is a string - if ( - queryType === "auto" && - (getRegistry().length() === 0 || instanceOfFullTextQuery(query)) - ) { - return this.query().fullTextSearch(query, { - columns: ftsColumns, + if (queryType === "auto") { + if (instanceOfFullTextQuery(query)) { + return this.query().fullTextSearch(query, { + columns: ftsColumns, + }); + } + + 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; + // 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); }); } diff --git a/nodejs/npm/darwin-arm64/package.json b/nodejs/npm/darwin-arm64/package.json index d89eb5ad1..ff6347c4d 100644 --- a/nodejs/npm/darwin-arm64/package.json +++ b/nodejs/npm/darwin-arm64/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-darwin-arm64", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.10", "os": ["darwin"], "cpu": ["arm64"], "main": "lancedb.darwin-arm64.node", diff --git a/nodejs/npm/linux-arm64-gnu/package.json b/nodejs/npm/linux-arm64-gnu/package.json index 1721c970e..ed99fee05 100644 --- a/nodejs/npm/linux-arm64-gnu/package.json +++ b/nodejs/npm/linux-arm64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-gnu", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.10", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-gnu.node", diff --git a/nodejs/npm/linux-arm64-musl/package.json b/nodejs/npm/linux-arm64-musl/package.json index 177804dce..5b215dcc0 100644 --- a/nodejs/npm/linux-arm64-musl/package.json +++ b/nodejs/npm/linux-arm64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-musl", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.10", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-musl.node", diff --git a/nodejs/npm/linux-x64-gnu/package.json b/nodejs/npm/linux-x64-gnu/package.json index 9af443583..e0f5a9f26 100644 --- a/nodejs/npm/linux-x64-gnu/package.json +++ b/nodejs/npm/linux-x64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-gnu", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.10", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-gnu.node", diff --git a/nodejs/npm/linux-x64-musl/package.json b/nodejs/npm/linux-x64-musl/package.json index c98aa3812..d42541707 100644 --- a/nodejs/npm/linux-x64-musl/package.json +++ b/nodejs/npm/linux-x64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-musl", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.10", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-musl.node", diff --git a/nodejs/npm/win32-arm64-msvc/package.json b/nodejs/npm/win32-arm64-msvc/package.json index 05b1086f5..496a40720 100644 --- a/nodejs/npm/win32-arm64-msvc/package.json +++ b/nodejs/npm/win32-arm64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-arm64-msvc", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.10", "os": [ "win32" ], diff --git a/nodejs/npm/win32-x64-msvc/package.json b/nodejs/npm/win32-x64-msvc/package.json index 50432bdea..734013343 100644 --- a/nodejs/npm/win32-x64-msvc/package.json +++ b/nodejs/npm/win32-x64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-x64-msvc", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.10", "os": ["win32"], "cpu": ["x64"], "main": "lancedb.win32-x64-msvc.node", diff --git a/nodejs/package-lock.json b/nodejs/package-lock.json index 219025534..732a7c01c 100644 --- a/nodejs/package-lock.json +++ b/nodejs/package-lock.json @@ -1,12 +1,12 @@ { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.10", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.10", "cpu": [ "x64", "arm64" diff --git a/nodejs/package.json b/nodejs/package.json index 63c51a799..262d4c4a7 100644 --- a/nodejs/package.json +++ b/nodejs/package.json @@ -11,7 +11,7 @@ "ann" ], "private": false, - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.10", "main": "dist/index.js", "exports": { ".": "./dist/index.js", diff --git a/nodejs/src/connection.rs b/nodejs/src/connection.rs index 44eb68f32..5cf676256 100644 --- a/nodejs/src/connection.rs +++ b/nodejs/src/connection.rs @@ -17,6 +17,7 @@ use lancedb::connection::{ConnectBuilder, Connection as LanceDBConnection, conne use lance_namespace::models::{ CreateNamespaceRequest, DescribeNamespaceRequest, DropNamespaceRequest, ListNamespacesRequest, + ListTablesRequest, }; use lancedb::ipc::{ipc_file_to_batches, ipc_file_to_schema}; @@ -36,6 +37,12 @@ pub struct ListNamespacesResponse { pub page_token: Option, } +#[napi(object)] +pub struct ListTablesResponse { + pub tables: Vec, + pub page_token: Option, +} + #[napi(object)] pub struct CreateNamespaceResponse { pub properties: Option>, @@ -206,6 +213,33 @@ impl Connection { op.execute().await.default_error() } + /// List a page of tables in the database. + #[napi(catch_unwind)] + pub async fn list_tables( + &self, + namespace_path: Option>, + page_token: Option, + limit: Option, + ) -> napi::Result { + let request = ListTablesRequest { + // The root namespace is an empty path, not an absent one: a namespace-backed + // database rejects a request that names no namespace. + id: Some(namespace_path.unwrap_or_default()), + page_token, + limit: limit.map(|limit| i32::try_from(limit).unwrap_or(i32::MAX)), + ..Default::default() + }; + let response = self + .get_inner()? + .list_tables(request) + .await + .default_error()?; + Ok(ListTablesResponse { + tables: response.tables, + page_token: response.page_token, + }) + } + /// Create table from a Apache Arrow IPC (file) buffer. /// /// Parameters: diff --git a/nodejs/src/table.rs b/nodejs/src/table.rs index 9d60b2056..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( @@ -554,6 +561,12 @@ impl Table { .default_error() } + #[napi(catch_unwind)] + pub async fn checkout_current(&self) -> napi::Result { + let table = self.inner_ref()?.checkout_current().await.default_error()?; + Ok(Self::new(table)) + } + #[napi(catch_unwind)] pub async fn checkout(&self, version: i64) -> napi::Result<()> { self.inner_ref()? diff --git a/python/Cargo.toml b/python/Cargo.toml index 2e4f8996f..3a8a05522 100644 --- a/python/Cargo.toml +++ b/python/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb-python" -version = "0.38.0-beta.6" +version = "0.38.0-beta.10" publish = false edition.workspace = true description = "Python bindings for LanceDB" diff --git a/python/pyproject.toml b/python/pyproject.toml index fad3b1001..22a41a8a9 100644 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -101,9 +101,12 @@ azure = ["adlfs>=2024.2.0"] [tool.maturin] python-source = "python" module-name = "lancedb._lancedb" +# uv installs the project as an editable package before `uv run`, so keep that +# bootstrap build consistent with `maturin develop`. +editable-profile = "dev" [build-system] -requires = ["maturin>=1.9.4"] +requires = ["maturin>=1.10"] build-backend = "maturin" [tool.ruff.lint] diff --git a/python/python/lancedb/__init__.py b/python/python/lancedb/__init__.py index df8950686..8cb85a3ed 100644 --- a/python/python/lancedb/__init__.py +++ b/python/python/lancedb/__init__.py @@ -179,6 +179,18 @@ def connect( ... }, ... ) + For Azure Blob Storage, credentials can be passed directly without setting + environment variables: + + >>> azure_storage_options = { + ... "account_name": "some-account", + ... "account_key": "some-key", + ... } + >>> db = lancedb.connect( # doctest: +SKIP + ... "az://my-container/my-database", + ... storage_options=azure_storage_options, + ... ) + For tests and temporary data, use an in-memory database: >>> db = lancedb.connect("memory://") @@ -465,6 +477,10 @@ async def connect_async( -------- >>> import lancedb + >>> azure_storage_options = { + ... "account_name": "some-account", + ... "account_key": "some-key", + ... } >>> async def doctest_example(): ... # For a local directory, provide a path to the database ... db = await lancedb.connect_async("~/.lancedb") @@ -472,6 +488,11 @@ async def connect_async( ... db = await lancedb.connect_async("s3://my-bucket/lancedb", ... storage_options={ ... "aws_access_key_id": "***"}) + ... # Azure credentials can also be passed directly + ... db = await lancedb.connect_async( + ... "az://my-container/my-database", + ... storage_options=azure_storage_options, + ... ) ... # For tests and temporary data, use an in-memory database ... db = await lancedb.connect_async("memory://") ... # Connect to LanceDB cloud diff --git a/python/python/lancedb/_lancedb.pyi b/python/python/lancedb/_lancedb.pyi index c9571f03a..f11182f5d 100644 --- a/python/python/lancedb/_lancedb.pyi +++ b/python/python/lancedb/_lancedb.pyi @@ -283,6 +283,7 @@ class Table: mode: Literal["append", "overwrite"], progress: Optional[Any] = None, write_parallelism: Optional[int] = None, + allow_external_blob_outside_bases: bool = False, ) -> AddResult: ... async def update( self, updates: Dict[str, str], where: Optional[str] diff --git a/python/python/lancedb/functions.py b/python/python/lancedb/functions.py index c4ff9a9b3..237ed73c2 100644 --- a/python/python/lancedb/functions.py +++ b/python/python/lancedb/functions.py @@ -4,7 +4,7 @@ """Canonical Function values exchanged with LanceDB Enterprise services. These immutable models contain client/wire state only. Catalog persistence, -environment bake, secret resolution, and execution are owned by Sophon. +environment bake, and execution are owned by Sophon. ``RefreshColumnResult`` is also the backend-neutral result of a local expression-backed refresh job. """ @@ -12,17 +12,19 @@ expression-backed refresh job. from __future__ import annotations import ast +import builtins import base64 import functools import hashlib +import importlib import inspect +import symtable import json import math import re import sys import textwrap import types -import uuid from collections.abc import Mapping from datetime import date, datetime from typing import ( @@ -226,7 +228,7 @@ class PythonEnvironmentSpec(_RemoteValue): class PythonRuntimeSpec(_RemoteValue): - """Remote runtime definition with non-secret environment values. + """Remote runtime definition with environment values. V1 supports ``kind="python"``. Newer runtime kinds remain readable, while their unknown payload fields are intentionally not retained by the client. @@ -265,7 +267,6 @@ class FunctionVersion(_RemoteValue): runtime: PythonRuntimeSpec runtime_digest: str environment_digest: str - required_secrets: tuple[str, ...] = () created_at: str def __call__(self, **inputs: Any) -> FunctionApplication: @@ -273,7 +274,7 @@ class FunctionVersion(_RemoteValue): Every input must be a direct [lancedb.col][lancedb.expr.col] reference. The returned application is immutable and retains a - named-struct output as one sibling group, so every row's sibling values + named-struct output as one binding, so every row's sibling values come from one logical Function evaluation. Map result fields to table columns with [FunctionApplication.rename][lancedb.functions.FunctionApplication.rename], @@ -323,22 +324,16 @@ class FunctionVersion(_RemoteValue): function=FunctionVersionRef(name=self.name, version=self.version), inputs=tuple(bindings), output=self.signature.output, - group_id=f"fg_{uuid.uuid4().hex}", ) class FunctionRegistrationRequest(_RemoteValue): - """Stable remote registration envelope produced by :func:`udf`. - - Only secret names are represented. Secret values are resolved inside the - remote service and have no client request field. - """ + """Stable remote registration envelope produced by :func:`udf`.""" name: str artifact: FunctionArtifactRequest signature: FunctionSignature runtime: PythonRuntimeSpec - required_secrets: tuple[str, ...] = () class FunctionVersionRef(_OpenRemoteValue): @@ -367,7 +362,7 @@ class ApplicationInput(_OpenRemoteValue): class FunctionApplication(_OpenRemoteValue): """Immutable pre-declaration application of an exact Function version. - A named-struct output remains one grouped application through table + A named-struct output remains one application through table declaration and execution. [FunctionApplication.rename][lancedb.functions.FunctionApplication.rename] records the result-field to table-column mapping without splitting sibling @@ -377,7 +372,6 @@ class FunctionApplication(_OpenRemoteValue): function: FunctionVersionRef inputs: tuple[ApplicationInput, ...] output: FunctionOutput - group_id: str columns: Mapping[str, str] = Field(default_factory=dict) def _known_dict(self) -> dict[str, Any]: @@ -449,12 +443,10 @@ class OutputMapping(_RemoteValue): class FunctionBinding(_RemoteValue): - """Immutable grouped binding persisted by the Enterprise table service.""" + """Immutable Function binding persisted by the Enterprise table service.""" binding_id: str - revision: _UInt64 function: FunctionVersionRef - group_id: str inputs: tuple[InputBinding, ...] outputs: tuple[OutputMapping, ...] input_schema: Optional[Mapping[str, Any]] = None @@ -486,62 +478,60 @@ class RefreshColumnResult(_RemoteValue): _FUNCTION_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_.-]*$") -_SECRET_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") + + +_GRAMMAR_PRIMITIVES = ( + (pa.bool_(), "bool"), + (pa.int8(), "int8"), + (pa.int16(), "int16"), + (pa.int32(), "int32"), + (pa.int64(), "int64"), + (pa.uint8(), "uint8"), + (pa.uint16(), "uint16"), + (pa.uint32(), "uint32"), + (pa.uint64(), "uint64"), + (pa.float16(), "float16"), + (pa.float32(), "float32"), + (pa.float64(), "float64"), + (pa.string(), "utf8"), + (pa.binary(), "binary"), + (pa.date32(), "date32"), + (pa.date64(), "date64"), +) def _canonical_arrow_type(data_type: pa.DataType) -> str: - primitive_types = ( - (pa.bool_(), "bool"), - (pa.int8(), "int8"), - (pa.int16(), "int16"), - (pa.int32(), "int32"), - (pa.int64(), "int64"), - (pa.uint8(), "uint8"), - (pa.uint16(), "uint16"), - (pa.uint32(), "uint32"), - (pa.uint64(), "uint64"), - (pa.float16(), "float16"), - (pa.float32(), "float32"), - (pa.float64(), "float64"), - (pa.string(), "utf8"), - (pa.large_utf8(), "large_utf8"), - (pa.binary(), "binary"), - (pa.large_binary(), "large_binary"), - (pa.date32(), "date32"), - (pa.date64(), "date64"), - ) - for candidate, name in primitive_types: + """The server's V1 Function type grammar. Anything outside it is rejected + here rather than at registration.""" + for candidate, name in _GRAMMAR_PRIMITIVES: if data_type == candidate: return name - if pa.types.is_fixed_size_binary(data_type): - return f"fixed_size_binary[{data_type.byte_width}]" - if pa.types.is_list(data_type): - return f"list<{_canonical_arrow_type(data_type.value_type)}>" - if pa.types.is_large_list(data_type): - return f"large_list<{_canonical_arrow_type(data_type.value_type)}>" - if pa.types.is_fixed_size_list(data_type): + if pa.types.is_list(data_type) or pa.types.is_large_list(data_type): + prefix = "list" if pa.types.is_list(data_type) else "large_list" + return f"{prefix}<{_canonical_list_item(data_type)}>" + if pa.types.is_fixed_size_list(data_type) and data_type.list_size > 0: return ( - f"fixed_size_list<{_canonical_arrow_type(data_type.value_type)}>" - f"[{data_type.list_size}]" + f"fixed_size_list<{_canonical_list_item(data_type)}, {data_type.list_size}>" ) - if pa.types.is_struct(data_type): - fields = ",".join( - f"{field.name}:{_canonical_arrow_type(field.type)}" for field in data_type - ) - return f"struct<{fields}>" - if pa.types.is_timestamp(data_type): - timezone = f",tz={data_type.tz}" if data_type.tz is not None else "" - return f"timestamp[{data_type.unit}{timezone}]" - if pa.types.is_time32(data_type) or pa.types.is_time64(data_type): - return f"time[{data_type.unit}]" - if pa.types.is_duration(data_type): - return f"duration[{data_type.unit}]" - if pa.types.is_decimal(data_type): - bit_width = data_type.bit_width - return f"decimal{bit_width}({data_type.precision},{data_type.scale})" raise TypeError(f"unsupported Arrow type for Function signature: {data_type}") +def _canonical_list_item(data_type: pa.DataType) -> str: + """The grammar names only the item type; it always means a non-nullable + child called `item`, so any other child metadata cannot be represented.""" + child = data_type.value_field + if child.name != "item" or child.nullable or child.metadata: + raise TypeError( + "unsupported Arrow type for Function signature: list items must be a " + f"non-nullable field named 'item', got {child}" + ) + return _canonical_arrow_type(child.type) + + +def _list_of(item: pa.DataType) -> pa.DataType: + return pa.list_(pa.field("item", item, nullable=False)) + + def _annotation_type(annotation: Any) -> tuple[pa.DataType, bool]: nullable = False origin = get_origin(annotation) @@ -589,7 +579,7 @@ def _annotation_type(annotation: Any) -> tuple[pa.DataType, bool]: value_type, value_nullable = _annotation_type(arguments[0]) if value_nullable: raise TypeError("nullable Function list elements are not supported") - return pa.list_(value_type), nullable + return _list_of(value_type), nullable raise TypeError(f"unsupported Function annotation: {annotation!r}") @@ -736,6 +726,104 @@ def _literal_source(value: Any) -> str: ) +_DYNAMIC_NAMESPACE_ACCESS = frozenset( + {"globals", "locals", "vars", "eval", "exec", "compile", "__import__"} +) +# Modules that hand out namespaces (`sys.modules`, `builtins`, importers, +# introspection). The artifact's module namespace holds only the names it was +# packaged with, so reaching around it cannot be represented. +_NAMESPACE_MODULES = frozenset( + {"sys", "builtins", "importlib", "inspect", "gc", "ctypes", "types"} +) + + +def _namespace_acquisition( + definition: ast.FunctionDef, references: set[str] +) -> list[str]: + found = set(references & _DYNAMIC_NAMESPACE_ACCESS) + for node in ast.walk(definition): + if isinstance(node, ast.Import): + found.update( + alias.name + for alias in node.names + if alias.name.split(".")[0] in _NAMESPACE_MODULES + ) + elif isinstance(node, ast.ImportFrom) and node.module: + if node.module.split(".")[0] in _NAMESPACE_MODULES: + found.add(node.module) + return sorted(found) + + +def _module_references(module_source: str) -> set[str]: + """Names any scope in `module_source` binds or loads at module scope. + Python's own scope analysis on the exact text that ships: free variables + belong to an enclosing scope inside the function, and postponed + annotations are not runtime loads.""" + + def visit(table: symtable.SymbolTable, found: set[str]) -> None: + for symbol in table.get_symbols(): + if symbol.is_global() and ( + symbol.is_referenced() or symbol.is_declared_global() + ): + found.add(symbol.get_name()) + for child in table.get_children(): + visit(child, found) + + found: set[str] = set() + for table in symtable.symtable(module_source, "", "exec").get_children(): + visit(table, found) + return found + + +def _global_source(name: str, value: Any) -> str: + """One module-level line that rebinds `name` to `value` in the artifact: + an import for modules and importable classes/functions, a literal otherwise.""" + if isinstance(value, types.ModuleType): + if value.__name__.split(".")[0] in _NAMESPACE_MODULES: + raise ValueError( + f"@udf cannot package dynamic namespace access: {value.__name__!r}" + ) + try: + imported = importlib.import_module(value.__name__) + except ImportError: + imported = None + if imported is not value: + raise TypeError( + f"Function source references module {name!r} that does not import " + f"as {value.__name__!r}" + ) + return f"import {value.__name__} as {name}" + module_name = getattr(value, "__module__", None) + qualname = getattr(value, "__qualname__", None) + if ( + isinstance(module_name, str) + and isinstance(qualname, str) + and module_name != "__main__" + and "." not in qualname + and "<" not in qualname + ): + try: + imported = getattr(importlib.import_module(module_name), qualname) + except (ImportError, AttributeError): + imported = None + if imported is value: + return f"from {module_name} import {qualname} as {name}" + return f"{name} = {_literal_source(value)}" + + +def _is_recursive_reference(function: Callable[..., Any], name: str) -> bool: + """`name` inside the body means the function itself unless the module has + since bound it to something else.""" + if name != function.__name__: + return False + bound = function.__globals__.get(name, function) + if bound is function: + return True + # The decorator's own result is the one wrapper known to call `function` + # unchanged; any other binding may behave differently from a self-call. + return type(bound) is UdfDefinition and bound._function is function + + def _package_source(function: Callable[..., Any]) -> bytes: if not inspect.isfunction(function) or inspect.iscoroutinefunction(function): raise TypeError("@udf requires a synchronous Python function") @@ -760,23 +848,46 @@ def _package_source(function: Callable[..., Any]) -> bytes: closure = inspect.getclosurevars(function) if closure.nonlocals: raise ValueError("@udf cannot package functions that capture closure values") - if closure.unbound: - raise ValueError( - f"@udf source contains unresolved global names: {sorted(closure.unbound)!r}" - ) - globals_source = [] - for name, value in sorted(closure.globals.items()): - if isinstance(value, types.ModuleType): - globals_source.append(f"import {value.__name__} as {name}") - else: - globals_source.append(f"{name} = {_literal_source(value)}") - function_source = ast.unparse(definition) - parts = ["from __future__ import annotations"] + module_header = "from __future__ import annotations" + references = _module_references(f"{module_header}\n\n{function_source}\n") + dynamic = _namespace_acquisition(definition, references) + if dynamic: + raise ValueError(f"@udf cannot package dynamic namespace access: {dynamic!r}") + # Resolve every module-scope reference the way the interpreter would: the + # function's own globals first (a module global may shadow a builtin, and + # nested scopes are not visible to getclosurevars), then its builtins. + # The artifact runs under the standard builtins; only the exact mapping is + # provably equivalent (a subclass or copy can change lookups and hooks). + if function.__builtins__ is not vars(builtins): + raise ValueError("@udf cannot package a non-standard builtins environment") + globals_source = [] + unresolved = [] + for name in sorted(references): + if name == function.__name__: + if not _is_recursive_reference(function, name): + raise ValueError( + f"@udf cannot package {name!r}: the module binds that name to " + "another value, which the artifact's own definition would shadow" + ) + continue + if name in function.__globals__: + globals_source.append(_global_source(name, function.__globals__[name])) + elif hasattr(builtins, name): + pass + else: + unresolved.append(name) + if unresolved: + raise ValueError( + f"@udf source contains unresolved global names: {unresolved!r}" + ) + + parts = [module_header] if globals_source: parts.extend(["", *globals_source]) parts.extend(["", function_source, ""]) - return "\n".join(parts).encode("utf-8") + packaged = "\n".join(parts) + return packaged.encode("utf-8") class UdfDefinition: @@ -797,7 +908,6 @@ class UdfDefinition: output_schema: Optional[pa.DataType | pa.Field | pa.Schema], pip: tuple[str, ...], env: Mapping[str, str], - secrets: tuple[str, ...], python_version: Optional[str], ): function_name = name or function.__name__ @@ -812,18 +922,6 @@ class UdfDefinition: for key, value in environment.items() ): raise TypeError("Function env keys and values must be strings") - required_secrets = tuple(sorted(set(secrets))) - invalid_secrets = [ - secret for secret in required_secrets if not _SECRET_NAME.fullmatch(secret) - ] - if invalid_secrets: - raise ValueError(f"invalid Function secret names: {invalid_secrets!r}") - overlap = set(environment) & set(required_secrets) - if overlap: - raise ValueError( - f"Function env and secret names must be disjoint: {sorted(overlap)!r}" - ) - signature = _infer_signature(function, input_schema, output_schema) source = _package_source(function) digest = f"sha256:{hashlib.sha256(source).hexdigest()}" @@ -852,7 +950,6 @@ class UdfDefinition: ), signature=signature, runtime=runtime, - required_secrets=required_secrets, ) functools.update_wrapper(self, function) @@ -878,7 +975,6 @@ def udf( output_schema: Optional[pa.DataType | pa.Field | pa.Schema] = None, pip: tuple[str, ...] | list[str] = (), env: Optional[Mapping[str, str]] = None, - secrets: tuple[str, ...] | list[str] = (), python_version: Optional[str] = None, ) -> Callable[[Callable[..., Any]], UdfDefinition]: ... @@ -891,7 +987,6 @@ def udf( output_schema: Optional[pa.DataType | pa.Field | pa.Schema] = None, pip: tuple[str, ...] | list[str] = (), env: Optional[Mapping[str, str]] = None, - secrets: tuple[str, ...] | list[str] = (), python_version: Optional[str] = None, ): """Prepare a scalar Python callable for remote Function registration. @@ -916,13 +1011,17 @@ def udf( pip : sequence of str, optional Pip requirements for the remote environment. env : mapping of str to str, optional - Non-secret environment variables. Use ``secrets`` for credentials. - secrets : sequence of str, optional - Names of secrets resolved by the remote service. Secret values are not - accepted by this API or included in the registration request. + Environment variables included in the Function definition. python_version : str, optional Remote Python major/minor version. Defaults to the client version. + The packaged artifact is a snapshot: the function source plus exactly + the module-level names it references (modules as imports, importable + classes and functions as imports, literals inline). Code that reaches the + module namespace another way -- ``globals()``/``eval``, ``sys.modules``, + ``builtins`` -- is rejected where it can be seen and otherwise + unsupported; closures and a non-standard ``__builtins__`` are rejected. + Returns ------- UdfDefinition @@ -934,7 +1033,7 @@ def udf( Examples -------- >>> from lancedb import udf - >>> @udf(pip=["numpy==2.2.0"], secrets=["MODEL_TOKEN"]) + >>> @udf(pip=["numpy==2.2.0"]) ... def score(value: float) -> float: ... return value * 2 >>> score(1.5) @@ -949,7 +1048,6 @@ def udf( output_schema=output_schema, pip=tuple(pip), env={} if env is None else env, - secrets=tuple(secrets), python_version=python_version, ) diff --git a/python/python/lancedb/permutation.py b/python/python/lancedb/permutation.py index 5d7a685ef..8287ef2fe 100644 --- a/python/python/lancedb/permutation.py +++ b/python/python/lancedb/permutation.py @@ -391,6 +391,15 @@ def _table_to_pickle_state(table: Table) -> dict[str, Any]: } +def _drop_base_version(permutation_data: pa.Table) -> pa.Table: + """Strip the recorded base version so the reader leaves the base table unpinned.""" + metadata = dict(permutation_data.schema.metadata or {}) + if metadata.pop(b"base_version", None) is None: + return permutation_data + metadata.pop(b"base_branch", None) + return permutation_data.replace_schema_metadata(metadata) + + def _table_from_pickle_state(state: dict[str, Any]) -> Table: from . import connect @@ -679,11 +688,15 @@ class Permutation: from . import connect connection_factory = state["connection_factory"] + rebuilt_base = False if connection_factory is not None: base_table = connection_factory(state["base_table_name"]) elif "base_table_state" in state: - base_table = _table_from_pickle_state(state["base_table_state"]) + base_state = state["base_table_state"] + rebuilt_base = base_state["kind"] == "memory" + base_table = _table_from_pickle_state(base_state) elif "base_table_data" in state: + rebuilt_base = True # In-memory base table inlined into the pickle; rebuild the same # way we rebuild the in-memory permutation table. mem_db = connect("memory://") @@ -701,11 +714,14 @@ class Permutation: ) permutation_table: Optional[Table] = None - if state["permutation_data"] is not None: + permutation_data = state["permutation_data"] + if permutation_data is not None: + if rebuilt_base: + # The base table was materialized from Arrow, so it is a fresh + # single-version dataset and the recorded pin cannot resolve on it. + permutation_data = _drop_base_version(permutation_data) mem_db = connect("memory://") - permutation_table = mem_db.create_table( - "permutation", state["permutation_data"] - ) + permutation_table = mem_db.create_table("permutation", permutation_data) self.base_table = base_table self.permutation_table = permutation_table diff --git a/python/python/lancedb/remote/table.py b/python/python/lancedb/remote/table.py index 41394ed71..d0bf9f67a 100644 --- a/python/python/lancedb/remote/table.py +++ b/python/python/lancedb/remote/table.py @@ -610,6 +610,7 @@ class RemoteTable(Table): fill_value: float = 0.0, progress: Optional[Union[bool, Callable, Any]] = None, write_parallelism: Optional[int] = None, + allow_external_blob_outside_bases: bool = False, ) -> AddResult: """Add more data to the [Table][lancedb.table.Table]. @@ -642,6 +643,8 @@ class RemoteTable(Table): data in flight. Defaults to an estimate based on the data size, capped at the number of CPU cores. Lower this if bulk ingestion is using too much memory. + allow_external_blob_outside_bases: bool, default False + Not supported on LanceDB Cloud. Setting this raises. Returns ------- @@ -658,6 +661,7 @@ class RemoteTable(Table): fill_value=fill_value, progress=progress, write_parallelism=write_parallelism, + allow_external_blob_outside_bases=allow_external_blob_outside_bases, ) ) finally: diff --git a/python/python/lancedb/streaming.py b/python/python/lancedb/streaming.py index c54940ab7..76b3702a1 100644 --- a/python/python/lancedb/streaming.py +++ b/python/python/lancedb/streaming.py @@ -19,6 +19,7 @@ above. """ import ctypes +import heapq import logging import os import random @@ -29,17 +30,18 @@ from collections import deque from concurrent.futures import ThreadPoolExecutor from copy import deepcopy from multiprocessing import RawArray -from typing import Any, Callable, cast, Iterator, Literal, Optional, Union +from typing import Any, Callable, cast, Iterator, Literal, NamedTuple, Optional, Union import pyarrow as pa import pyarrow.compute as pc import torch -from torch.utils.data import IterableDataset, get_worker_info +from torch.utils.data import DataLoader, IterableDataset, get_worker_info from .permutation import ( Permutation, Transforms, permutation_builder, + _drop_base_version, _table_from_pickle_state, _table_to_pickle_state, ) @@ -55,6 +57,155 @@ DEFAULT_READ_BATCH_SIZE = 64 DEFAULT_PREFETCH_BATCHES = 4 +class _WorkerSample(NamedTuple): + data: Any + dataset: "StreamingDataset" + + +class _WorkerBatch(NamedTuple): + data: Any + state: dict + + +class _ConsumerIteratorLease(NamedTuple): + owner_token: int + owner_thread: int + + +class _CheckpointCollate: + """Attach the worker's post-fetch state to a collated batch.""" + + def __init__(self, collate_fn: Callable): + self._collate_fn = collate_fn + + def __call__(self, samples): + try: + if isinstance(samples, list): + if not samples: + return _WorkerBatch(self._collate_fn(samples), {}) + worker_samples = samples + data = self._collate_fn([sample.data for sample in worker_samples]) + dataset = worker_samples[-1].dataset + else: + data = self._collate_fn(samples.data) + dataset = samples.dataset + except StopIteration as exc: + raise RuntimeError( + "collate_fn raised StopIteration before returning a batch" + ) from exc + return _WorkerBatch(data, dataset._checkpoint_snapshot()) + + +class _StreamingDatasetAdapter(IterableDataset): + """Yield private sample wrappers for :class:`StreamingDataLoader`.""" + + def __init__(self, dataset: "StreamingDataset"): + super().__init__() + self.dataset = dataset + + def __iter__(self): + for sample in self.dataset._iter(consumer_checkpoint_transport=True): + yield _WorkerSample(sample, self.dataset) + + def __getattr__(self, name): + dataset = self.__dict__.get("dataset") + if dataset is None: + raise AttributeError(name) + return getattr(dataset, name) + + +class _ConsumerCommitIterator: + def __init__( + self, + iterator, + dataset: "StreamingDataset", + *, + owner_token: int, + require_uniform: bool, + ): + self._iterator = iterator + self._dataset = dataset + self._owner_token = owner_token + self._require_uniform = require_uniform + self._released = False + self._terminal = False + + def __iter__(self): + return self + + def __next__(self): + if self._terminal: + raise StopIteration + try: + batch = next(self._iterator) + except StopIteration: + self._terminal = True + self._release() + raise + except BaseException as exc: + self._dataset._invalidate_checkpoint( + f"a DataLoader batch failed before it was returned: {exc}" + ) + raise + try: + if not isinstance(batch, _WorkerBatch): + raise RuntimeError( + "StreamingDataLoader did not receive worker checkpoint metadata" + ) + self._dataset._commit_worker_state( + batch.state, require_uniform=self._require_uniform + ) + return batch.data + except BaseException as exc: + self._dataset._invalidate_checkpoint( + f"a DataLoader batch failed before it was returned: {exc}" + ) + raise + + def _release(self) -> None: + if self.__dict__.get("_released", True): + return + self._released = True + dataset = self.__dict__.get("_dataset") + if dataset is not None: + dataset._release_consumer_iterator(self._owner_token) + + def _shutdown_workers(self): + if self.__dict__.get("_released", True): + return None + self._terminal = True + iterator = self.__dict__.get("_iterator") + shutdown = getattr(iterator, "_shutdown_workers", None) + try: + if shutdown is not None: + shutdown() + else: + fetcher = getattr(iterator, "_dataset_fetcher", None) + dataset_iterator = getattr(fetcher, "dataset_iter", None) + close = getattr(dataset_iterator, "close", None) + if close is None: + raise RuntimeError( + "StreamingDataLoader could not close its inner iterator" + ) + close() + except BaseException as exc: + self._dataset._invalidate_checkpoint( + f"a DataLoader iterator could not be shut down safely: {exc}" + ) + raise + else: + self._release() + + def __del__(self): + try: + self._shutdown_workers() + except BaseException: + pass + + def __getattr__(self, name): + return getattr(self._iterator, name) + + class StreamingDataset(IterableDataset): """An elastic, resumable PyTorch IterableDataset backed by a LanceDB table. @@ -384,6 +535,22 @@ class StreamingDataset(IterableDataset): # rows_skipped] self._worker_stats: RawArray = RawArray(ctypes.c_int64, 8) + # A standard multi-process DataLoader cannot report which prefetched + # batches were actually returned to its consumer. Workers set this + # shared flag so state_dict() can reject a stale parent checkpoint + # unless StreamingDataLoader installed the consumer-commit transport. + self._untracked_worker_iteration: RawArray = RawArray(ctypes.c_int64, 1) + + # Parent-side checkpoint lifecycle. A failed DataLoader task creates + # a permanent hole in that iterator's delivery stream, while a + # multi-worker checkpoint is safe to restore only after all splits + # reach the same logical step boundary. + self._checkpoint_invalid_reason: Optional[str] = None + self._consumer_checkpoint_requires_uniform = False + self._consumer_iterator_lock = threading.Lock() + self._consumer_iterator_generation = 0 + self._consumer_iterator_lease: Optional[_ConsumerIteratorLease] = None + # Cumulative bytes of Arrow buffer data fetched across all iterations. self._bytes_loaded: int = 0 # Cumulative seconds spent in LanceDB I/O and in transform functions. @@ -396,6 +563,10 @@ class StreamingDataset(IterableDataset): # step boundaries all splits have consumed this many samples, so a # single scalar captures the topology-independent checkpoint state. self._resume_offset: int = 0 + # Exact yielded-sample counts for splits this process has advanced. + # Missing entries use _resume_offset, which remains the lower-bound + # checkpoint inherited from an earlier uniform/global state. + self._resume_samples: dict[int, int] = {} # Permutation position each split has consumed through, keyed by # global split index. Equal to _resume_offset for every split unless # on_transform_error skipped rows, in which case skipped positions @@ -521,11 +692,45 @@ class StreamingDataset(IterableDataset): return self._rank_splits[start : start + splits_per_worker] def __iter__(self) -> Iterator[dict[str, Any]]: + return self._iter() + + def _iter( + self, *, consumer_checkpoint_transport: bool = False + ) -> Iterator[dict[str, Any]]: + owner_token = None + previous_lease = self._consumer_iterator_lease + if consumer_checkpoint_transport: + if not self._consumer_iterator_active: + raise RuntimeError( + "StreamingDataLoader worker transport requires an active " + "parent iterator reservation" + ) + else: + try: + owner_token = self._acquire_consumer_iterator() + except BaseException: + self._release_consumer_iterator_after_failed_acquire(previous_lease) + raise + try: + yield from self._iter_owned( + consumer_checkpoint_transport=consumer_checkpoint_transport + ) + finally: + if owner_token is not None: + self._release_consumer_iterator(owner_token) + + def _iter_owned( + self, *, consumer_checkpoint_transport: bool + ) -> Iterator[dict[str, Any]]: if self._raw_batches_ref is not None: raise RuntimeError( "StreamingDataset does not support concurrent iteration. " "Only one active iterator per dataset instance is allowed." ) + real_worker = get_worker_info() is not None + if real_worker and not consumer_checkpoint_transport: + self._untracked_worker_iteration[0] = 1 + my_splits = self._resolve_my_splits() if not my_splits: return @@ -533,6 +738,7 @@ class StreamingDataset(IterableDataset): # Set identity transform on each Permutation so __getitems__ returns # the raw RecordBatch. Stage 2 applies the real transform. permutations: list[Permutation] = [] + initial_samples: list[int] = [] initial_positions: list[int] = [] for split_idx in my_splits: perm = Permutation.from_tables( @@ -541,21 +747,22 @@ class StreamingDataset(IterableDataset): if self._columns is not None: perm = perm.select_columns(self._columns) perm = perm.with_transform(Transforms.arrow2arrow) + sample_count = self._resume_samples.get(split_idx, self._resume_offset) # Both modes resume from absolute permutation positions. Packing # stores them separately because it also checkpoints partial blocks. start_pos = ( self._pack_consumed[split_idx] if self._pack_sequences is not None - else self._resume_positions.get(split_idx, self._resume_offset) + else self._resume_positions.get(split_idx, sample_count) ) if start_pos > 0: perm = perm.with_skip(start_pos) + initial_samples.append(sample_count) initial_positions.append(start_pos) permutations.append(perm) n = len(permutations) split_sizes = [perm.num_rows for perm in permutations] - initial_offset = self._resume_offset local_consumed = [0] * n # Permutation position each split has consumed through (absolute, # i.e. counted from the start of the unskipped split). Runs ahead of @@ -853,6 +1060,27 @@ class StreamingDataset(IterableDataset): for i in range(n): _fill_io(i) + def _yield_row(i: int): + pos, row = cooked[i].popleft() + # Surface any completed prefetched failure before the + # current row becomes durable checkpoint progress. + _advance(i) + local_consumed[i] += 1 + pos_consumed[i] = pos + 1 + split_idx = my_splits[i] + self._resume_samples[split_idx] = ( + initial_samples[i] + local_consumed[i] + ) + self._resume_positions[split_idx] = pos_consumed[i] + return row + + def _update_progress_stats() -> None: + if not real_worker: + self._resume_offset = min( + initial_samples[j] + local_consumed[j] for j in range(n) + ) + _update_stats() + if self._pack_sequences is not None: first_count = pack_blocks_emitted[my_splits[0]] if any( @@ -878,12 +1106,38 @@ class StreamingDataset(IterableDataset): tokens.extend([pad_id] * (pack_len - len(tokens))) block = _emit_block(i) pack_blocks_emitted[my_splits[i]] += 1 + # Checkpoint state must advance before yielding so + # StreamingDataLoader can attach the exact state to + # the batch it transports to the parent process. + _commit_pack_state() if i == n - 1: - _commit_pack_state() _update_stats() yield block return + # A checkpoint taken between round-robin split turns has + # non-uniform counts. Resume lagging splits first so the + # exact canonical sequence continues without replaying + # already-consumed rows. + if len(set(initial_samples)) > 1: + catch_up_to = max(initial_samples) + pending = [ + (initial_samples[i], my_splits[i], i) + for i in range(n) + if initial_samples[i] < catch_up_to + ] + heapq.heapify(pending) + while pending: + consumed, _, i = heapq.heappop(pending) + _ensure_cooked(i) + if not cooked[i]: + return + row = _yield_row(i) + if consumed + 1 < catch_up_to: + heapq.heappush(pending, (consumed + 1, my_splits[i], i)) + _update_progress_stats() + yield row + while True: # A cycle only runs if every split can still produce a # row. Without skips all splits exhaust simultaneously @@ -904,20 +1158,14 @@ class StreamingDataset(IterableDataset): break for i in range(n): - pos, row = cooked[i].popleft() - local_consumed[i] += 1 - pos_consumed[i] = pos + 1 - _advance(i) + row = _yield_row(i) # After the last split in each cycle: update the # global offset and refresh the shared-memory stats # so the main process can observe pipeline depth # even when __iter__ runs in a worker process. if i == n - 1: - self._resume_offset = initial_offset + local_consumed[i] - for j, split_idx in enumerate(my_splits): - self._resume_positions[split_idx] = pos_consumed[j] - _update_stats() + _update_progress_stats() yield row finally: @@ -1064,6 +1312,7 @@ class StreamingDataset(IterableDataset): "_local_consumed_ref", ): state[key] = None + state["_consumer_iterator_lock"] = None return state def __setstate__(self, state): @@ -1074,19 +1323,31 @@ class StreamingDataset(IterableDataset): table_state = state.pop("_table") perm_name, perm_data = state.pop("_perm_table") self.__dict__.update(state) + self._consumer_iterator_lock = threading.Lock() if self._connection_factory is not None: self._table = self._connection_factory(table_name) else: self._table = _table_from_pickle_state(table_state) + if table_state["kind"] == "memory": + # Rebuilt from Arrow, so the recorded pin cannot resolve on it. + perm_data = _drop_base_version(perm_data) self._perm_table = _connect("memory://").create_table(perm_name, perm_data) def state_dict(self) -> dict: """Snapshot the dataset's consumption state. + When using DataLoader workers, construct a + [StreamingDataLoader][lancedb.streaming.StreamingDataLoader]. It + commits worker state only when a prefetched batch is returned to the + trainer. A standard multi-process ``DataLoader`` cannot expose that + boundary, so calling this method after one has started raises + ``RuntimeError`` instead of returning stale producer state. + In row mode, the returned dict is topology-independent at global step boundaries. ``positions_consumed_per_split`` records how far each split's permutation has advanced, which can differ from the sample - count when ``on_transform_error`` skips rows. Combine state dicts from + count when ``on_transform_error`` skips rows. ``StreamingDataLoader`` + combines worker state in its parent process. Combine state dicts from every rank with [merge_state_dicts][lancedb.streaming.StreamingDataset.merge_state_dicts] before resuming on a different topology. @@ -1095,6 +1356,43 @@ class StreamingDataset(IterableDataset): for every logical split. When packing is sharded, merge every rank state with ``merge_state_dicts`` before loading it. """ + if self._untracked_worker_iteration[0] and get_worker_info() is None: + raise RuntimeError( + "StreamingDataset cannot checkpoint a standard DataLoader with " + "num_workers > 0 because prefetched worker progress is not " + "consumer-committed. Use StreamingDataLoader instead." + ) + if self._checkpoint_invalid_reason is not None: + raise RuntimeError( + "StreamingDataset checkpointing is invalid because " + f"{self._checkpoint_invalid_reason}. Load the last valid " + "checkpoint into a fresh dataset before continuing." + ) + state = self._checkpoint_snapshot() + if self._pack_sequences is not None: + rank_blocks = [ + state["blocks_emitted_per_split"][split] for split in self._rank_splits + ] + if len(set(rank_blocks)) > 1: + raise RuntimeError( + "Packed StreamingDataset checkpointing is only safe at a " + "complete logical step boundary, when every split assigned " + "to this rank has emitted the same block count. Consume more " + "batches before calling state_dict()." + ) + elif self._consumer_checkpoint_requires_uniform: + samples = state["samples_consumed_per_split"] + rank_samples = [samples[split] for split in self._rank_splits] + if len(set(rank_samples)) > 1: + raise RuntimeError( + "StreamingDataLoader checkpointing with multiple workers is " + "only safe at a complete logical step boundary, when every " + "split assigned to this rank has the same consumed-sample " + "count. Consume more batches before calling state_dict()." + ) + return state + + def _checkpoint_snapshot(self) -> dict: if self._pack_sequences is not None: return { "shuffle_seed": self._shuffle_seed, @@ -1108,18 +1406,141 @@ class StreamingDataset(IterableDataset): "blocks_emitted_per_split": list(self._pack_blocks_emitted), "pack_buffers": deepcopy(self._pack_buffers), } + samples = [ + self._resume_samples.get(split, self._resume_offset) + for split in range(self._num_splits) + ] positions = [ - self._resume_positions.get(split, self._resume_offset) + self._resume_positions.get(split, samples[split]) for split in range(self._num_splits) ] return { "shuffle_seed": self._shuffle_seed, "num_splits": self._num_splits, "epoch": self._epoch, - "samples_consumed_per_split": [self._resume_offset] * self._num_splits, + "samples_consumed_per_split": samples, "positions_consumed_per_split": positions, } + def _invalidate_checkpoint(self, reason: str) -> None: + if self._checkpoint_invalid_reason is None: + self._checkpoint_invalid_reason = reason + + @property + def _consumer_iterator_active(self) -> bool: + return self._consumer_iterator_lease is not None + + @property + def _consumer_iterator_owner(self) -> Optional[int]: + lease = self._consumer_iterator_lease + return lease.owner_token if lease is not None else None + + @property + def _consumer_iterator_owner_thread(self) -> Optional[int]: + lease = self._consumer_iterator_lease + return lease.owner_thread if lease is not None else None + + def _acquire_consumer_iterator(self) -> int: + """Reserve this parent dataset for one checkpoint-aware iterator.""" + with self._consumer_iterator_lock: + if self._consumer_iterator_active or self._raw_batches_ref is not None: + raise RuntimeError( + "StreamingDataset does not support concurrent iteration. " + "Only one active iterator per dataset instance is allowed." + ) + owner_thread = threading.get_ident() + owner_token = self._consumer_iterator_generation + 1 + lease = _ConsumerIteratorLease(owner_token, owner_thread) + self._consumer_iterator_generation = owner_token + self._consumer_iterator_lease = lease + return owner_token + + def _release_consumer_iterator(self, owner_token: int) -> None: + with self._consumer_iterator_lock: + lease = self._consumer_iterator_lease + if lease is not None and lease.owner_token == owner_token: + self._consumer_iterator_lease = None + + def _release_consumer_iterator_after_failed_acquire( + self, previous_lease: Optional[_ConsumerIteratorLease] + ) -> None: + """Clean up when an interrupted acquire set a lease but did not return it.""" + owner_thread = threading.current_thread().ident + with self._consumer_iterator_lock: + lease = self._consumer_iterator_lease + if ( + lease is not None + and lease is not previous_lease + and lease.owner_thread == owner_thread + ): + self._consumer_iterator_lease = None + + def _commit_worker_state(self, state: dict, *, require_uniform: bool) -> None: + """Merge one trainer-consumed worker batch into parent state.""" + for key, expected in ( + ("shuffle_seed", self._shuffle_seed), + ("num_splits", self._num_splits), + ("epoch", self._epoch), + ): + if state.get(key) != expected: + raise ValueError( + f"{key} mismatch in worker checkpoint: " + f"{state.get(key)} != {expected}" + ) + packed = "pack_buffers" in state + if packed != (self._pack_sequences is not None): + raise ValueError("worker checkpoint mode does not match the dataset") + if packed: + for key in ("pack_sequences", "eos_id", "pad_id", "blocks_per_epoch"): + expected = getattr(self, f"_{key}") + if state.get(key) != expected: + raise ValueError( + f"{key} mismatch in worker checkpoint: " + f"{state.get(key)} != {expected}" + ) + samples = state["samples_consumed_per_split"] + emitted = state["blocks_emitted_per_split"] + if len(samples) != self._num_splits or len(emitted) != self._num_splits: + raise ValueError( + "packed worker checkpoint must contain one entry per split" + ) + buffers = state["pack_buffers"] + for split, (count, blocks) in enumerate(zip(samples, emitted)): + incoming = (int(blocks), int(count)) + current = ( + self._pack_blocks_emitted[split], + self._pack_consumed[split], + ) + if incoming > current: + self._pack_blocks_emitted[split] = incoming[0] + self._pack_consumed[split] = incoming[1] + buffer = buffers.get(split, buffers.get(str(split))) + if buffer is None: + self._pack_buffers.pop(split, None) + else: + self._pack_buffers[split] = { + "tokens": list(buffer["tokens"]), + "starts": list(buffer["starts"]), + } + self._consumer_checkpoint_requires_uniform |= require_uniform + return + + samples = state["samples_consumed_per_split"] + positions = state.get("positions_consumed_per_split", samples) + for split, count in enumerate(samples): + current = self._resume_samples.get(split, self._resume_offset) + self._resume_samples[split] = max(current, int(count)) + for split, position in enumerate(positions): + current = self._resume_positions.get( + split, self._resume_samples.get(split, self._resume_offset) + ) + self._resume_positions[split] = max(current, int(position)) + self._resume_offset = min( + self._resume_samples.get(split, self._resume_offset) + for split in range(self._num_splits) + ) + self._consumer_checkpoint_requires_uniform |= require_uniform + def load_state_dict(self, state: dict) -> None: """Resume from a previously snapshotted state. @@ -1139,6 +1560,7 @@ class StreamingDataset(IterableDataset): f"shuffle_seed mismatch: checkpoint has {state['shuffle_seed']}, " f"current dataset has {self._shuffle_seed}" ) + self._consumer_checkpoint_requires_uniform = False if "pack_buffers" in state or self._pack_sequences is not None: for key in ( @@ -1165,14 +1587,17 @@ class StreamingDataset(IterableDataset): return consumed = state["samples_consumed_per_split"] - # All entries are equal at step boundaries; use the first. if isinstance(consumed, list): - self._resume_offset = consumed[0] if consumed else 0 + self._resume_offset = min(consumed) if consumed else 0 + self._resume_samples = { + split: int(count) for split, count in enumerate(consumed) + } else: self._resume_offset = int(consumed) + self._resume_samples = {} # Older checkpoints predate positions_consumed_per_split; without # skipped rows positions equal sample counts, so falling back to - # _resume_offset (the .get default in __iter__) is exact. + # the per-split sample count (the .get default in __iter__) is exact. positions = state.get("positions_consumed_per_split") if positions is None: self._resume_positions = {} @@ -1185,10 +1610,11 @@ class StreamingDataset(IterableDataset): def merge_state_dicts(states: list[dict]) -> dict: """Merge state dicts saved by different ranks into one exact state. - For row mode, the elementwise maximum of permutation positions recovers - splits advanced by different ranks after transform failures. For packed - mode, the state that emitted the most blocks for each logical split - supplies that split's permutation position and partial token buffer. Packed + In row mode, each rank records exact consumer-committed progress for + its own splits and lower bounds for the rest, so elementwise maxima + recover both sample counts and permutation positions. In packed mode, + the state that emitted the most blocks for each logical split supplies + that split's permutation position and partial token buffer. Packed states must cover every rank at the same global step. Raises ``ValueError`` if the states are empty, were not produced by @@ -1299,17 +1725,13 @@ class StreamingDataset(IterableDataset): merged["pack_buffers"] = merged_buffers return merged - for state in states[1:]: - if ( - state["samples_consumed_per_split"] - != first["samples_consumed_per_split"] - ): - raise ValueError( - "samples_consumed_per_split mismatch across state dicts; " - "state_dict() must be called at the same global step " - "boundary on every rank" - ) merged = dict(first) + merged["samples_consumed_per_split"] = [ + max(per_split) + for per_split in zip( + *(state["samples_consumed_per_split"] for state in states) + ) + ] all_positions = [ state.get( "positions_consumed_per_split", state["samples_consumed_per_split"] @@ -1320,3 +1742,113 @@ class StreamingDataset(IterableDataset): max(per_split) for per_split in zip(*all_positions) ] return merged + + +class StreamingDataLoader(DataLoader): + """A PyTorch DataLoader with consumer-committed dataset checkpoints. + + PyTorch workers prefetch batches ahead of the trainer, so worker-local + producer progress is not a safe checkpoint. This loader carries a state + snapshot alongside every internal batch and applies it to the parent + [StreamingDataset][lancedb.streaming.StreamingDataset] only when that batch + is returned by ``next()``. + The trainer receives the same collated batch it would receive from a + standard ``torch.utils.data.DataLoader``. + + With more than one worker, row-mode ``state_dict()`` is available only at + complete logical step boundaries, when every split assigned to the rank has + the same consumed-sample count. Packed checkpoints require equal emitted-block + counts across the rank's splits for any worker count. ``persistent_workers=True`` + is not supported because prefetched worker copies cannot be restored from + parent-committed state. If batch collation raises, checkpointing remains + invalid for that dataset instance; restore the last valid checkpoint into a + fresh dataset before continuing. + Only one active iterator may own a dataset at a time, including when worker + processes are used. Exhausting or explicitly shutting down the iterator + releases that ownership. ``drop_last=True`` is not supported because worker + replicas discard incomplete tails independently, which cannot produce a + topology-independent checkpoint. + + Parameters are the same as ``torch.utils.data.DataLoader`` except that + ``dataset`` must be a + [StreamingDataset][lancedb.streaming.StreamingDataset]. + Subclasses that override ``StreamingDataset.__iter__`` are not supported + because the custom iterator cannot provide the exact per-yield checkpoint + snapshots required by this loader. + + Examples + -------- + >>> # dataset = StreamingDataset(table, num_splits=2) + >>> # loader = StreamingDataLoader(dataset, batch_size=8, num_workers=2) + >>> # batch = next(iter(loader)) + >>> # checkpoint = dataset.state_dict() + """ + + def __init__(self, dataset: StreamingDataset, *args, **kwargs): + if not isinstance(dataset, StreamingDataset): + raise TypeError("StreamingDataLoader requires a StreamingDataset") + if type(dataset).__iter__ is not StreamingDataset.__iter__: + raise TypeError( + "StreamingDataLoader does not support StreamingDataset subclasses " + "that override __iter__ because they cannot provide exact " + "per-yield checkpoint state" + ) + if kwargs.get("in_order", True) is False: + raise ValueError( + "StreamingDataLoader requires in_order=True for deterministic " + "consumer checkpoints" + ) + if kwargs.get("persistent_workers", False): + raise ValueError( + "StreamingDataLoader does not support persistent_workers=True " + "because worker prefetch state cannot be reset from a checkpoint" + ) + self._streaming_dataset = dataset + super().__init__(_StreamingDatasetAdapter(dataset), *args, **kwargs) + if self.drop_last: + raise ValueError( + "StreamingDataLoader does not support drop_last=True because " + "discarded worker tails cannot be checkpointed " + "topology-independently" + ) + self.collate_fn = _CheckpointCollate(self.collate_fn) + + def __iter__(self): + dataset = self._streaming_dataset + previous_lease = dataset._consumer_iterator_lease + owner_token = None + try: + owner_token = dataset._acquire_consumer_iterator() + state = dataset._checkpoint_snapshot() + packed = dataset._pack_sequences is not None + if packed: + blocks = state["blocks_emitted_per_split"] + rank_blocks = [blocks[split] for split in dataset._rank_splits] + if len(set(rank_blocks)) > 1: + raise RuntimeError( + "StreamingDataLoader cannot start from a partial packed " + "logical step; resume from a checkpoint whose splits " + "assigned to this rank have equal emitted-block counts" + ) + elif self.num_workers > 1: + samples = state["samples_consumed_per_split"] + rank_samples = [samples[split] for split in dataset._rank_splits] + if len(set(rank_samples)) > 1: + raise RuntimeError( + "StreamingDataLoader cannot start multiple workers from a " + "partial logical step; resume from a checkpoint whose " + "splits assigned to this rank have equal consumed-sample " + "counts" + ) + return _ConsumerCommitIterator( + super().__iter__(), + dataset, + owner_token=owner_token, + require_uniform=self.num_workers > 1 or packed, + ) + except BaseException: + if owner_token is not None: + dataset._release_consumer_iterator(owner_token) + else: + dataset._release_consumer_iterator_after_failed_acquire(previous_lease) + raise diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index fd0c721a5..aa3021317 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -1269,6 +1269,7 @@ class Table(ABC): fill_value: float = 0.0, progress: Optional[Union[bool, Callable, Any]] = None, write_parallelism: Optional[int] = None, + allow_external_blob_outside_bases: bool = False, ) -> AddResult: """Add more data to the [Table][lancedb.table.Table]. @@ -1320,6 +1321,10 @@ class Table(ABC): data in flight. Defaults to an estimate based on the data size, capped at the number of CPU cores. Lower this if bulk ingestion is using too much memory. + allow_external_blob_outside_bases: bool, default False + Store blob URIs that sit outside registered blob bases. The row + keeps a reference, so the object has to stay readable. Local + tables only. Returns ------- @@ -1972,7 +1977,7 @@ class Table(ABC): A mapping with one ``FunctionApplication`` value keeps its scalar or named-struct result in the named table column. A bare named-struct application expands its ordered result fields as one - atomic sibling group; aliases come from ``rename(columns=...)``. + atomic binding; aliases come from ``rename(columns=...)``. Function columns are supported only on LanceDB Cloud and Enterprise. computed: Dict[str, str], optional @@ -3409,6 +3414,7 @@ class LanceTable(Table): fill_value: float = 0.0, progress: Optional[Union[bool, Callable, Any]] = None, write_parallelism: Optional[int] = None, + allow_external_blob_outside_bases: bool = False, ) -> AddResult: """Add data to the table. If vector columns are missing and the table @@ -3436,6 +3442,9 @@ class LanceTable(Table): data in flight. Defaults to an estimate based on the data size, capped at the number of CPU cores. Lower this if bulk ingestion is using too much memory. + allow_external_blob_outside_bases: bool, default False + Allow blob URIs outside registered bases. See :meth:`Table.add`. + Local tables only. Returns ------- @@ -3452,6 +3461,7 @@ class LanceTable(Table): fill_value=fill_value, progress=progress, write_parallelism=write_parallelism, + allow_external_blob_outside_bases=allow_external_blob_outside_bases, ) ) finally: @@ -5366,6 +5376,7 @@ class AsyncTable: fill_value: Optional[float] = None, progress: Optional[Union[bool, Callable, Any]] = None, write_parallelism: Optional[int] = None, + allow_external_blob_outside_bases: bool = False, ) -> AddResult: """Add more data to the [AsyncTable][lancedb.table.AsyncTable]. @@ -5396,6 +5407,9 @@ class AsyncTable: data in flight. Defaults to an estimate based on the data size, capped at the number of CPU cores. Lower this if bulk ingestion is using too much memory. + allow_external_blob_outside_bases: bool, default False + Allow blob URIs outside registered bases. See :meth:`Table.add`. + Local tables only. """ schema = await self.schema() @@ -5432,6 +5446,7 @@ class AsyncTable: mode or "append", progress=progress, write_parallelism=write_parallelism, + allow_external_blob_outside_bases=allow_external_blob_outside_bases, ) except RuntimeError as e: if "Cast error" in str(e): @@ -6056,7 +6071,7 @@ class AsyncTable: A mapping with one ``FunctionApplication`` value keeps its scalar or named-struct result in the named table column. A bare named-struct application expands its ordered result fields as one - atomic sibling group; aliases come from ``rename(columns=...)``. + atomic binding; aliases come from ``rename(columns=...)``. Function columns are supported only on LanceDB Cloud and Enterprise. computed: Dict[str, str], optional @@ -6093,7 +6108,7 @@ class AsyncTable: isinstance(value, FunctionApplication) for value in transforms.values() ): raise ValueError( - "one add_columns call declares exactly one Function sibling group" + "one add_columns call declares exactly one Function binding" ) function_output_name, function_application = next(iter(transforms.items())) diff --git a/python/python/tests/test_blob.py b/python/python/tests/test_blob.py index b205c20a6..5d7682f24 100644 --- a/python/python/tests/test_blob.py +++ b/python/python/tests/test_blob.py @@ -617,3 +617,71 @@ def test_fetch_blobs_nested_path_survives_sort_after_query(): def _identifiable_payload(size: int) -> bytes: block = 256 return b"".join(bytes([i % 256]) * block for i in range(size // block)) + + +def _external_uri_blob_array(uris): + blob_type = lancedb.blob("image").type + storage_type = blob_type.storage_type + child_names = [field.name for field in storage_type] + assert "uri" in child_names, "blob layout no longer has a uri child" + children = [ + pa.array(uris if field.name == "uri" else [None] * len(uris), type=field.type) + for field in storage_type + ] + storage = pa.StructArray.from_arrays(children, fields=list(storage_type)) + return pa.ExtensionArray.from_storage(blob_type, storage) + + +def _external_uri_table_and_rows(name, uris): + db = lancedb.connect("memory:///") + schema = pa.schema([pa.field("id", pa.int64()), lancedb.blob("image")]) + table = db.create_table(name, schema=schema) + rows = pa.Table.from_arrays( + [ + pa.array(range(len(uris)), type=pa.int64()), + _external_uri_blob_array(uris), + ], + schema=schema, + ) + return table, rows + + +def test_add_external_uri_struct_round_trips_with_flag(tmp_path): + payload = b"external-uri-bytes" + blob_path = tmp_path / "payload.bin" + blob_path.write_bytes(payload) + + table, rows = _external_uri_table_and_rows("external_struct", [blob_path.as_uri()]) + table.add(rows, allow_external_blob_outside_bases=True) + + hits = table.search().to_arrow() + blobs = table.fetch_blobs("image", hits) + assert blobs[0].as_py() == payload + + +def test_add_external_uri_without_flag_raises(tmp_path): + blob_path = tmp_path / "payload.bin" + blob_path.write_bytes(b"unreachable") + + table, rows = _external_uri_table_and_rows("external_no_flag", [blob_path.as_uri()]) + with pytest.raises(ValueError, match="allow_external_blob_outside_bases"): + table.add(rows) + assert table.count_rows() == 0 + + +def test_add_external_uri_string_round_trips_with_flag(tmp_path): + payload = b"external-uri-bytes" + blob_path = tmp_path / "payload.bin" + blob_path.write_bytes(payload) + + db = lancedb.connect("memory:///") + schema = pa.schema([pa.field("id", pa.int64()), lancedb.blob("image")]) + table = db.create_table("external_string", schema=schema) + table.add( + [{"id": 1, "image": blob_path.as_uri()}], + allow_external_blob_outside_bases=True, + ) + + hits = table.search().to_arrow() + blobs = table.fetch_blobs("image", hits) + assert blobs[0].as_py() == payload diff --git a/python/python/tests/test_elastic_dataloader.py b/python/python/tests/test_elastic_dataloader.py index ff65d100c..da27e5bfc 100644 --- a/python/python/tests/test_elastic_dataloader.py +++ b/python/python/tests/test_elastic_dataloader.py @@ -32,6 +32,7 @@ Parameters used throughout: import dataclasses import logging +import threading from unittest.mock import patch import lancedb @@ -46,6 +47,7 @@ from utils import ( torch = pytest.importorskip("torch") streaming = pytest.importorskip("lancedb.streaming") StreamingDataset = streaming.StreamingDataset +StreamingDataLoader = streaming.StreamingDataLoader # --------------------------------------------------------------------------- # Dataset parameters @@ -92,6 +94,27 @@ class FakeWorkerInfo: num_workers: int +def _collate_with_first_batch_error(samples): + ids = [sample["id"] for sample in samples] + if ids == [0, 1]: + raise ValueError("first batch fails") + return ids + + +def _collate_with_first_batch_stop(samples): + ids = [sample["id"] for sample in samples] + if ids == [0, 1]: + raise StopIteration("first batch stopped") + return ids + + +def _collate_with_first_batch_interrupt(samples): + ids = [sample["id"] for sample in samples] + if ids == [0, 1]: + raise KeyboardInterrupt("first batch interrupted") + return ids + + # --------------------------------------------------------------------------- # Fixtures # --------------------------------------------------------------------------- @@ -1008,6 +1031,565 @@ def test_multi_worker_elastic_det_across_worker_counts(lance_table): # ── Resumability with num_workers ───────────────────────────────────────────── +def test_streaming_dataloader_commits_only_consumed_worker_batches(tmp_path): + """Prefetched worker state is committed only as the trainer receives it.""" + db = lancedb.connect(tmp_path) + table = db.create_table( + "worker_commit", pa.table({"id": [1, 2, 3, 4, 10, 20, 30, 40]}) + ) + dataset = StreamingDataset(table, num_splits=2, shuffle=False) + loader = StreamingDataLoader( + dataset, + batch_size=2, + num_workers=2, + multiprocessing_context="spawn", + prefetch_factor=4, + ) + iterator = iter(loader) + try: + first = next(iterator)["id"].tolist() + + assert first == [1, 2] + assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [2, 0] + with pytest.raises(RuntimeError, match="complete logical step boundary"): + dataset.state_dict() + + second = next(iterator)["id"].tolist() + assert second == [10, 20] + checkpoint = dataset.state_dict() + assert checkpoint["samples_consumed_per_split"] == [2, 2] + uninterrupted = [batch["id"].tolist() for batch in iterator] + finally: + iterator._shutdown_workers() + + resumed = StreamingDataset(table, num_splits=2, shuffle=False) + resumed.load_state_dict(checkpoint) + resumed_loader = StreamingDataLoader( + resumed, + batch_size=2, + num_workers=2, + multiprocessing_context="spawn", + prefetch_factor=4, + ) + resumed_iterator = iter(resumed_loader) + try: + remaining = [batch["id"].tolist() for batch in resumed_iterator] + finally: + resumed_iterator._shutdown_workers() + assert remaining == uninterrupted == [[3, 4], [30, 40]] + + +def test_distributed_checkpoint_uses_rank_local_worker_boundary(tmp_path): + db = lancedb.connect(tmp_path) + table = db.create_table("rank_boundary", pa.table({"id": list(range(8))})) + dataset = StreamingDataset( + table, + num_splits=4, + shuffle=False, + rank=0, + world_size=2, + ) + loader = StreamingDataLoader( + dataset, + batch_size=1, + num_workers=2, + multiprocessing_context="spawn", + ) + iterator = iter(loader) + try: + assert next(iterator)["id"].tolist() == [0] + assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [ + 1, + 0, + 0, + 0, + ] + with pytest.raises(RuntimeError, match="complete logical step boundary"): + dataset.state_dict() + + assert next(iterator)["id"].tolist() == [2] + checkpoint = dataset.state_dict() + remaining = [batch["id"].tolist() for batch in iterator] + finally: + iterator._shutdown_workers() + + assert checkpoint["samples_consumed_per_split"] == [1, 1, 0, 0] + assert remaining == [[1], [3]] + + +def test_standard_dataloader_rejects_stale_parent_checkpoint(tmp_path): + """A standard DataLoader must not expose prefetched producer progress.""" + db = lancedb.connect(tmp_path) + table = db.create_table("untracked_workers", pa.table({"id": [1, 2, 10, 20]})) + dataset = StreamingDataset(table, num_splits=2, shuffle=False) + # Merely constructing the checkpoint-aware loader must not authorize a + # later plain DataLoader's worker progress. + StreamingDataLoader(dataset, batch_size=2, num_workers=0) + loader = torch.utils.data.DataLoader( + dataset, + batch_size=2, + num_workers=2, + multiprocessing_context="spawn", + ) + iterator = iter(loader) + try: + assert next(iterator)["id"].tolist() == [1, 2] + with pytest.raises(RuntimeError, match="Use StreamingDataLoader"): + dataset.state_dict() + list(iterator) + finally: + iterator._shutdown_workers() + + +def test_streaming_dataloader_rejects_persistent_workers(tmp_path): + db = lancedb.connect(tmp_path) + table = db.create_table("persistent_workers", pa.table({"id": [1, 2]})) + dataset = StreamingDataset(table, num_splits=2, shuffle=False) + + with pytest.raises(ValueError, match="persistent_workers=True"): + StreamingDataLoader( + dataset, + batch_size=1, + num_workers=2, + persistent_workers=True, + ) + + +def test_collate_failure_invalidates_consumer_checkpoint(tmp_path): + db = lancedb.connect(tmp_path) + table = db.create_table( + "collate_failure", pa.table({"id": [0, 1, 2, 3, 100, 101, 102, 103]}) + ) + dataset = StreamingDataset(table, num_splits=2, shuffle=False) + loader = StreamingDataLoader( + dataset, + batch_size=2, + num_workers=2, + multiprocessing_context="spawn", + collate_fn=_collate_with_first_batch_error, + prefetch_factor=2, + ) + iterator = iter(loader) + try: + with pytest.raises(ValueError, match="first batch fails"): + next(iterator) + assert next(iterator) == [100, 101] + assert next(iterator) == [2, 3] + with pytest.raises(RuntimeError, match="failed before it was returned"): + dataset.state_dict() + list(iterator) + finally: + iterator._shutdown_workers() + + +def test_collate_stop_iteration_invalidates_consumer_checkpoint(tmp_path): + db = lancedb.connect(tmp_path) + table = db.create_table("collate_stop", pa.table({"id": list(range(6))})) + dataset = StreamingDataset(table, num_splits=1, shuffle=False) + loader = StreamingDataLoader( + dataset, + batch_size=2, + num_workers=0, + collate_fn=_collate_with_first_batch_stop, + ) + iterator = iter(loader) + + with pytest.raises(RuntimeError, match="collate_fn raised StopIteration"): + next(iterator) + assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [2] + with pytest.raises(RuntimeError, match="failed before it was returned"): + dataset.state_dict() + assert list(iterator) == [[2, 3], [4, 5]] + + +def test_batch_base_exception_invalidates_consumer_checkpoint(tmp_path): + db = lancedb.connect(tmp_path) + table = db.create_table("collate_interrupt", pa.table({"id": list(range(6))})) + dataset = StreamingDataset(table, num_splits=1, shuffle=False) + loader = StreamingDataLoader( + dataset, + batch_size=2, + num_workers=0, + collate_fn=_collate_with_first_batch_interrupt, + ) + iterator = iter(loader) + + with pytest.raises(KeyboardInterrupt, match="first batch interrupted"): + next(iterator) + assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [2] + with pytest.raises(RuntimeError, match="failed before it was returned"): + dataset.state_dict() + assert list(iterator) == [[2, 3], [4, 5]] + + +def test_parent_commit_base_exception_invalidates_consumer_checkpoint(tmp_path): + db = lancedb.connect(tmp_path) + table = db.create_table("commit_interrupt", pa.table({"id": list(range(4))})) + dataset = StreamingDataset(table, num_splits=1, shuffle=False) + loader = StreamingDataLoader(dataset, batch_size=2, num_workers=0) + iterator = iter(loader) + real_commit = dataset._commit_worker_state + + def interrupt_after_commit(state, *, require_uniform): + real_commit(state, require_uniform=require_uniform) + raise KeyboardInterrupt("after parent commit") + + with patch.object( + dataset, "_commit_worker_state", side_effect=interrupt_after_commit + ): + with pytest.raises(KeyboardInterrupt, match="after parent commit"): + next(iterator) + + assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [2] + with pytest.raises(RuntimeError, match="failed before it was returned"): + dataset.state_dict() + + +def test_direct_iteration_surfaces_prefetch_failure_before_committing_row( + tmp_path, monkeypatch +): + db = lancedb.connect(tmp_path) + table = db.create_table("prefetch_failure", pa.table({"id": list(range(4))})) + release = threading.Event() + failed = threading.Event() + real_getitems = streaming.Permutation.__getitems__ + + def controlled_getitems(permutation, indices): + if indices and indices[0] >= 2: + assert release.wait(timeout=5) + failed.set() + raise RuntimeError("later prefetched I/O failed") + return real_getitems(permutation, indices) + + class SignalDict(dict): + def __setitem__(self, key, value): + super().__setitem__(key, value) + release.set() + assert failed.wait(timeout=5) + + monkeypatch.setattr(streaming.Permutation, "__getitems__", controlled_getitems) + dataset = StreamingDataset( + table, + num_splits=1, + shuffle=False, + read_batch_size=2, + io_queue_depth=2, + ) + dataset._resume_positions = SignalDict() + iterator = iter(dataset) + + assert next(iterator)["id"] == 0 + with pytest.raises(RuntimeError, match="later prefetched I/O failed"): + next(iterator) + + checkpoint = dataset.state_dict() + assert checkpoint["samples_consumed_per_split"] == [1] + assert checkpoint["positions_consumed_per_split"] == [1] + + +@pytest.mark.parametrize("workers", [0, 1, 2]) +def test_streaming_dataloader_rejects_drop_last(tmp_path, workers): + db = lancedb.connect(tmp_path) + table = db.create_table("drop_last", pa.table({"id": [0, 1, 2]})) + dataset = StreamingDataset(table, num_splits=1, shuffle=False) + worker_options = {"multiprocessing_context": "spawn"} if workers else {} + + with pytest.raises(ValueError, match="drop_last=True"): + StreamingDataLoader( + dataset, + batch_size=2, + num_workers=workers, + drop_last=True, + **worker_options, + ) + + +def test_streaming_dataloader_owns_one_iterator_until_teardown(tmp_path): + db = lancedb.connect(tmp_path) + table = db.create_table("iterator_owner", pa.table({"id": list(range(4))})) + dataset = StreamingDataset(table, num_splits=1, shuffle=False) + loader = StreamingDataLoader( + dataset, + batch_size=2, + num_workers=1, + multiprocessing_context="spawn", + ) + + first = iter(loader) + try: + assert next(first)["id"].tolist() == [0, 1] + with pytest.raises(RuntimeError, match="concurrent iteration"): + iter(loader) + finally: + first._shutdown_workers() + + second = iter(loader) + try: + assert [batch["id"].tolist() for batch in second] == [[2, 3]] + except BaseException: + second._shutdown_workers() + raise + + # Natural exhaustion releases ownership too. + third = iter(loader) + try: + assert list(third) == [] + finally: + third._shutdown_workers() + + +def test_zero_worker_shutdown_closes_inner_iterator_before_release(tmp_path): + db = lancedb.connect(tmp_path) + table = db.create_table("zero_worker_shutdown", pa.table({"id": list(range(6))})) + dataset = StreamingDataset(table, num_splits=1, shuffle=False) + loader = StreamingDataLoader(dataset, batch_size=2, num_workers=0) + + first = iter(loader) + assert next(first)["id"].tolist() == [0, 1] + first._shutdown_workers() + + assert dataset._consumer_iterator_active is False + assert dataset._raw_batches_ref is None + second = iter(loader) + try: + with pytest.raises(StopIteration): + next(first) + assert next(second)["id"].tolist() == [2, 3] + finally: + second._shutdown_workers() + + +def test_direct_and_loader_admission_share_one_atomic_lease(tmp_path, monkeypatch): + db = lancedb.connect(tmp_path) + table = db.create_table("direct_loader_lease", pa.table({"id": list(range(4))})) + dataset = StreamingDataset(table, num_splits=1, shuffle=False) + loader = StreamingDataLoader(dataset, batch_size=2, num_workers=0) + entered = threading.Event() + release = threading.Event() + direct_result = [] + direct_error = [] + contender = [] + real_resolve = dataset._resolve_my_splits + + def controlled_resolve(): + if threading.current_thread().name == "direct-start": + entered.set() + assert release.wait(timeout=5) + return real_resolve() + + def advance_direct(iterator): + try: + direct_result.append(next(iterator)["id"]) + except BaseException as exc: + direct_error.append(exc) + + monkeypatch.setattr(dataset, "_resolve_my_splits", controlled_resolve) + direct = iter(dataset) + thread = threading.Thread( + target=advance_direct, args=(direct,), name="direct-start" + ) + thread.start() + assert entered.wait(timeout=5) + try: + with pytest.raises(RuntimeError, match="concurrent iteration"): + contender.append(iter(loader)) + finally: + release.set() + thread.join(timeout=5) + if contender: + contender[0]._shutdown_workers() + direct.close() + + assert not thread.is_alive() + assert direct_error == [] + assert direct_result == [0] + + +def test_loader_acquires_before_snapshot_and_cleans_interrupted_acquire( + tmp_path, monkeypatch +): + db = lancedb.connect(tmp_path) + table = db.create_table("lease_snapshot", pa.table({"id": list(range(4))})) + dataset = StreamingDataset(table, num_splits=1, shuffle=False) + loader = StreamingDataLoader(dataset, batch_size=2, num_workers=0) + first = iter(loader) + assert next(first)["id"].tolist() == [0, 1] + + entered = threading.Event() + release = threading.Event() + pending = [] + pending_errors = [] + observed_snapshots = [] + real_acquire = dataset._acquire_consumer_iterator + real_snapshot = dataset._checkpoint_snapshot + + def controlled_acquire(): + if threading.current_thread().name == "stale-start": + entered.set() + assert release.wait(timeout=5) + return real_acquire() + + def recording_snapshot(): + state = real_snapshot() + if threading.current_thread().name == "stale-start": + observed_snapshots.append(state["samples_consumed_per_split"]) + return state + + def create_pending_iterator(): + try: + pending.append(iter(loader)) + except BaseException as exc: + pending_errors.append(exc) + + monkeypatch.setattr(dataset, "_acquire_consumer_iterator", controlled_acquire) + monkeypatch.setattr(dataset, "_checkpoint_snapshot", recording_snapshot) + thread = threading.Thread(target=create_pending_iterator, name="stale-start") + thread.start() + assert entered.wait(timeout=5) + assert next(first)["id"].tolist() == [2, 3] + with pytest.raises(StopIteration): + next(first) + release.set() + thread.join(timeout=5) + + assert not thread.is_alive() + assert pending_errors == [] + assert observed_snapshots == [[4]] + assert len(pending) == 1 + assert list(pending[0]) == [] + assert dataset.state_dict()["samples_consumed_per_split"] == [4] + + def interrupted_acquire(): + real_acquire() + raise KeyboardInterrupt("after acquire") + + monkeypatch.setattr(dataset, "_acquire_consumer_iterator", interrupted_acquire) + with pytest.raises(KeyboardInterrupt, match="after acquire"): + iter(loader) + assert dataset._consumer_iterator_active is False + + +def test_consumer_iterator_lease_publication_is_atomic(tmp_path, monkeypatch): + db = lancedb.connect(tmp_path) + table = db.create_table("atomic_lease", pa.table({"id": [0, 1]})) + dataset = StreamingDataset(table, num_splits=1, shuffle=False) + loader = StreamingDataLoader(dataset, batch_size=1, num_workers=0) + real_get_ident = streaming.threading.get_ident + calls = 0 + + def interrupt_during_publication(): + nonlocal calls + calls += 1 + if calls == 1: + raise KeyboardInterrupt("during lease mutation") + return real_get_ident() + + monkeypatch.setattr(streaming.threading, "get_ident", interrupt_during_publication) + with pytest.raises(KeyboardInterrupt, match="during lease mutation"): + iter(loader) + monkeypatch.setattr(streaming.threading, "get_ident", real_get_ident) + + assert dataset._consumer_iterator_active is False + iterator = iter(loader) + try: + assert next(iterator)["id"].tolist() == [0] + finally: + iterator._shutdown_workers() + + +def test_streaming_dataloader_rejects_dataset_iter_override(tmp_path): + class CustomizedDataset(StreamingDataset): + def __iter__(self): + return iter([1000, 1001]) + + db = lancedb.connect(tmp_path) + table = db.create_table("custom_iteration", pa.table({"id": [0, 1, 2]})) + dataset = CustomizedDataset(table, num_splits=1, shuffle=False) + + assert list(dataset) == [1000, 1001] + with pytest.raises(TypeError, match="override __iter__"): + StreamingDataLoader( + dataset, + batch_size=2, + num_workers=0, + collate_fn=list, + ) + + +def test_interleaved_adapters_do_not_authorize_plain_iteration(tmp_path): + db = lancedb.connect(tmp_path) + table_a = db.create_table("adapter_a", pa.table({"id": [0, 1]})) + table_b = db.create_table("adapter_b", pa.table({"id": [10, 11]})) + dataset_a = StreamingDataset(table_a, num_splits=1, shuffle=False) + dataset_b = StreamingDataset(table_b, num_splits=1, shuffle=False) + initial_state = dataset_a.state_dict() + + owner_a = dataset_a._acquire_consumer_iterator() + owner_b = dataset_b._acquire_consumer_iterator() + try: + iterator_a = iter(streaming._StreamingDatasetAdapter(dataset_a)) + iterator_b = iter(streaming._StreamingDatasetAdapter(dataset_b)) + assert next(iterator_a).data["id"] == 0 + assert next(iterator_b).data["id"] == 10 + assert [sample.data["id"] for sample in iterator_a] == [1] + assert [sample.data["id"] for sample in iterator_b] == [11] + finally: + dataset_a._release_consumer_iterator(owner_a) + dataset_b._release_consumer_iterator(owner_b) + + dataset_a.load_state_dict(initial_state) + with patch( + "lancedb.streaming.get_worker_info", + return_value=FakeWorkerInfo(id=0, num_workers=1), + ): + plain_iterator = iter(dataset_a) + assert next(plain_iterator)["id"] == 0 + plain_iterator.close() + + assert dataset_a._untracked_worker_iteration[0] == 1 + with pytest.raises(RuntimeError, match="Use StreamingDataLoader"): + dataset_a.state_dict() + + +def test_resume_from_partial_split_cycle_preserves_remaining_order(tmp_path): + db = lancedb.connect(tmp_path) + table = db.create_table("partial_cycle", pa.table({"id": [1, 2, 10, 20]})) + dataset = StreamingDataset(table, num_splits=2, shuffle=False) + iterator = iter(dataset) + + assert next(iterator)["id"] == 1 + checkpoint = dataset.state_dict() + iterator.close() + assert checkpoint["samples_consumed_per_split"] == [1, 0] + + resumed = StreamingDataset(table, num_splits=2, shuffle=False) + resumed.load_state_dict(checkpoint) + assert [row["id"] for row in resumed] == [10, 2, 20] + + +def test_partial_cycle_resume_preserves_skip_truncation(tmp_path): + db = lancedb.connect(tmp_path) + table = db.create_table( + "partial_skip", pa.table({"id": [0, 1, 2, 3, 100, 101, 102, 103]}) + ) + kwargs = dict( + num_splits=2, + shuffle=False, + transform=_failing_transform({1, 2, 3}), + on_transform_error="skip", + ) + dataset = StreamingDataset(table, **kwargs) + iterator = iter(dataset) + + assert next(iterator)["id"] == 0 + checkpoint = dataset.state_dict() + uninterrupted = [row["id"] for row in iterator] + + resumed = StreamingDataset(table, **kwargs) + resumed.load_state_dict(checkpoint) + assert [row["id"] for row in resumed] == uninterrupted == [100] + + def test_multi_worker_resumability_same_topology(lance_table): """Checkpoint with num_workers=2, resume with num_workers=2: exact continuation.""" world_size = 1 @@ -2018,6 +2600,23 @@ def test_merge_state_dicts_validates_consistency(lance_table): StreamingDataset.merge_state_dicts([]) +def test_merge_state_dicts_combines_nonuniform_consumer_progress(lance_table): + dataset = StreamingDataset( + lance_table, num_splits=2, shuffle=False, shuffle_seed=SHUFFLE_SEED + ) + rank0 = dataset.state_dict() + rank0["samples_consumed_per_split"] = [2, 0] + rank0["positions_consumed_per_split"] = [2, 0] + rank1 = dataset.state_dict() + rank1["samples_consumed_per_split"] = [0, 2] + rank1["positions_consumed_per_split"] = [0, 2] + + merged = StreamingDataset.merge_state_dicts([rank0, rank1]) + + assert merged["samples_consumed_per_split"] == [2, 2] + assert merged["positions_consumed_per_split"] == [2, 2] + + def test_load_state_dict_without_positions_key(lance_table): """Checkpoints from before positions_consumed_per_split existed still resume exactly (positions equal sample counts when nothing is skipped).""" @@ -2254,6 +2853,65 @@ def test_pack_sequences_checkpoint_resumes_on_new_topology(tmp_path): ] +def test_packed_checkpoint_requires_complete_split_cycle(tmp_path): + table = _create_token_table(tmp_path, [[1], [2], [10], [20]]) + dataset = _packed_dataset(table, pack_sequences=3, blocks_per_epoch=4, num_splits=2) + iterator = iter(dataset) + + next(iterator) + with pytest.raises(RuntimeError, match="complete logical step boundary"): + dataset.state_dict() + + next(iterator) + assert dataset.state_dict()["blocks_emitted_per_split"] == [1, 1] + iterator.close() + + +def test_streaming_dataloader_commits_consumed_packed_batches(tmp_path): + table = _create_token_table( + tmp_path, + [[1], [2], [3], [4], [10], [20], [30], [40]], + ) + kwargs = dict(pack_sequences=4, blocks_per_epoch=4, num_splits=2) + dataset = _packed_dataset(table, **kwargs) + loader = StreamingDataLoader( + dataset, + batch_size=1, + num_workers=2, + multiprocessing_context="spawn", + prefetch_factor=2, + ) + iterator = iter(loader) + try: + next(iterator) + with pytest.raises(RuntimeError, match="complete logical step boundary"): + dataset.state_dict() + + next(iterator) + checkpoint = dataset.state_dict() + uninterrupted = [batch["input_ids"].tolist() for batch in iterator] + finally: + iterator._shutdown_workers() + + resumed = _packed_dataset(table, **kwargs) + resumed.load_state_dict(checkpoint) + resumed_loader = StreamingDataLoader( + resumed, + batch_size=1, + num_workers=2, + multiprocessing_context="spawn", + prefetch_factor=2, + ) + resumed_iterator = iter(resumed_loader) + try: + remaining = [batch["input_ids"].tolist() for batch in resumed_iterator] + finally: + resumed_iterator._shutdown_workers() + + assert checkpoint["blocks_emitted_per_split"] == [1, 1] + assert remaining == uninterrupted + + def test_pack_sequences_validates_configuration_and_tokens(tmp_path): table = _create_token_table(tmp_path, [[1, 2]]) diff --git a/python/python/tests/test_first_class_function_slice1.py b/python/python/tests/test_first_class_function_slice1.py index 1bb8feaf4..89172ba3f 100644 --- a/python/python/tests/test_first_class_function_slice1.py +++ b/python/python/tests/test_first_class_function_slice1.py @@ -37,21 +37,6 @@ def job_result(name: str) -> dict: return json.loads(fixture(name))["result"] -def assert_no_secret_values(value): - if isinstance(value, dict): - for key, child in value.items(): - assert key not in { - "secret_value", - "secret_values", - "resolved_secret", - "resolved_secrets", - } - assert_no_secret_values(child) - elif isinstance(value, list): - for child in value: - assert_no_secret_values(child) - - def test_public_function_values_are_in_api_reference(): docs = Path(__file__).parents[3] / "docs" / "src" / "python" / "python.md" rendered = docs.read_text() @@ -109,7 +94,6 @@ def test_function_version_identity_is_immutable_and_exact(): version = FunctionVersion.from_json(json.dumps(value)) assert version.name == "embed" assert version.version == "fv_01K3EXACT" - assert version.required_secrets == ("HF_TOKEN",) with pytest.raises((TypeError, ValueError)): version.version = "fv_changed" @@ -121,7 +105,7 @@ def test_function_version_identity_is_immutable_and_exact(): assert FunctionVersion(**changed) != version -def test_function_version_binds_named_columns_as_one_immutable_group(): +def test_function_version_binds_named_columns_as_one_immutable_application(): version = FunctionVersion.from_json( json.dumps(job_result("remote_function_job.json")) ) @@ -131,13 +115,10 @@ def test_function_version_binds_named_columns_as_one_immutable_group(): assert application.function.name == version.name assert application.function.version == version.version assert application.output is version.signature.output - assert application.group_id.startswith("fg_") assert [ (value.parameter, value.kind, value.value["path"]) for value in application.inputs ] == [("text", "column", "documents.body")] - with pytest.raises((TypeError, ValueError)): - application.group_id = "fg_changed" def test_function_version_binding_validates_names_and_direct_columns(): @@ -156,7 +137,7 @@ def test_function_version_binding_validates_names_and_direct_columns(): def test_function_version_keeps_named_struct_outputs_in_one_application(): value = job_result("remote_function_job.json") value["name"] = "text_features" - value["version"] = "fv_grouped" + value["version"] = "fv_multi_output" value["signature"] = { "inputs": [ {"name": "title", "arrow_type": "utf8", "nullable": True}, @@ -221,7 +202,6 @@ def test_function_application_uses_rename_columns_only(): assert application.columns["normalized_text"] == "search_text" assert renamed.columns["normalized_text"] == "body_normalized" assert renamed.function == application.function - assert renamed.group_id == application.group_id assert not hasattr(application, "rename_outputs") with pytest.raises(TypeError, match="immutable"): renamed.columns["normalized_text"] = "changed" @@ -242,7 +222,6 @@ def test_function_application_uses_rename_columns_only(): def test_binding_and_refresh_result_keep_stable_remote_fields(): binding = FunctionBinding.from_json(fixture("remote_function_binding.json")) - assert binding.revision == 3 assert binding.function.version == "fv_01K3TEXT" assert [output.output_ordinal for output in binding.outputs] == [0, 1] assert binding.input_schema is not None @@ -297,15 +276,6 @@ def test_refresh_result_rejects_non_u64_values(field): RefreshColumnResult.from_json(json.dumps(value)) -def test_canonical_client_values_contain_secret_names_only(): - version = FunctionVersion.from_json( - json.dumps(job_result("remote_function_job.json")) - ) - canonical = json.loads(version.to_canonical_json()) - assert canonical["required_secrets"] == ["HF_TOKEN"] - assert_no_secret_values(canonical) - - class _FunctionDeclarationInner: def __init__(self): self.calls = [] @@ -322,7 +292,7 @@ def known_application() -> FunctionApplication: @pytest.mark.asyncio -async def test_add_columns_routes_struct_as_one_and_grouped_expansion_atomically(): +async def test_add_columns_routes_struct_as_one_and_multi_output_binding_atomically(): inner = _FunctionDeclarationInner() table = AsyncTable(inner) application = known_application() @@ -343,12 +313,12 @@ async def test_add_columns_routes_struct_as_one_and_grouped_expansion_atomically @pytest.mark.asyncio -async def test_add_columns_rejects_mixed_groups_and_unknown_newer_application(): +async def test_add_columns_rejects_multiple_bindings_and_unknown_newer_application(): inner = _FunctionDeclarationInner() table = AsyncTable(inner) application = known_application() - with pytest.raises(ValueError, match="exactly one Function sibling group"): + with pytest.raises(ValueError, match="exactly one Function binding"): await table.add_columns({"a": application, "b": application}) future = json.loads(fixture("remote_function_application.json")) @@ -376,7 +346,6 @@ def test_rename_requires_named_struct_and_keeps_partial_mapping_immutable(): "arrow_type": "list", "nullable": False, }, - "group_id": "fg_scalar", } ) ) diff --git a/python/python/tests/test_first_class_function_slice2.py b/python/python/tests/test_first_class_function_slice2.py index a20347771..55257d322 100644 --- a/python/python/tests/test_first_class_function_slice2.py +++ b/python/python/tests/test_first_class_function_slice2.py @@ -3,7 +3,12 @@ from __future__ import annotations +import base64 import contextlib +import functools +import importlib.util +import types +from datetime import date import http.server import json from pathlib import Path @@ -16,6 +21,9 @@ import pytest import lancedb from lancedb.functions import UdfDefinition, udf +THRESHOLD = 20 +_CACHE = None + FIXTURES = ( Path(__file__).parents[3] @@ -31,28 +39,12 @@ FIXTURES = ( @udf( pip=["numpy>=2"], env={"MODE": "test"}, - secrets=["API_TOKEN"], python_version="3.12", ) def normalize_score(value: float) -> float: return value / 100.0 -def _assert_no_secret_values(value): - if isinstance(value, dict): - for key, child in value.items(): - assert key not in { - "secret_value", - "secret_values", - "resolved_secret", - "resolved_secrets", - } - _assert_no_secret_values(child) - elif isinstance(value, list): - for child in value: - _assert_no_secret_values(child) - - def test_scalar_udf_matches_shared_registration_golden_and_remains_callable(): assert isinstance(normalize_score, UdfDefinition) assert normalize_score(25.0) == 0.25 @@ -67,13 +59,397 @@ def test_scalar_udf_matches_shared_registration_golden_and_remains_callable(): "kind": "scalar_to_arrow_batch", "version": 1, } - assert request["required_secrets"] == ["API_TOKEN"] - _assert_no_secret_values(request) + + +def _run_packaged(definition, *args): + """Execute the shipped artifact in a fresh namespace, as a worker would.""" + source = base64.b64decode(definition.registration_request.artifact.content.data) + namespace: dict = {} + exec(compile(source, "", "exec"), namespace) + return namespace[definition.registration_request.artifact.entrypoint](*args) + + +def test_udf_packages_attribute_access_and_body_imports(): + @udf + def word_norm(body: str) -> float: + import numpy as np + + try: + words = body.split() + except AttributeError as error: + raise ValueError(str(error)) from error + return float(np.linalg.norm([len(w) for w in words])) + + assert _run_packaged(word_norm, "aa bb") == pytest.approx(8**0.5) + + +def test_udf_packages_module_globals_and_global_caches(): + @udf + def label(value: int) -> str: + return "big" if value >= THRESHOLD else "small" + + assert _run_packaged(label, 21) == "big" + + @udf + def cached(value: int) -> int: + global _CACHE + if _CACHE is None: + _CACHE = 40 + return _CACHE + value + + assert _run_packaged(cached, 2) == 42 + + +def test_udf_annotations_are_not_runtime_names(): + @udf + def identity(value: date) -> date: + return value + + assert _run_packaged(identity, date(2026, 8, 25)) == date(2026, 8, 25) + + +def test_udf_nested_scopes_resolve_lexically(): + @udf + def score(value: int) -> int: + offset = 2 + + def add_offset() -> int: + return value + offset + + return add_offset() + sum(v for v in [0]) + + assert _run_packaged(score, 3) == 5 + + +def test_udf_resolves_module_globals_before_builtins(tmp_path): + module_path = tmp_path / "shadowing_udfs.py" + module_path.write_text( + "max = 7\n" + "len = lambda _: 99\n" + "\n" + "def uses_literal_shadow(value: int) -> int:\n" + " def nested() -> int:\n" + " return max\n" + " return nested() + value\n" + "\n" + "def uses_callable_shadow(value: int) -> int:\n" + " def nested() -> int:\n" + " return len([1])\n" + " return nested() + value\n" + ) + spec = importlib.util.spec_from_file_location("shadowing_udfs", module_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + + # The module's `max = 7` is what the interpreter would use, so it ships. + assert _run_packaged(udf(module.uses_literal_shadow), 1) == 8 + # A callable global cannot ship; it must not be silently swapped for the builtin. + with pytest.raises(TypeError, match="unsupported global value of type function"): + udf(module.uses_callable_shadow) + + +def test_canonical_arrow_type_is_exactly_the_grammar(): + from lancedb.functions import _GRAMMAR_PRIMITIVES, _canonical_arrow_type + + golden = json.loads( + ( + Path(__file__).parents[3] + / "rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json" + ).read_text() + ) + primitives = [ + case["arrow_type"] for case in golden["valid"] if "<" not in case["arrow_type"] + ] + assert [name for _, name in _GRAMMAR_PRIMITIVES] == primitives + for outside in [ + pa.timestamp("us"), + pa.decimal128(10, 2), + pa.large_string(), + pa.large_binary(), + pa.binary(4), + pa.duration("s"), + pa.struct([pa.field("a", pa.int32())]), + pa.list_(pa.float32(), 0), + pa.list_(pa.timestamp("us")), + ]: + with pytest.raises(TypeError, match="unsupported Arrow type"): + _canonical_arrow_type(outside) + + +def test_udf_nested_annotations_are_postponed_in_the_artifact(): + @udf + def score(value: int) -> int: + def identity(item: date) -> date: + return item + + identity(date(2026, 8, 25)) + return value + + assert _run_packaged(score, 3) == 3 + + +def test_udf_ships_globals_the_body_deletes(): + @udf + def clear(value: int) -> int: + global _CACHE + del _CACHE + return value + + assert _run_packaged(clear, 3) == 3 + + +def test_udf_rejects_a_module_global_that_does_not_import_as_itself(tmp_path): + module_path = tmp_path / "fake_module_udfs.py" + module_path.write_text( + "import types\n" + "np = types.ModuleType('numpy')\n" + "np.sqrt = lambda x: 0\n" + "\n" + "def score(value: int) -> int:\n" + " return int(np.sqrt(value))\n" + ) + spec = importlib.util.spec_from_file_location("fake_module_udfs", module_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + with pytest.raises(TypeError, match="does not import as 'numpy'"): + udf(module.score) + + +def test_udf_rejects_a_module_level_namespace_alias(tmp_path): + module_path = tmp_path / "aliasing_udfs.py" + module_path.write_text( + "import builtins as b\n" + "THRESHOLD = 5\n" + "\n" + "def score(value: int) -> int:\n" + " return value + b.vars(b.__import__('aliasing_udfs'))['THRESHOLD']\n" + ) + spec = importlib.util.spec_from_file_location("aliasing_udfs", module_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + with pytest.raises(ValueError, match="dynamic namespace access"): + udf(module.score) + + +@pytest.mark.parametrize( + "access", + [ + "globals()['THRESHOLD']", + "eval('THRESHOLD')", + "(lambda g: g()['THRESHOLD'])(globals)", + "__import__('sys').modules[__name__].THRESHOLD", + "sys.modules[__name__].THRESHOLD", + ], +) +def test_udf_rejects_dynamic_namespace_access(access): + namespace: dict = {} + exec( + f"def score(value: int) -> int:\n return value + {access}\n", + {"THRESHOLD": 5}, + namespace, + ) + with pytest.raises(ValueError, match="dynamic namespace access"): + _package_from_text( + "def score(value: int) -> int:\n" + " import sys\n" + f" return value + {access}\n" + ) + + +def _package_from_text(source: str, module_globals: dict | None = None): + """Load `source` as a real module file so the packager can inspect it.""" + import tempfile + + directory = tempfile.mkdtemp() + path = Path(directory) / "generated_udf_module.py" + path.write_text(source) + spec = importlib.util.spec_from_file_location(f"generated_udf_{id(source)}", path) + module = importlib.util.module_from_spec(spec) + if module_globals: + module.__dict__.update(module_globals) + spec.loader.exec_module(module) + functions = [ + value + for value in vars(module).values() + if callable(value) and getattr(value, "__module__", None) == module.__name__ + ] + return udf(functions[0]) + + +def test_udf_rejects_a_non_standard_builtins_environment(): + def score(value: int) -> int: + return len([1]) + value + + score.__globals__ # noqa: B018 -- real function, real globals + import builtins + + patched = types.FunctionType( + score.__code__, + {"__builtins__": {**vars(builtins), "len": lambda _: 99}}, + "score", + ) + patched.__annotations__ = score.__annotations__ + assert patched(3) == 102 + with pytest.raises(ValueError, match="non-standard builtins environment"): + udf(patched) + + class ReportingDict(dict): # reports standard entries, resolves differently + def __missing__(self, key): + return vars(builtins)[key] + + disguised = types.FunctionType( + score.__code__, {"__builtins__": ReportingDict(len=lambda _: 99)}, "score" + ) + disguised.__annotations__ = score.__annotations__ + assert disguised(3) == 102 + with pytest.raises(ValueError, match="non-standard builtins environment"): + udf(disguised) + + hooked = types.FunctionType( + score.__code__, + {"__builtins__": {**vars(builtins), "__import__": lambda *a, **k: None}}, + "score", + ) + hooked.__annotations__ = score.__annotations__ + with pytest.raises(ValueError, match="non-standard builtins environment"): + udf(hooked) + + +def test_udf_recursion_versus_a_rebound_module_name(tmp_path): + module_path = tmp_path / "rebound_udfs.py" + module_path.write_text( + "def fact(value: int) -> int:\n" + " return 1 if value <= 1 else value * fact(value - 1)\n" + "\n" + "def score(value: int) -> int:\n" + " return score + value\n" + ) + spec = importlib.util.spec_from_file_location("rebound_udfs", module_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + assert _run_packaged(udf(module.fact), 5) == 120 + raw = module.score + module.score = 10 + with pytest.raises(ValueError, match="binds that name to another value"): + udf(raw) + # A wrapper that merely exposes __wrapped__ is not the function. + module.score = functools.wraps(raw)(lambda value: 41) + with pytest.raises(ValueError, match="binds that name to another value"): + udf(raw) + # The decorator's own result is; a subclass of it is not. + module.fact = udf(module.fact) + assert _run_packaged(module.fact, 4) == 24 + + class Twisted(UdfDefinition): + def __call__(self, *args, **kwargs): + return 41 + + raw_fact = module.fact._function + module.fact = Twisted( + raw_fact, + name=None, + input_schema=None, + output_schema=None, + pip=(), + env={}, + python_version=None, + ) + with pytest.raises(ValueError, match="binds that name to another value"): + udf(raw_fact) + + +def test_canonical_arrow_type_rejects_unrepresentable_list_children(): + from lancedb.functions import _canonical_arrow_type + + for outside in [ + pa.list_(pa.float32()), # pyarrow default: nullable child + pa.list_(pa.field("custom", pa.float32(), nullable=False)), + pa.list_(pa.field("item", pa.float32(), nullable=False, metadata={"k": "v"})), + pa.list_(pa.field("item", pa.float32(), nullable=False), 0), + ]: + with pytest.raises(TypeError, match="unsupported Arrow type"): + _canonical_arrow_type(outside) + assert ( + _canonical_arrow_type( + pa.list_(pa.field("item", pa.float32(), nullable=False), 3) + ) + == "fixed_size_list" + ) + + +def _calls_missing(value: int) -> int: + return missing(value) # noqa: F821 + + +def _shadows_missing_in_a_comprehension(value: int) -> int: + return missing(value) + sum(missing for missing in ()) # noqa: F821 + + +def _shadows_missing_in_a_lambda(value: int) -> int: + return (lambda missing: missing)(value) + missing # noqa: F821 + + +@pytest.mark.parametrize( + "function", + [_calls_missing, _shadows_missing_in_a_comprehension, _shadows_missing_in_a_lambda], +) +def test_udf_rejects_a_truly_unresolved_global(function): + with pytest.raises(ValueError, match=r"unresolved global names: \['missing'\]"): + udf(function) + + +def _arrow_type_from_golden(spec: dict) -> pa.DataType: + kind = spec["type"] + if kind in ("list", "large_list", "fixed_size_list"): + item = _arrow_type_from_golden(spec["fields"][0]["type"]) + field = pa.field("item", item, nullable=False) + if kind == "list": + return pa.list_(field) + if kind == "large_list": + return pa.large_list(field) + return pa.list_(field, spec["length"]) + return { + "null": pa.null(), + "bool": pa.bool_(), + "utf8": pa.string(), + "binary": pa.binary(), + "float16": pa.float16(), + "float32": pa.float32(), + "float64": pa.float64(), + "date32": pa.date32(), + "date64": pa.date64(), + }.get(kind) or getattr(pa, kind)() + + +def test_arrow_type_grammar_matches_the_shared_golden(): + golden = json.loads( + ( + Path(__file__).parents[3] + / "rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json" + ).read_text() + ) + from lancedb.functions import _canonical_arrow_type + + emitted = { + case["arrow_type"]: _canonical_arrow_type(_arrow_type_from_golden(case["json"])) + for case in golden["valid"] + } + assert emitted == { + case["arrow_type"]: case["arrow_type"] for case in golden["valid"] + } + assert not set(emitted) & set(golden["invalid"]) + for case in golden["server_only"]: + with pytest.raises(TypeError, match="unsupported Arrow type"): + _canonical_arrow_type(_arrow_type_from_golden(case["json"])) def test_explicit_arrow_schema_is_deterministic(): input_schema = pa.schema([pa.field("value", pa.float32(), nullable=True)]) - output_schema = pa.field("embedding", pa.list_(pa.float32(), 3), nullable=False) + output_schema = pa.field( + "embedding", + pa.list_(pa.field("item", pa.float32(), nullable=False), 3), + nullable=False, + ) @udf(input_schema=input_schema, output_schema=output_schema) def explicit(value): @@ -82,7 +458,7 @@ def test_explicit_arrow_schema_is_deterministic(): signature = explicit.registration_request.signature assert signature.inputs[0].arrow_type == "float32" assert signature.inputs[0].nullable is True - assert signature.output.arrow_type == "fixed_size_list[3]" + assert signature.output.arrow_type == "fixed_size_list" assert signature.output.nullable is False @@ -130,14 +506,6 @@ def test_annotation_and_explicit_schema_validation_fail_closed(): return value -def test_environment_rejects_secret_value_overlap(): - with pytest.raises(ValueError, match="must be disjoint"): - - @udf(env={"TOKEN": "plaintext"}, secrets=["TOKEN"]) - def overlapping(value: int) -> int: - return value - - def test_local_function_catalog_operations_are_not_supported(tmp_path): db = lancedb.connect(tmp_path) message = "Function catalog operations are not supported by this database" @@ -174,7 +542,6 @@ def _mock_remote_function_catalog(): "runtime": body["runtime"], "runtime_digest": "sha256:runtime", "environment_digest": "sha256:environment", - "required_secrets": body.get("required_secrets", []), "created_at": "2026-08-21T00:00:00Z", } response = {"job_id": "job-register"} @@ -233,7 +600,6 @@ def test_remote_registration_job_and_exact_version_reopen_round_trip(): assert create_request == json.loads( normalize_score.registration_request.to_canonical_json() ) - _assert_no_secret_values(create_request) def test_blocking_remote_registration_returns_function_version(): diff --git a/python/python/tests/test_permutation.py b/python/python/tests/test_permutation.py index 135742c84..142fc84f1 100644 --- a/python/python/tests/test_permutation.py +++ b/python/python/tests/test_permutation.py @@ -56,6 +56,31 @@ def test_execute_does_not_reenter_background_loop(tmp_path, monkeypatch): assert permutation_tbl._conn.read_consistency_interval is None +def test_pickled_permutation_reads_pinned_version(tmp_path): + """An unpickled copy must still read the pinned version, which also covers the + version surviving the ``to_arrow()`` round trip in ``__getstate__``.""" + import pickle + + db = connect(tmp_path) + tbl = db.create_table("base", pa.table({"idx": range(20)})) + permutation_tbl = permutation_builder(tbl).execute() + perm = Permutation.from_tables(tbl, permutation_tbl) + + payload = pickle.dumps(perm) + + # Compact so the stored row addresses no longer describe these rows at latest. + tbl.delete("true") + tbl.optimize() + assert tbl.count_rows() == 0 + + # Unpickle after the mutation: __setstate__ reopens at latest, so this only + # passes if the recorded version is applied on reopen. + restored = pickle.loads(payload) + assert len(restored) == 20 + rows = restored.__getitems__(list(range(20))) + assert sorted(row["idx"] for row in rows) == list(range(20)) + + def test_split_random_counts(mem_db): """Test random splitting with absolute counts.""" tbl = mem_db.create_table( diff --git a/python/src/table.rs b/python/src/table.rs index b225b191f..784d29136 100644 --- a/python/src/table.rs +++ b/python/src/table.rs @@ -780,15 +780,19 @@ impl Table { }) } - #[pyo3(signature = (data, mode, progress=None, write_parallelism=None))] + #[pyo3(signature = (data, mode, progress=None, write_parallelism=None, allow_external_blob_outside_bases=false))] pub fn add<'a>( self_: PyRef<'a, Self>, data: PyScannable, mode: String, progress: Option>, write_parallelism: Option, + allow_external_blob_outside_bases: bool, ) -> PyResult> { - let mut op = self_.inner_ref()?.add(data); + let mut op = self_ + .inner_ref()? + .add(data) + .allow_external_blob_outside_bases(allow_external_blob_outside_bases); if mode == "append" { op = op.mode(AddDataMode::Append); } else if mode == "overwrite" { diff --git a/rust/lancedb/Cargo.toml b/rust/lancedb/Cargo.toml index d32d88ecd..8276e5bb3 100644 --- a/rust/lancedb/Cargo.toml +++ b/rust/lancedb/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb" -version = "0.38.0-beta.6" +version = "0.38.0-beta.10" edition.workspace = true description = "LanceDB: A serverless, low-latency vector database for AI applications" license.workspace = true diff --git a/rust/lancedb/src/dataloader/permutation/builder.rs b/rust/lancedb/src/dataloader/permutation/builder.rs index a3700d0ef..ba0d0e180 100644 --- a/rust/lancedb/src/dataloader/permutation/builder.rs +++ b/rust/lancedb/src/dataloader/permutation/builder.rs @@ -27,6 +27,12 @@ pub const SRC_ROW_ID_COL: &str = "row_id"; pub const SPLIT_NAMES_CONFIG_KEY: &str = "split_names"; +/// Base table version the permutation was built against. +pub const BASE_VERSION_CONFIG_KEY: &str = "base_version"; + +/// Base table branch the permutation was built against. Absent means main. +pub const BASE_BRANCH_CONFIG_KEY: &str = "base_branch"; + pub const DEFAULT_MEMORY_LIMIT: usize = 100 * 1024 * 1024; /// Where to store the permutation table @@ -214,21 +220,11 @@ impl PermutationBuilder { Ok(Box::pin(SimpleRecordBatchStream { schema, stream })) } - fn add_split_names( + fn add_config_metadata( data: SendableRecordBatchStream, - split_names: &[String], + metadata: HashMap, ) -> Result { - let schema = data - .schema() - .as_ref() - .clone() - .with_metadata(HashMap::from([( - SPLIT_NAMES_CONFIG_KEY.to_string(), - serde_json::to_string(split_names).map_err(|e| Error::Other { - message: format!("Failed to serialize split names: {}", e), - source: Some(e.into()), - })?, - )])); + let schema = data.schema().as_ref().clone().with_metadata(metadata); let schema = Arc::new(schema); let schema_clone = schema.clone(); let stream = data.map_ok(move |batch| batch.with_schema(schema.clone()).unwrap()); @@ -239,7 +235,20 @@ impl PermutationBuilder { } /// Builds the permutation table and stores it in the given database. - pub async fn build(self) -> Result { + pub async fn build(mut self) -> Result
{ + // Remote tables resolve latest independently for each request. Use a + // separate pinned handle so count, projection, and scan all refer to one + // snapshot without changing the caller's table checkout state. Native + // tables return `None` here and retain their existing behavior. + if let Some(snapshot) = self + .base_table + .base_table() + .snapshot_at_current_version() + .await? + { + self.base_table = Table::from(snapshot); + } + // Unflushed rows have no row id, so a permutation cannot address them. match self.base_table.base_table().get_lsm_write_spec().await { Ok(Some(_)) => { @@ -256,9 +265,14 @@ impl PermutationBuilder { Err(err) => return Err(err), } + // The handle above is already pinned to one version. Record which one, so a + // reader -- in a DataLoader worker, against a table that has since moved -- + // resolves these row addresses against the same snapshot. + let base_version = self.base_table.version().await?; + let base_branch = self.base_table.current_branch(); + // First pass, apply filter and load row ids. `Shuffler` permutes positions, so // every rank must scan the rows in the same order to build the same permutation. - // TODO: pin the version resolved here; remote does not implement Lazy. let mut rows = self.base_table.query().select(Select::columns(&[ROW_ID])); if let Some(filter) = &self.config.filter { @@ -318,11 +332,24 @@ impl PermutationBuilder { // Rename _rowid to row_id let renamed = rename_column(sorted, ROW_ID, SRC_ROW_ID_COL)?; - let streaming_data = if let Some(split_names) = &self.config.split_names { - Self::add_split_names(renamed, split_names)? - } else { - renamed - }; + let mut metadata = HashMap::from([( + BASE_VERSION_CONFIG_KEY.to_string(), + base_version.to_string(), + )]); + // Version numbers are per-branch, so the branch is part of the coordinate. + if let Some(branch) = &base_branch { + metadata.insert(BASE_BRANCH_CONFIG_KEY.to_string(), branch.clone()); + } + if let Some(split_names) = &self.config.split_names { + metadata.insert( + SPLIT_NAMES_CONFIG_KEY.to_string(), + serde_json::to_string(split_names).map_err(|e| Error::Other { + message: format!("Failed to serialize split names: {}", e), + source: Some(e.into()), + })?, + ); + } + let streaming_data = Self::add_config_metadata(renamed, metadata)?; let (name, database) = match &self.config.destination { PermutationDestination::Permanent(database, table_name) => { @@ -409,6 +436,253 @@ mod tests { assert!(table.base_table().scan_order_is_deterministic()); } + #[cfg(feature = "remote")] + #[tokio::test] + async fn test_remote_permutation_builder_pins_snapshot() { + use std::sync::{ + Mutex, + atomic::{AtomicU64, Ordering}, + }; + + use arrow_array::{RecordBatch, UInt64Array}; + use arrow_schema::{DataType, Field, Schema}; + + let row_ids = RecordBatch::try_new( + Arc::new(Schema::new(vec![Field::new( + ROW_ID, + DataType::UInt64, + false, + )])), + vec![Arc::new(UInt64Array::from(vec![100]))], + ) + .unwrap(); + let mut query_body = Vec::new(); + { + let mut writer = + arrow_ipc::writer::FileWriter::try_new(&mut query_body, &row_ids.schema()).unwrap(); + writer.write(&row_ids).unwrap(); + writer.finish().unwrap(); + } + + let latest = Arc::new(AtomicU64::new(7)); + let expected_snapshot = Arc::new(AtomicU64::new(7)); + let planning_versions = Arc::new(Mutex::new(Vec::new())); + let latest_ref = latest.clone(); + let expected_snapshot_ref = expected_snapshot.clone(); + let planning_versions_ref = planning_versions.clone(); + let table = Table::new_with_handler("remote_base", move |request| { + let path = request.url().path(); + let body = request + .body() + .and_then(|body| body.as_bytes()) + .map(|body| serde_json::from_slice::(body).unwrap()); + + match path { + "/v1/table/remote_base/describe/" => { + let requested = body.as_ref().and_then(|body| body["version"].as_u64()); + let version = requested.unwrap_or_else(|| latest_ref.load(Ordering::SeqCst)); + http::Response::builder() + .status(200) + .body( + format!(r#"{{"version":{version},"schema":{{"fields":[]}}}}"#) + .into_bytes(), + ) + .unwrap() + } + "/v1/table/remote_base/get_lsm_write_spec/" => http::Response::builder() + .status(200) + .body(br#"{"lsm_write_spec":null}"#.to_vec()) + .unwrap(), + "/v1/table/remote_base/count_rows/" => { + let body = body.unwrap(); + let version = body["version"].as_u64().unwrap(); + assert_eq!(version, expected_snapshot_ref.load(Ordering::SeqCst)); + assert_eq!(body["predicate"], "value > 0"); + planning_versions_ref.lock().unwrap().push(version); + + // Simulate a concurrent append after count_rows. An unpinned + // scan would now resolve version 8 and include different rows. + latest_ref.store(8, Ordering::SeqCst); + http::Response::builder() + .status(200) + .body(b"1".to_vec()) + .unwrap() + } + "/v1/table/remote_base/query/" => { + let body = body.unwrap(); + let version = body["version"].as_u64().unwrap(); + assert_eq!(version, expected_snapshot_ref.load(Ordering::SeqCst)); + assert_eq!(body["filter"], "value > 0"); + assert_eq!(body["columns"], serde_json::json!([ROW_ID])); + planning_versions_ref.lock().unwrap().push(version); + http::Response::builder() + .status(200) + .header("content-type", "application/vnd.apache.arrow.file") + .body(query_body.clone()) + .unwrap() + } + _ => panic!("unexpected request: {path}"), + } + }); + + let permutation = PermutationBuilder::new(table.clone()) + .with_filter("value > 0".to_string()) + .build() + .await + .unwrap(); + assert_eq!(permutation.count_rows(None).await.unwrap(), 1); + + // Building uses a separate handle and must not pin the caller's table. + assert_eq!(table.version().await.unwrap(), 8); + + // An explicit checkout is copied as-is and remains checked out afterward. + expected_snapshot.store(6, Ordering::SeqCst); + table.checkout(6).await.unwrap(); + let permutation = PermutationBuilder::new(table.clone()) + .with_filter("value > 0".to_string()) + .build() + .await + .unwrap(); + assert_eq!(permutation.count_rows(None).await.unwrap(), 1); + assert_eq!(table.version().await.unwrap(), 6); + assert_eq!(*planning_versions.lock().unwrap(), vec![7, 7, 6, 6]); + } + + #[tokio::test] + async fn test_permutation_records_base_version() { + let temp_dir = tempfile::tempdir().unwrap(); + + let db = connect(temp_dir.path().to_str().unwrap()) + .execute() + .await + .unwrap(); + + let initial_data = lance_datagen::gen_batch() + .col("col_a", lance_datagen::array::step::()) + .into_ldb_stream(RowCount::from(100), BatchCount::from(2)); + let data_table = db + .create_table("base_tbl", initial_data) + .execute() + .await + .unwrap(); + + let build_version = data_table.version().await.unwrap(); + let permutation_table = PermutationBuilder::new(data_table.clone()) + .build() + .await + .unwrap(); + + let recorded = permutation_table + .schema() + .await + .unwrap() + .metadata + .get(BASE_VERSION_CONFIG_KEY) + .expect("permutation should record the base version") + .parse::() + .unwrap(); + assert_eq!(recorded, build_version); + + // Advancing the base table must not move the recorded version. + let more_data = lance_datagen::gen_batch() + .col("col_a", lance_datagen::array::step::()) + .into_ldb_stream(RowCount::from(50), BatchCount::from(1)); + data_table.add(more_data).execute().await.unwrap(); + assert!(data_table.version().await.unwrap() > recorded); + assert_eq!( + permutation_table + .schema() + .await + .unwrap() + .metadata + .get(BASE_VERSION_CONFIG_KEY) + .unwrap() + .parse::() + .unwrap(), + recorded, + ); + } + + /// Version numbers are per-branch, so a permutation built on a branch must record + /// it -- a worker reopens by name and lands on main at the same number. + #[tokio::test] + async fn test_permutation_records_base_branch() { + let temp_dir = tempfile::tempdir().unwrap(); + let db = connect(temp_dir.path().to_str().unwrap()) + .execute() + .await + .unwrap(); + + let initial_data = lance_datagen::gen_batch() + .col("col_a", lance_datagen::array::step::()) + .into_ldb_stream(RowCount::from(10), BatchCount::from(1)); + let data_table = db + .create_table("base_tbl", initial_data) + .execute() + .await + .unwrap(); + + let branch = data_table + .create_branch("exp", lance::dataset::refs::Ref::from(("main", 1))) + .await + .unwrap(); + let permutation_table = PermutationBuilder::new(branch.clone()) + .build() + .await + .unwrap(); + + let metadata = permutation_table.schema().await.unwrap().metadata.clone(); + assert_eq!( + metadata.get(BASE_BRANCH_CONFIG_KEY).map(String::as_str), + Some("exp") + ); + + // Main records nothing, so an absent key keeps meaning main. + let main_permutation = PermutationBuilder::new(data_table.clone()) + .build() + .await + .unwrap(); + assert!( + !main_permutation + .schema() + .await + .unwrap() + .metadata + .contains_key(BASE_BRANCH_CONFIG_KEY) + ); + } + + #[tokio::test] + async fn test_build_does_not_pin_the_callers_table() { + let temp_dir = tempfile::tempdir().unwrap(); + + let db = connect(temp_dir.path().to_str().unwrap()) + .execute() + .await + .unwrap(); + + let initial_data = lance_datagen::gen_batch() + .col("col_a", lance_datagen::array::step::()) + .into_ldb_stream(RowCount::from(100), BatchCount::from(1)); + let data_table = db + .create_table("base_tbl", initial_data) + .execute() + .await + .unwrap(); + + PermutationBuilder::new(data_table.clone()) + .build() + .await + .unwrap(); + + // The builder pins its own handle; the caller's must still track latest. + let more_data = lance_datagen::gen_batch() + .col("col_a", lance_datagen::array::step::()) + .into_ldb_stream(RowCount::from(50), BatchCount::from(1)); + data_table.add(more_data).execute().await.unwrap(); + assert_eq!(data_table.count_rows(None).await.unwrap(), 150); + } + #[tokio::test] async fn test_permutation_builder() { let temp_dir = tempfile::tempdir().unwrap(); diff --git a/rust/lancedb/src/dataloader/permutation/reader.rs b/rust/lancedb/src/dataloader/permutation/reader.rs index 03dacb9c4..015c3a17d 100644 --- a/rust/lancedb/src/dataloader/permutation/reader.rs +++ b/rust/lancedb/src/dataloader/permutation/reader.rs @@ -8,7 +8,9 @@ //! the rows from a source table that correspond to row IDs stored in a separate table. use crate::arrow::{SendableRecordBatchStream, SimpleRecordBatchStream}; -use crate::dataloader::permutation::builder::SRC_ROW_ID_COL; +use crate::dataloader::permutation::builder::{ + BASE_BRANCH_CONFIG_KEY, BASE_VERSION_CONFIG_KEY, SRC_ROW_ID_COL, +}; use crate::dataloader::permutation::split::SPLIT_ID_COLUMN; use crate::error::Error; use crate::query::{ @@ -23,6 +25,7 @@ use arrow_array::{RecordBatch, UInt64Array}; use arrow_schema::SchemaRef; use datafusion_expr::{Expr, col, lit}; use futures::{StreamExt, TryStreamExt}; +use lance::dataset::refs::MAIN_BRANCH; use lance::dataset::scanner::DatasetRecordBatchStream; use lance::io::RecordBatchStream; use lance_arrow::RecordBatchExt; @@ -69,6 +72,10 @@ impl PermutationReader { permutation_table: Option>, split: u64, ) -> Result { + let base_table = match &permutation_table { + Some(permutation_table) => Self::pin_base_table(base_table, permutation_table).await?, + None => base_table, + }; let mut slf = Self { base_table, permutation_table, @@ -89,6 +96,34 @@ impl PermutationReader { Ok(slf) } + /// Pins the base table to the version the permutation was built against. + /// Permutations written before that was recorded carry no key and stay unpinned. + async fn pin_base_table( + base_table: Arc, + permutation_table: &Arc, + ) -> Result> { + let schema = permutation_table.schema().await?; + let Some(raw) = schema.metadata.get(BASE_VERSION_CONFIG_KEY) else { + return Ok(base_table); + }; + let version = raw.parse::().map_err(|e| Error::InvalidInput { + message: format!( + "Permutation table has an unreadable {} of {:?}: {}", + BASE_VERSION_CONFIG_KEY, raw, e + ), + })?; + // The recorded branch, not the handle's: a worker reopens by name and lands + // on main, and version numbers are per-branch. + let branch = schema + .metadata + .get(BASE_BRANCH_CONFIG_KEY) + .map(String::as_str) + .unwrap_or(MAIN_BRANCH); + base_table + .checkout_branch_version(branch, Some(version)) + .await + } + pub async fn try_from_tables( base_table: Arc, permutation_table: Arc, @@ -519,9 +554,13 @@ mod tests { use lance_datagen::{BatchCount, RowCount}; use rand::seq::SliceRandom; + // Aliased: `test_utils::datagen` exports a trait of the same name. + use crate::arrow::LanceDbDatagenExt as _; use crate::{ Table, arrow::SendableRecordBatchStream, + connect, + dataloader::permutation::builder::PermutationBuilder, query::{ExecutableQuery, QueryBase}, test_utils::datagen::{LanceDbDatagenExt, virtual_table}, }; @@ -553,6 +592,58 @@ mod tests { .await } + /// Compaction moves row addresses, so the reader must read the pinned version. + #[tokio::test] + async fn test_reader_pins_base_version() { + let temp_dir = tempfile::tempdir().unwrap(); + let db = connect(temp_dir.path().to_str().unwrap()) + .execute() + .await + .unwrap(); + + let data = lance_datagen::gen_batch() + .col("idx", lance_datagen::array::step::()) + .into_ldb_stream(RowCount::from(20), BatchCount::from(1)); + let base_table = db.create_table("base_tbl", data).execute().await.unwrap(); + + let permutation_table = PermutationBuilder::new(base_table.clone()) + .build() + .await + .unwrap(); + + base_table.delete("true").await.unwrap(); + base_table + .optimize(crate::table::OptimizeAction::All) + .await + .unwrap(); + assert_eq!(base_table.count_rows(None).await.unwrap(), 0); + + let reader = PermutationReader::try_from_tables( + base_table.base_table().clone(), + permutation_table.base_table().clone(), + 0, + ) + .await + .unwrap(); + + let values = collect_from_stream::( + reader + .read( + Select::Columns(vec!["idx".to_string()]), + QueryExecutionOptions::default(), + ) + .await + .unwrap(), + "idx", + ) + .await; + assert_eq!( + values.len(), + 20, + "reader should still see the pinned version" + ); + } + #[tokio::test] async fn test_permutation_reader() { let base_table = lance_datagen::gen_batch() diff --git a/rust/lancedb/src/function.rs b/rust/lancedb/src/function.rs index f02158871..52c70a4b1 100644 --- a/rust/lancedb/src/function.rs +++ b/rust/lancedb/src/function.rs @@ -5,7 +5,7 @@ //! backend-neutral terminal result of a computed-column refresh. //! //! This module contains client/wire values only. Catalog persistence, -//! environment bake, secret resolution, and execution are owned by Sophon. +//! environment bake, and execution are owned by Sophon. use std::collections::BTreeMap; @@ -195,9 +195,6 @@ pub struct PythonEnvironmentSpec { } /// Reproducible Python runtime definition understood by Sophon. -/// -/// `env` contains non-secret values. Secret values have no client model; -/// [`FunctionVersion::required_secrets`] contains names only. #[derive(Debug, Clone, PartialEq, Eq)] #[non_exhaustive] pub enum PythonRuntimeSpec { @@ -239,7 +236,7 @@ impl PythonRuntimeSpec { } } - /// Non-secret environment variables, or `None` for an unknown kind. + /// Environment variables, or `None` for an unknown kind. pub fn env(&self) -> Option<&BTreeMap> { match self { Self::Python { env, .. } => Some(env), @@ -324,8 +321,6 @@ pub struct FunctionVersion { runtime: PythonRuntimeSpec, runtime_digest: String, environment_digest: String, - #[serde(default, skip_serializing_if = "Vec::is_empty")] - required_secrets: Vec, created_at: String, } @@ -358,11 +353,6 @@ impl FunctionVersion { &self.environment_digest } - /// Required secret names. Resolved values exist only inside Sophon. - pub fn required_secrets(&self) -> &[String] { - &self.required_secrets - } - pub fn created_at(&self) -> &str { &self.created_at } @@ -404,18 +394,12 @@ pub struct FunctionArtifactRequest { } /// Stable request envelope for remote immutable Function registration. -/// -/// Secret values deliberately have no field in this model. The only secret -/// material the client may send is the ordered set of names Sophon resolves -/// inside the remote runtime. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct FunctionRegistrationRequest { pub name: String, pub artifact: FunctionArtifactRequest, pub signature: FunctionSignature, pub runtime: PythonRuntimeSpec, - #[serde(default, skip_serializing_if = "Vec::is_empty")] - pub required_secrets: Vec, } impl_json!(FunctionRegistrationRequest); @@ -446,7 +430,6 @@ pub struct FunctionApplication { function: FunctionVersionRef, inputs: Vec, output: FunctionOutput, - group_id: String, #[serde(default, skip_serializing_if = "BTreeMap::is_empty")] columns: BTreeMap, #[serde(default, flatten, skip_serializing)] @@ -468,10 +451,6 @@ impl FunctionApplication { &self.output } - pub fn group_id(&self) -> &str { - &self.group_id - } - pub fn columns(&self) -> &BTreeMap { &self.columns } @@ -513,7 +492,7 @@ pub struct InputBinding { pub nullable: bool, } -/// Ordered result-field to table-field mapping for a grouped binding. +/// Ordered result-field to table-field mapping for a Function binding. /// /// Assignment state is not part of the Slice 1 client contract. During the /// NULL transition there is no public Lance cell-flag identifier to persist. @@ -527,20 +506,18 @@ pub struct OutputMapping { pub nullable: bool, } -/// Immutable grouped binding persisted by the Enterprise table service. +/// Immutable Function binding persisted by the Enterprise table service. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct FunctionBinding { binding_id: String, - revision: u64, function: FunctionVersionRef, - group_id: String, inputs: Vec, outputs: Vec, /// Exact Arrow schema presented to the Function, encoded with the Lance /// Namespace Arrow JSON representation. #[serde(default, skip_serializing_if = "Option::is_none")] input_schema: Option, - /// Exact physical Arrow schema of the grouped table outputs. + /// Exact physical Arrow schema of the binding's table outputs. #[serde(default, skip_serializing_if = "Option::is_none")] output_schema: Option, } @@ -550,18 +527,10 @@ impl FunctionBinding { &self.binding_id } - pub fn revision(&self) -> u64 { - self.revision - } - pub fn function(&self) -> &FunctionVersionRef { &self.function } - pub fn group_id(&self) -> &str { - &self.group_id - } - pub fn inputs(&self) -> &[InputBinding] { &self.inputs } 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/materialized_view.rs b/rust/lancedb/src/materialized_view.rs index 1d7cbb5e7..b28d52931 100644 --- a/rust/lancedb/src/materialized_view.rs +++ b/rust/lancedb/src/materialized_view.rs @@ -37,6 +37,13 @@ pub use refresh::{RefreshMaterializedViewResult, RefreshMode}; /// Schema metadata key holding the view definition, as kind-tagged JSON. pub const DEFINITION_META_KEY: &str = "mv.definition"; +/// Schema metadata key holding the view's incarnation: a token minted at each +/// physical creation of a view table, so a view dropped and recreated under +/// the same name and definition is still told apart from the one a caller +/// captured. A view whose metadata was replaced wholesale, or one declared +/// before tokens existed, carries none until its next refresh mints one. +pub const INCARNATION_META_KEY: &str = "mv.incarnation"; + /// Schema metadata key holding the source table version the view was last /// refreshed to. Absent until the first refresh. pub const SOURCE_VERSION_META_KEY: &str = "mv.source_version"; @@ -612,8 +619,17 @@ impl PreparedDeclaration { pub async fn create(self, name: &str) -> Result { let empty: Vec> = vec![]; + // Minted here, not at preparation: a declaration can be cloned and + // create more than one physical table, and each needs its own token. + let incarnation = uuid::Uuid::new_v4().to_string(); + let mut metadata = self.schema.metadata().clone(); + metadata.insert(INCARNATION_META_KEY.to_string(), incarnation.clone()); + let schema = Arc::new(ArrowSchema::new_with_metadata( + self.schema.fields().clone(), + metadata, + )); let reader: Box = - Box::new(arrow_array::RecordBatchIterator::new(empty, self.schema)); + Box::new(arrow_array::RecordBatchIterator::new(empty, schema)); let mut request = CreateTableRequest::new(name.to_string(), Box::new(reader)); let write_params = request .write_options @@ -648,6 +664,7 @@ impl PreparedDeclaration { Ok(MaterializedView { table, definition: self.definition, + incarnation: Some(incarnation), }) } } @@ -878,6 +895,7 @@ impl CreateMaterializedViewBuilder { pub struct MaterializedView { table: Table, definition: MaterializedViewDefinition, + incarnation: Option, } impl MaterializedView { @@ -893,8 +911,13 @@ impl MaterializedView { }); } let schema = table.schema().await?; + let incarnation = schema.metadata().get(INCARNATION_META_KEY).cloned(); match materialized_view_kind(schema.metadata())? { - Some(MaterializedViewKind::Select(definition)) => Ok(Self { table, definition }), + Some(MaterializedViewKind::Select(definition)) => Ok(Self { + table, + definition, + incarnation, + }), Some(MaterializedViewKind::Unrecognized { kind }) => Err(Error::NotSupported { message: format!( "materialized view '{}' is defined by '{kind}', which this version of \ @@ -923,6 +946,13 @@ impl MaterializedView { &self.definition } + /// The view's incarnation token as of when this handle was opened; see + /// [`RefreshMaterializedViewBuilder::expect_incarnation`]. `None` for a + /// view that has none yet (see [`INCARNATION_META_KEY`]). + pub fn incarnation(&self) -> Option<&str> { + self.incarnation.as_deref() + } + /// Recompute the view from its source. /// /// By default the refresh is incremental when the source's changes can be @@ -943,6 +973,7 @@ impl MaterializedView { view: self.clone(), full: false, source_version: None, + expected_incarnation: None, } } } @@ -952,6 +983,7 @@ pub struct RefreshMaterializedViewBuilder { view: MaterializedView, full: bool, source_version: Option, + expected_incarnation: Option, } impl RefreshMaterializedViewBuilder { @@ -967,8 +999,28 @@ impl RefreshMaterializedViewBuilder { self } + /// Refresh only if the view is still the incarnation that minted `token` + /// (see [`MaterializedView::incarnation`]): a refresh requested against + /// one declaration must not land in a view dropped and recreated since, + /// even under the same name and definition. + /// + /// Best effort. The token is read from the latest stored manifest before + /// planning and again immediately before every commit, but it is not part + /// of the commit's own condition, so a recreation that lands between that + /// final read and the commit is not caught. + pub fn expect_incarnation(mut self, token: impl Into) -> Self { + self.expected_incarnation = Some(token.into()); + self + } + pub async fn execute(self) -> Result { - refresh::execute_refresh(&self.view.table, self.full, self.source_version).await + refresh::execute_refresh( + &self.view.table, + self.full, + self.source_version, + self.expected_incarnation.as_deref(), + ) + .await } } diff --git a/rust/lancedb/src/materialized_view/refresh.rs b/rust/lancedb/src/materialized_view/refresh.rs index 24d828e01..735751c27 100644 --- a/rust/lancedb/src/materialized_view/refresh.rs +++ b/rust/lancedb/src/materialized_view/refresh.rs @@ -46,8 +46,8 @@ use lance_table::format::Fragment; use serde::{Deserialize, Serialize}; use super::{ - MaterializedViewDefinition, REFRESHED_AT_MS_META_KEY, SOURCE_ROW_ID_COLUMN, - SOURCE_VERSION_META_KEY, + INCARNATION_META_KEY, MaterializedViewDefinition, REFRESHED_AT_MS_META_KEY, + SOURCE_ROW_ID_COLUMN, SOURCE_VERSION_META_KEY, }; use crate::database::OpenTableRequest; use crate::table::{NativeTable, NativeTableExt, Table}; @@ -108,6 +108,7 @@ pub(crate) async fn execute_refresh( view: &Table, full: bool, pinned: Option, + expected_incarnation: Option<&str>, ) -> Result { let view_native = view.as_native().ok_or_else(|| Error::NotSupported { message: "materialized views are supported only on local tables".into(), @@ -122,6 +123,8 @@ pub(crate) async fn execute_refresh( view_native.dataset.reload().await?; let view_ds = view_native.dataset.get().await?.as_ref().clone(); + ensure_incarnation(&view_ds, expected_incarnation, view.name()).await?; + // The definition a handle cached at open may since have been replaced; // what refresh executes and what it stamps must be one generation. let definition = match super::materialized_view_kind(&view_ds.schema().metadata)? { @@ -240,6 +243,7 @@ pub(crate) async fn execute_refresh( increment, definition, watermark, + expected_incarnation, ) .await?; match reconciled { @@ -253,6 +257,7 @@ pub(crate) async fn execute_refresh( source_version, source_ts, definition, + expected_incarnation, ) .await } @@ -266,6 +271,7 @@ pub(crate) async fn execute_refresh( source_version, source_ts, definition, + expected_incarnation, ) .await } @@ -595,6 +601,7 @@ async fn incremental( increment: Increment, definition: &MaterializedViewDefinition, watermark: Option, + expected_incarnation: Option<&str>, ) -> Result> { let new_fragments = increment.appended; let watermark_version = watermark.unwrap_or(0); @@ -671,15 +678,35 @@ async fn incremental( }; let nothing_to_add = (new_fragments.is_empty() && !updated_rows) || remaining == Some(0); if nothing_to_add && eviction.is_none() { - result.version = - stamp_watermark(view_native, view_ds.clone(), source_version, source_ts).await?; + result.version = stamp_watermark( + view_native, + view_ds.clone(), + source_version, + source_ts, + expected_incarnation, + ) + .await?; return Ok(Some(result)); } // Rows left but none arrive: the removals still have to be published. if nothing_to_add { let filter = refresh_filter(&empty_keys(view_ds)?)?; - let published = publish(view_ds, eviction, Vec::new(), Some(filter)).await?; - result.version = stamp_watermark(view_native, published, source_version, source_ts).await?; + let published = publish( + view_ds, + eviction, + Vec::new(), + Some(filter), + expected_incarnation, + ) + .await?; + result.version = stamp_watermark( + view_native, + published, + source_version, + source_ts, + expected_incarnation, + ) + .await?; return Ok(Some(result)); } @@ -737,12 +764,20 @@ async fn incremental( eviction, Vec::new(), Some(refresh_filter(&empty_keys(view_ds)?)?), + expected_incarnation, ) .await? } else { view_ds.clone() }; - result.version = stamp_watermark(view_native, published, source_version, source_ts).await?; + result.version = stamp_watermark( + view_native, + published, + source_version, + source_ts, + expected_incarnation, + ) + .await?; return Ok(Some(result)); }; let stream: SendableRecordBatchStream = Box::pin(RecordBatchStreamAdapter::new( @@ -775,9 +810,23 @@ async fn incremental( }); }; let filter = refresh_filter(&keys)?; - let appended = publish(view_ds, eviction, new_fragments, Some(filter)).await?; + let appended = publish( + view_ds, + eviction, + new_fragments, + Some(filter), + expected_incarnation, + ) + .await?; result.rows_written = rows_written.load(Ordering::Relaxed); - result.version = stamp_watermark(view_native, appended, source_version, source_ts).await?; + result.version = stamp_watermark( + view_native, + appended, + source_version, + source_ts, + expected_incarnation, + ) + .await?; Ok(Some(result)) } @@ -788,6 +837,7 @@ async fn rebuild( source_version: u64, source_ts: u128, definition: &MaterializedViewDefinition, + expected_incarnation: Option<&str>, ) -> Result { let rows_written = Arc::new(AtomicU64::new(0)); let schema = Arc::new(ArrowSchema::from(view_ds.schema())); @@ -810,8 +860,16 @@ async fn rebuild( // carries no schema metadata, so it cannot erase a definition update // that raced in the way an overwrite (which adopts its stream's schema) // durably would -- and it must land on the planned generation or abort. - let replaced = replace_retaining_indices(view_ds.clone(), stream, keys).await?; - let version = stamp_watermark(view_native, replaced, source_version, source_ts).await?; + let replaced = + replace_retaining_indices(view_ds.clone(), stream, keys, expected_incarnation).await?; + let version = stamp_watermark( + view_native, + replaced, + source_version, + source_ts, + expected_incarnation, + ) + .await?; Ok(RefreshMaterializedViewResult { mode: RefreshMode::Rebuild, rows_written: rows_written.load(Ordering::Relaxed), @@ -828,11 +886,13 @@ async fn replace_retaining_indices( view_ds: Dataset, stream: SendableRecordBatchStream, keys: Arc>, + expected_incarnation: Option<&str>, ) -> Result { let ds = Arc::new(view_ds); let read_version = ds.version().version; #[cfg(test)] tests::hold_before_publish(ds.uri()).await; + ensure_incarnation(&ds, expected_incarnation, ds.uri()).await?; let removed_fragment_ids: Vec = ds.get_fragments().iter().map(|f| f.id() as u64).collect(); let write_txn = InsertBuilder::new(WriteDestination::Dataset(ds.clone())) @@ -886,6 +946,32 @@ async fn replace_retaining_indices( } /// Record that the view now reflects `source_version`, including the view +/// Refuse to act on a view that is not `expected`'s incarnation, judged from +/// the latest stored manifest. Not a commit condition; see +/// `RefreshMaterializedViewBuilder::expect_incarnation`. +async fn ensure_incarnation(view_ds: &Dataset, expected: Option<&str>, what: &str) -> Result<()> { + let Some(expected) = expected else { + return Ok(()); + }; + let mut latest = view_ds.clone(); + latest.checkout_latest().await?; + match latest.schema().metadata.get(INCARNATION_META_KEY) { + Some(actual) if actual == expected => Ok(()), + Some(_) => Err(Error::Runtime { + message: format!( + "materialized view '{what}' is not the incarnation this refresh was \ + requested for: it was dropped and recreated" + ), + }), + None => Err(Error::Runtime { + message: format!( + "materialized view '{what}' carries no incarnation token: its schema \ + metadata was replaced since the token was captured" + ), + }), + } +} + /// version this very commit produces. The version is predicted and then /// verified; on a mismatch another commit raced in between, and the stamp /// ABORTS rather than certify that commit as the refresh's own generation. @@ -895,10 +981,21 @@ async fn stamp_watermark( mut dataset: Dataset, source_version: u64, source_ts: u128, + expected_incarnation: Option<&str>, ) -> Result { + ensure_incarnation(&dataset, expected_incarnation, dataset.uri()).await?; let predicted = dataset.version().version + 1; + // A view with no token (declared before tokens existed, or its metadata + // replaced wholesale) starts a new incarnation here. + let incarnation = dataset + .schema() + .metadata + .get(INCARNATION_META_KEY) + .cloned() + .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); dataset .update_schema_metadata([ + (INCARNATION_META_KEY.to_string(), Some(incarnation)), ( SOURCE_VERSION_META_KEY.to_string(), Some(source_version.to_string()), @@ -1052,12 +1149,14 @@ async fn publish( eviction: Option<(Vec, Vec)>, new_fragments: Vec, keys: Option, + expected_incarnation: Option<&str>, ) -> Result { let planned = view_ds.version().version; #[cfg(test)] tests::hold_before_publish(view_ds.uri()).await; #[cfg(test)] tests::hold_until_peers_planned(); + ensure_incarnation(view_ds, expected_incarnation, view_ds.uri()).await?; let (updated_fragments, removed_fragment_ids) = eviction.unwrap_or_default(); let committed = CommitBuilder::new(WriteDestination::Dataset(Arc::new(view_ds.clone()))) .execute(Transaction::new( @@ -1830,7 +1929,9 @@ mod tests { ); let staged = eviction.finish().await.unwrap(); assert!(staged.is_some(), "four ids over a chunk of two stage twice"); - publish(&view_ds, staged, Vec::new(), None).await.unwrap(); + publish(&view_ds, staged, Vec::new(), None, None) + .await + .unwrap(); native.dataset.reload().await.unwrap(); assert_eq!(read(view.table(), "x").await, vec![5, 6]); @@ -2486,6 +2587,144 @@ mod tests { assert_eq!(read(view.table(), "twice").await, vec![14]); } + /// A refresh bound to an incarnation refuses a view dropped and recreated + /// since, even under the same name and definition; the recreated view's + /// own token is accepted, and the token survives a refresh's stamp. + #[tokio::test] + async fn test_refresh_refuses_a_recreated_view_incarnation() { + let (conn, _, view) = refreshed_doubled(vec![1]).await; + let token = view.incarnation().unwrap().to_string(); + view.refresh() + .expect_incarnation(&token) + .execute() + .await + .unwrap(); + let reopened = conn.open_materialized_view("doubled").await.unwrap(); + assert_eq!(reopened.incarnation(), Some(token.as_str())); + + conn.drop_table("doubled", &[]).await.unwrap(); + let recreated = doubled_view(&conn).await; + assert_ne!(recreated.incarnation(), Some(token.as_str())); + + let err = recreated + .refresh() + .expect_incarnation(&token) + .execute() + .await + .unwrap_err(); + assert!(err.to_string().contains("dropped and recreated"), "{err}"); + assert_eq!(read(recreated.table(), "twice").await, Vec::::new()); + + recreated + .refresh() + .expect_incarnation(recreated.incarnation().unwrap()) + .execute() + .await + .unwrap(); + assert_eq!(read(recreated.table(), "twice").await, vec![2]); + } + + /// A cloned declaration creates two physical tables; each gets its own + /// token. + #[tokio::test] + async fn test_cloned_declaration_mints_a_fresh_incarnation_per_create() { + let (conn, source) = db_with_source(vec![1]).await; + let prepared = crate::materialized_view::prepare_declaration( + &source, + &[("x".into(), "x".into()), ("twice".into(), "x * 2".into())], + None, + None, + ) + .await + .unwrap(); + let replacement = prepared.clone(); + let first = prepared.create("cloned").await.unwrap(); + let first_token = first.incarnation().unwrap().to_string(); + + conn.drop_table("cloned", &[]).await.unwrap(); + let second = replacement.create("cloned").await.unwrap(); + assert_ne!(second.incarnation(), Some(first_token.as_str())); + } + + /// A recreation that lands after planning but before publication is + /// caught by the pre-commit read: the stale refresh fails and the + /// replacement stays empty under its own token. + #[tokio::test(flavor = "multi_thread")] + async fn test_bound_refresh_cannot_publish_into_a_raced_recreation() { + let _serial = DRIFT_LOCK.lock().await; + let (conn, _) = db_with_source(vec![1]).await; + let view = doubled_view(&conn).await; + let token = view.incarnation().unwrap().to_string(); + let uri = view + .table() + .as_native() + .unwrap() + .dataset + .get() + .await + .unwrap() + .uri() + .to_string(); + + *DRIFT_TARGET.lock().unwrap() = Some(uri); + let refreshing = + tokio::spawn(async move { view.refresh().expect_incarnation(token).execute().await }); + tokio::time::timeout(std::time::Duration::from_secs(30), DRIFT_PLANNED.notified()) + .await + .expect("refresh never reached publication"); + + conn.drop_table("doubled", &[]).await.unwrap(); + let replacement = doubled_view(&conn).await; + let replacement_token = replacement.incarnation().unwrap().to_string(); + DRIFT_RELEASED.notify_one(); + + let result = refreshing.await.unwrap(); + assert!(result.is_err(), "the stale refresh unexpectedly succeeded"); + let reopened = conn.open_materialized_view("doubled").await.unwrap(); + assert_eq!(reopened.incarnation(), Some(replacement_token.as_str())); + assert_eq!(read(reopened.table(), "twice").await, Vec::::new()); + } + + /// Replacing the schema metadata wholesale drops the token. A refresh + /// bound to the old token is refused for that reason, not as a + /// recreation; an unbound refresh mints the view a fresh one. + #[tokio::test] + async fn test_a_view_whose_metadata_was_replaced_starts_a_new_incarnation() { + let (conn, _, view) = refreshed_doubled(vec![1]).await; + let token = view.incarnation().unwrap().to_string(); + let mut metadata = HashMap::new(); + metadata.insert( + crate::materialized_view::DEFINITION_META_KEY.to_string(), + crate::materialized_view::definition_to_metadata(view.definition()).unwrap(), + ); + view.table() + .as_native() + .unwrap() + .replace_schema_metadata(metadata) + .await + .unwrap(); + assert_eq!( + conn.open_materialized_view("doubled") + .await + .unwrap() + .incarnation(), + None + ); + + let err = view + .refresh() + .expect_incarnation(&token) + .execute() + .await + .unwrap_err(); + assert!(err.to_string().contains("no incarnation token"), "{err}"); + + view.refresh().execute().await.unwrap(); + let reopened = conn.open_materialized_view("doubled").await.unwrap(); + assert!(reopened.incarnation().is_some()); + assert_ne!(reopened.incarnation(), Some(token.as_str())); + } + /// In-process refreshes of one view serialize: the loser of the race /// observes the winner's watermark instead of appending the same rows. #[tokio::test(flavor = "multi_thread")] @@ -2528,7 +2767,7 @@ mod tests { let stale = view_native.dataset.get().await.unwrap().as_ref().clone(); view.table().delete("x = 1").await.unwrap(); - let err = stamp_watermark(view_native, stale, 99, 99).await; + let err = stamp_watermark(view_native, stale, 99, 99, None).await; assert!(err.is_err()); let result = view.refresh().execute().await.unwrap(); diff --git a/rust/lancedb/src/query.rs b/rust/lancedb/src/query.rs index 2d1159eff..1834efb73 100644 --- a/rust/lancedb/src/query.rs +++ b/rust/lancedb/src/query.rs @@ -2392,11 +2392,14 @@ mod tests { .postfilter(); let result = query.execute().await; let mut stream = result.expect("should have result"); - // should only have one batch + let mut num_rows = 0; while let Some(batch) = stream.next().await { - // post filter should have removed some rows - assert!(batch.expect("should be Ok").num_rows() < 10); + let batch = batch.expect("should be Ok"); + let ids: &Int32Array = batch["id"].as_primitive(); + assert!(ids.iter().all(|id| id.unwrap() % 2 == 0)); + num_rows += batch.num_rows(); } + assert!(num_rows <= 10); let query = table .query() diff --git a/rust/lancedb/src/remote/db.rs b/rust/lancedb/src/remote/db.rs index 8265d22c0..3f4216bc8 100644 --- a/rust/lancedb/src/remote/db.rs +++ b/rust/lancedb/src/remote/db.rs @@ -344,6 +344,62 @@ impl RemoteDatabase { self.table_cache.remove(&cache_key).await; Ok((request_id, resp)) } + + /// Collect the tables of a namespace in name order, for `table_names`. + /// + /// `table_names` promises name order and resumes after a table name, but the namespace + /// route's `page_token` is opaque -- it belongs to the store the listing walks, and a + /// token this client invented would resume from the wrong place. So the whole namespace is + /// walked by handing each response's token straight back, and the name semantics are + /// applied here. Constructing no token is what makes this work against a server on either + /// side of the change: it only ever repeats what the server said. + /// + /// This is the cost `table_names` already paid -- the server used to enumerate and sort the + /// namespace on every request -- and it is why `list_tables` replaces it. + async fn table_names_in_namespace( + &self, + request: &TableNamesRequest, + ) -> Result<(Vec, ServerVersion)> { + let namespace_id = + build_namespace_identifier(&request.namespace_path, &self.client.id_delimiter); + let path = format!("/v1/namespace/{}/table/list", namespace_id); + + let mut names = Vec::new(); + // Every page reports the same server, so keep the first page's version. + let mut version: Option = None; + let mut page_token: Option = None; + loop { + let mut req = self.client.get(&path); + if let Some(ref token) = page_token { + req = req.query(&[("page_token", token)]); + } + let (request_id, rsp) = self.client.send_with_retry(req, None, true).await?; + let rsp = self.client.check_response(&request_id, rsp).await?; + if version.is_none() { + version = Some(parse_server_version(&request_id, &rsp)?); + } + let response: ListTablesResponse = rsp.json().await.err_to_http(request_id)?; + names.extend(response.tables); + // An empty token is the end of the listing, not a token to send back: a server + // that reads an empty token as "start from the beginning" would hand back the + // first page again. + match response.page_token.filter(|token| !token.is_empty()) { + // A server that repeated a token would never finish; treat that as the end + // rather than looping on it. + Some(token) if Some(&token) != page_token.as_ref() => page_token = Some(token), + _ => break, + } + } + + names.sort(); + if let Some(ref start_after) = request.start_after { + names.retain(|name| name > start_after); + } + if let Some(limit) = request.limit { + names.truncate(limit as usize); + } + Ok((names, version.unwrap_or_default())) + } } #[cfg(all(test, feature = "remote"))] @@ -621,29 +677,29 @@ impl Database for RemoteDatabase { } async fn table_names(&self, request: TableNamesRequest) -> Result> { - let mut req = if !request.namespace_path.is_empty() { - let namespace_id = - build_namespace_identifier(&request.namespace_path, &self.client.id_delimiter); - self.client - .get(&format!("/v1/namespace/{}/table/list", namespace_id)) + let (tables, version) = if request.namespace_path.is_empty() { + // The flat route resumes after a table name and orders by name, which is exactly + // what `start_after` means, so the server does the paging. + let mut req = self.client.get("/v1/table/"); + if let Some(limit) = request.limit { + req = req.query(&[("limit", limit)]); + } + if let Some(ref start_after) = request.start_after { + req = req.query(&[("page_token", start_after)]); + } + let (request_id, rsp) = self.client.send_with_retry(req, None, true).await?; + let rsp = self.client.check_response(&request_id, rsp).await?; + let version = parse_server_version(&request_id, &rsp)?; + let tables = rsp + .json::() + .await + .err_to_http(request_id)? + .tables; + (tables, version) } else { - self.client.get("/v1/table/") + self.table_names_in_namespace(&request).await? }; - if let Some(limit) = request.limit { - req = req.query(&[("limit", limit)]); - } - if let Some(start_after) = request.start_after { - req = req.query(&[("page_token", start_after)]); - } - let (request_id, rsp) = self.client.send_with_retry(req, None, true).await?; - let rsp = self.client.check_response(&request_id, rsp).await?; - let version = parse_server_version(&request_id, &rsp)?; - let tables = rsp - .json::() - .await - .err_to_http(request_id)? - .tables; for table in &tables { let table_identifier = build_table_identifier(table, &request.namespace_path, &self.client.id_delimiter); @@ -1227,6 +1283,101 @@ mod tests { assert_eq!(names, vec!["table1", "table2"]); } + #[tokio::test] + async fn test_table_names_in_a_namespace_never_invents_a_page_token() { + // The namespace route's token belongs to the store, so `table_names` cannot build one + // from `start_after`. It walks the namespace on the server's own tokens and applies the + // name semantics itself, which is what keeps it working either side of the change. + let page = Arc::new(AtomicUsize::new(0)); + let conn = Connection::new_with_handler(move |request| { + assert_eq!(request.url().path(), "/v1/namespace/ns/table/list"); + let query = request.url().query().unwrap_or(""); + assert!( + !query.contains("page_token=users"), + "a table name must never be sent as a page token: {query}" + ); + match page.fetch_add(1, Ordering::SeqCst) { + 0 => { + assert!( + !query.contains("page_token"), + "the walk starts with no token" + ); + http::Response::builder() + .status(200) + .body(r#"{"tables": ["users", "orders"], "page_token": "opaque-1"}"#) + .unwrap() + } + _ => { + assert!(query.contains("page_token=opaque-1")); + http::Response::builder() + .status(200) + .body(r#"{"tables": ["widgets"]}"#) + .unwrap() + } + } + }); + + let names = conn + .table_names() + .namespace(vec!["ns".to_string()]) + .start_after("users") + .execute() + .await + .unwrap(); + // Name order, resumed after "users": "orders" sorts before it and is dropped. + assert_eq!(names, vec!["widgets"]); + } + + #[tokio::test] + async fn test_table_names_in_a_namespace_stops_on_a_repeated_token() { + // A server that handed back the token it was given would never finish the walk. + let conn = Connection::new_with_handler(|_request| { + http::Response::builder() + .status(200) + .body(r#"{"tables": ["a"], "page_token": "same"}"#) + .unwrap() + }); + + let names = conn + .table_names() + .namespace(vec!["ns".to_string()]) + .execute() + .await + .unwrap(); + // The guard bounds the walk instead of letting it run forever. The repeat is the + // server breaking the token contract and is not papered over here. + assert_eq!(names, vec!["a", "a"]); + } + + #[tokio::test] + async fn test_table_names_in_a_namespace_stops_on_an_empty_token() { + // An empty token ends the listing. Sending it back would ask a server that reads it + // as "start from the beginning" for the first page a second time, and every name on + // that page would be collected twice. + let requests = Arc::new(AtomicUsize::new(0)); + let seen = requests.clone(); + let conn = Connection::new_with_handler(move |request| { + seen.fetch_add(1, Ordering::SeqCst); + assert!( + !request.url().query().unwrap_or("").contains("page_token"), + "an empty token must never be sent back" + ); + http::Response::builder() + .status(200) + .body(r#"{"tables": ["a"], "page_token": ""}"#) + .unwrap() + }); + + let names = conn + .table_names() + .namespace(vec!["ns".to_string()]) + .execute() + .await + .unwrap(); + assert_eq!(names, vec!["a"]); + assert_eq!(requests.load(Ordering::SeqCst), 1); + } + #[tokio::test] async fn test_table_names_pagination() { let conn = Connection::new_with_handler(|request| { diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index ba8f0aaae..34b22d138 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -87,22 +87,31 @@ 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"; /// 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 @@ -119,14 +128,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 { @@ -141,6 +158,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 @@ -151,6 +223,8 @@ struct FreshnessJob { inner: RemoteJob, freshness: Arc>, version: Arc>>, + track_refresh_result: bool, + freshness_request: FreshnessHeaders, } #[async_trait] @@ -167,7 +241,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) } @@ -194,6 +292,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 { @@ -211,8 +341,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) @@ -242,12 +373,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 } @@ -277,7 +408,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?; @@ -469,6 +600,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)); } @@ -511,23 +643,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 { @@ -566,7 +725,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 { @@ -593,21 +752,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 { @@ -620,14 +792,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, @@ -1004,24 +1203,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**. @@ -1036,46 +1243,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( @@ -1083,7 +1269,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. @@ -1093,15 +1281,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); @@ -1126,8 +1316,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); @@ -1231,14 +1424,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); @@ -1258,6 +1452,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 { @@ -1272,6 +1467,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)) @@ -1402,18 +1603,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()); @@ -1432,7 +1640,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); } @@ -1458,6 +1666,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 @@ -1470,7 +1679,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) => { @@ -1526,18 +1735,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(); @@ -1600,6 +1815,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 @@ -1611,6 +1858,7 @@ impl RemoteTable { body: &str, request_id: &str, schema: &SchemaRef, + read_snapshot: ReadSnapshot, ) -> Result> { use crate::index::IndexType; @@ -1679,7 +1927,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, @@ -1726,9 +1977,39 @@ 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) } + + async fn checkout_current(&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 checkout(&self, version: u64) -> Result<()> { // Validate the version exists. The describe is sent without freshness // headers so a stale `min_version` from a previous write doesn't ride @@ -1736,7 +2017,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 @@ -1749,46 +2030,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 read_snapshot = self.snapshot_read_state().await; + let version = match read_snapshot.version { + Some(version) => 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(()) @@ -1796,7 +2083,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?; @@ -1860,31 +2148,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( @@ -1926,7 +2225,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 { @@ -1987,8 +2286,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())?; @@ -2042,7 +2343,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), @@ -2068,6 +2369,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!( @@ -2079,7 +2386,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 { @@ -2097,7 +2404,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 @@ -2105,7 +2412,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 { @@ -2113,23 +2428,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?; @@ -2141,7 +2461,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 { @@ -2153,6 +2472,15 @@ impl BaseTable for RemoteTable { async fn add(&self, mut add: AddDataBuilder) -> Result { self.check_mutable().await?; + if add.allow_external_blob_outside_bases { + return Err(Error::NotSupported { + message: "allow_external_blob_outside_bases is only supported on local tables" + .to_string(), + }); + } + // String blob values still coerce to the uri child in into_plan. + // Remote and local share that input shape. + let table_schema = self.schema().await?; crate::table::computed_columns::ensure_supported_function_metadata(table_schema.as_ref())?; let table_def = TableDefinition::try_from_rich_schema(table_schema.clone())?; @@ -2299,9 +2627,12 @@ impl BaseTable for RemoteTable { return crate::query::explain_take_offsets_plan(self, request, offsets, verbose).await; } - 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| { @@ -2315,7 +2646,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())?; @@ -2357,7 +2690,9 @@ impl BaseTable for RemoteTable { }; let query = prepared_query.as_ref().unwrap_or(query); - 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(&[( @@ -2366,14 +2701,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())?; @@ -2417,7 +2755,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())?; @@ -2436,7 +2777,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) } @@ -2452,7 +2793,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() { @@ -2468,7 +2812,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) } @@ -2478,7 +2822,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(), }) } @@ -2518,16 +2868,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, @@ -2545,7 +2902,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) => { @@ -2584,7 +2941,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, })); @@ -2658,7 +3016,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 )); @@ -2724,18 +3082,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(()) } @@ -2772,7 +3129,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())?; @@ -2789,7 +3149,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) } @@ -2825,7 +3185,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())?; @@ -2841,7 +3204,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) } @@ -2884,7 +3247,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())?; @@ -2899,7 +3265,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) } @@ -2921,7 +3287,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?; @@ -2941,6 +3308,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(), }))) } @@ -2972,7 +3341,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())?; @@ -2988,7 +3360,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) } @@ -3007,7 +3379,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())?; @@ -3019,7 +3394,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) } @@ -3031,7 +3406,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())?; @@ -3047,55 +3425,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<()> { @@ -3185,7 +3542,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 })); } @@ -3207,15 +3566,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, + ), + )) } } @@ -3306,7 +3671,10 @@ mod tests { use arrow::{array::AsArray, compute::concat_batches, datatypes::Int32Type}; use arrow_array::Array; use arrow_array::builder::LargeBinaryBuilder; - use arrow_array::{BinaryArray, Int32Array, RecordBatch, RecordBatchIterator, record_batch}; + use arrow_array::{ + BinaryArray, Int32Array, Int64Array, RecordBatch, RecordBatchIterator, StringArray, + StructArray, record_batch, + }; use arrow_schema::{DataType, Field, Schema}; use chrono::{DateTime, Utc}; use futures::{StreamExt, TryFutureExt, TryStreamExt, future::BoxFuture}; @@ -3658,6 +4026,88 @@ mod tests { assert_eq!(&body, &expected_body); } + #[tokio::test] + async fn add_rejects_external_blob_flag_before_any_request() { + let table = Table::new_with_handler::("my_table", |request| { + panic!("Unexpected request: {}", request.url().path()) + }); + let data = RecordBatch::try_new( + Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])), + vec![Arc::new(Int32Array::from(vec![1]))], + ) + .unwrap(); + + let err = table + .add(data) + .allow_external_blob_outside_bases(true) + .execute() + .await + .unwrap_err(); + + assert!(matches!(err, Error::NotSupported { .. }), "got {err:?}"); + assert!(err.to_string().contains("local tables")); + } + + #[tokio::test] + async fn add_string_blob_becomes_uri_struct_without_the_local_flag() { + let table_schema = Schema::new(vec![ + Field::new("id", DataType::Int64, false), + crate::blob("image", true), + ]); + let describe_body = describe_response(&table_schema); + let input = RecordBatch::try_new( + Arc::new(Schema::new(vec![ + Field::new("id", DataType::Int64, false), + Field::new("image", DataType::Utf8, true), + ])), + vec![ + Arc::new(Int64Array::from(vec![1])), + Arc::new(StringArray::from(vec![Some("s3://bucket/key")])), + ], + ) + .unwrap(); + + let (sender, receiver) = std::sync::mpsc::channel(); + let table = + Table::new_with_handler("my_table", move |mut request| match request.url().path() { + "/v1/table/my_table/describe/" => http::Response::builder() + .status(200) + .body(describe_body.clone()) + .unwrap(), + "/v1/table/my_table/insert/" => { + let mut body_out = reqwest::Body::from(Vec::new()); + std::mem::swap(request.body_mut().as_mut().unwrap(), &mut body_out); + sender.send(body_out).unwrap(); + http::Response::builder() + .status(200) + .body(r#"{"version": 2}"#.to_string()) + .unwrap() + } + path => panic!("Unexpected path: {path}"), + }); + + table.add(input).execute().await.unwrap(); + + let body = collect_body(receiver.recv().unwrap()).await; + let mut reader = + arrow_ipc::reader::StreamReader::try_new(std::io::Cursor::new(body), None).unwrap(); + let batch = reader.next().unwrap().unwrap(); + let image = batch + .column_by_name("image") + .unwrap() + .as_any() + .downcast_ref::() + .expect("remote add should send the coerced blob struct"); + let uri: &StringArray = image + .column_by_name("uri") + .unwrap() + .as_any() + .downcast_ref() + .unwrap(); + assert_eq!(uri.value(0), "s3://bucket/key"); + assert!(image.column_by_name("data").unwrap().is_null(0)); + } + #[rstest] #[case(true)] #[case(false)] @@ -4386,6 +4836,43 @@ mod tests { assert!(!table.base_table().scan_order_is_deterministic()); } + #[tokio::test] + async fn test_checkout_branch_pins_without_touching_the_original() { + let seen = Arc::new(std::sync::Mutex::new(Vec::new())); + let recorder = seen.clone(); + let table = Table::new_with_handler_version( + "my_table", + semver::Version::new(0, 5, 0), + move |request| match request.url().path() { + "/v1/table/my_table/describe/" => http::Response::builder() + .status(200) + .body(br#"{"version": 42, "schema": {"fields": []}}"#.to_vec()) + .unwrap(), + "/v1/table/my_table/count_rows/" => { + let body = request_body_json(&request); + recorder.lock().unwrap().push(body["version"].clone()); + http::Response::builder() + .status(200) + .body(b"0".to_vec()) + .unwrap() + } + path => panic!("unexpected request path: {path}"), + }, + ); + + let pinned = table.checkout_branch("main", Some(42)).await.unwrap(); + pinned.count_rows(None).await.unwrap(); + table.count_rows(None).await.unwrap(); + + let seen = seen.lock().unwrap(); + assert_eq!(seen[0], 42, "the pinned handle must send its version"); + assert!( + seen[1].is_null(), + "the original handle must still track latest, got {:?}", + seen[1] + ); + } + #[tokio::test] async fn test_fetch_blobs_sends_the_checked_out_version() { let ipc = one_row_blob_ipc_stream("image"); @@ -6895,8 +7382,7 @@ mod tests { r#"{ "function":{"name":"embed","version":"fv_01K3EXACT"}, "inputs":[{"parameter":"text","kind":"column","value":{"path":"description"}}], - "output":{"kind":"scalar","arrow_type":"list","nullable":false}, - "group_id":"fg_scalar" + "output":{"kind":"scalar","arrow_type":"list","nullable":false} }"#, ) .unwrap(); @@ -6911,7 +7397,53 @@ mod tests { } #[tokio::test] - async fn test_add_named_struct_function_expands_one_atomic_sibling_group() { + async fn test_add_fixed_size_list_function_column_declares_the_vector_type() { + let table = Table::new_with_handler("my_table", |request| { + match request.url().path() { + "/v1/table/my_table/describe/" => http::Response::builder() + .status(200) + .body( + r#"{"version":1,"schema":{"fields":[{"name":"description","nullable":true,"type":{"type":"string"}}]}}"#, + ) + .unwrap(), + "/v1/table/my_table/add_columns/" => { + let actual: serde_json::Value = serde_json::from_slice( + request.body().unwrap().as_bytes().unwrap(), + ) + .unwrap(); + let expected: serde_json::Value = serde_json::from_str(include_str!( + "../../tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json" + )) + .unwrap(); + assert_eq!(actual, expected); + http::Response::builder() + .status(200) + .body(r#"{"version":8}"#) + .unwrap() + } + path => panic!("Unexpected path: {path}"), + } + }); + let application = crate::function::FunctionApplication::from_json( + r#"{ + "function":{"name":"embed","version":"fv_01K3EXACT"}, + "inputs":[{"parameter":"text","kind":"column","value":{"path":"description"}}], + "output":{"kind":"scalar","arrow_type":"fixed_size_list","nullable":false} + }"#, + ) + .unwrap(); + + let result = table + .add_columns() + .function_as("embedding", application) + .execute() + .await + .unwrap(); + assert_eq!(result.version, 8); + } + + #[tokio::test] + async fn test_add_named_struct_function_expands_one_atomic_binding() { let table = Table::new_with_handler("my_table", |request| match request.url().path() { "/v1/table/my_table/describe/" => http::Response::builder() .status(200) @@ -6926,7 +7458,7 @@ mod tests { let actual: serde_json::Value = serde_json::from_slice(request.body().unwrap().as_bytes().unwrap()).unwrap(); let expected: serde_json::Value = serde_json::from_str(include_str!( - "../../tests/fixtures/first_class_functions/v1/remote_grouped_declaration_request.json" + "../../tests/fixtures/first_class_functions/v1/remote_multi_output_declaration_request.json" )) .unwrap(); assert_eq!(actual, expected); @@ -6948,7 +7480,6 @@ mod tests { {"name":"normalized_text","arrow_type":"utf8","nullable":false}, {"name":"token_count","arrow_type":"int64","nullable":false} ]}, - "group_id":"fg_01K3TEXT", "columns":{"normalized_text":"search_text"} }"#, ) @@ -7053,12 +7584,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() @@ -7071,7 +7602,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() @@ -7088,8 +7623,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" ); } @@ -7264,12 +7799,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::<()>(); @@ -7293,10 +7827,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()) @@ -7317,25 +7848,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") ); } @@ -8709,6 +9233,28 @@ mod tests { } } + /// A pinned snapshot should reuse the version and schema returned by its + /// initial describe instead of issuing two more describe requests. + #[tokio::test] + async fn test_checkout_current_seeds_schema_from_single_describe() { + let describe_calls = Arc::new(AtomicUsize::new(0)); + let calls = describe_calls.clone(); + let table = Table::new_with_handler("my_table", move |request| { + assert_eq!(request.url().path(), "/v1/table/my_table/describe/"); + calls.fetch_add(1, Ordering::SeqCst); + http::Response::builder() + .status(200) + .body( + r#"{"version":42,"schema":{"fields":[{"name":"a","type":{"type":"int32"},"nullable":false}]}}"#, + ) + .unwrap() + }); + + let snapshot = table.checkout_current().await.unwrap(); + assert_eq!(snapshot.schema().await.unwrap().fields().len(), 1); + assert_eq!(describe_calls.load(Ordering::SeqCst), 1); + } + /// Test that schema cache is invalidated after checkout #[tokio::test] async fn test_schema_cache_invalidation_on_checkout() { @@ -8771,6 +9317,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() { @@ -10005,6 +10595,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)); @@ -10030,6 +10621,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), @@ -10042,6 +10634,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), @@ -10118,6 +10711,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()); @@ -10186,6 +10883,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". @@ -10232,6 +10985,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. @@ -10647,6 +11558,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| { @@ -10738,6 +11729,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 323790596..f0c08f0ba 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -560,6 +560,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. @@ -759,7 +766,7 @@ pub trait BaseTable: std::fmt::Display + std::fmt::Debug + Send + Sync { message: "computed columns are not supported on this table type".into(), }) } - /// Declare one immutable registered-Function output group. + /// Declare one immutable registered-Function binding. async fn add_function_columns( &self, _application: &crate::function::FunctionApplication, @@ -793,6 +800,12 @@ pub trait BaseTable: std::fmt::Display + std::fmt::Debug + Send + Sync { async fn drop_columns(&self, columns: &[&str]) -> Result; /// Get the version of the table. async fn version(&self) -> Result; + /// Return a new table handle pinned to the exact revision currently visible. + async fn checkout_current(&self) -> Result> { + Err(Error::NotSupported { + message: "checkout_current is not supported on this table type".into(), + }) + } /// Checkout a specific version of the table. async fn checkout(&self, version: u64) -> Result<()>; /// Checkout a table version referenced by a tag. @@ -800,6 +813,14 @@ pub trait BaseTable: std::fmt::Display + std::fmt::Debug + Send + Sync { async fn checkout_tag(&self, tag: &str) -> Result<()>; /// Checkout the latest version of the table. async fn checkout_latest(&self) -> Result<()>; + /// Return an independent handle pinned to the version currently selected. + /// + /// Backends that can advance between requests should override this for + /// multi-request operations that need snapshot consistency. Backends whose + /// existing handles already provide the desired behavior return `None`. + async fn snapshot_at_current_version(&self) -> Result>> { + Ok(None) + } /// Whether repeated identical scans return rows in the same order. /// /// Callers that assign meaning to a row's position must order the results @@ -1133,6 +1154,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 @@ -1944,6 +1975,20 @@ impl Table { self.inner.version().await } + /// Return a new table handle pinned to the exact revision currently visible. + /// + /// This is used when asynchronous preparation must remain consistent with + /// the revision used for a later read. + #[doc(hidden)] + pub async fn checkout_current(&self) -> Result { + let inner = self.inner.checkout_current().await?; + Ok(Self { + inner, + database: self.database.clone(), + embedding_registry: self.embedding_registry.clone(), + }) + } + /// Checks out a specific version of the Table /// /// Any read operation on the table will now access the data at the checked out version. @@ -3039,10 +3084,33 @@ 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) } + async fn checkout_current(&self) -> Result> { + let current = self.dataset.get().await?; + let dataset = dataset::DatasetConsistencyWrapper::new_time_travel( + current.as_ref().clone(), + self.read_consistency_interval, + ); + Ok(Arc::new(Self { + dataset, + ..self.clone() + })) + } + async fn checkout(&self, version: u64) -> Result<()> { self.dataset.as_time_travel(version).await } @@ -3208,7 +3276,7 @@ impl BaseTable for NativeTable { let output = add.into_plan(&table_schema, &table_def)?; - let lance_params = output + let mut lance_params = output .write_options .lance_write_params .unwrap_or(WriteParams { @@ -3218,6 +3286,9 @@ impl BaseTable for NativeTable { }, ..Default::default() }); + if output.allow_external_blob_outside_bases { + lance_params.allow_external_blob_outside_bases = true; + } // Repartition for write parallelism if beneficial. let plan = if num_partitions > 1 { @@ -4083,6 +4154,14 @@ mod tests { parent_list_calls: self.parent_list_calls.clone(), }) } + + fn wrap_paginated( + &self, + _store_prefix: &str, + _original: Arc, + ) -> Option> { + None + } } #[tokio::test] @@ -4186,6 +4265,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/add_columns.rs b/rust/lancedb/src/table/add_columns.rs index 3b91c30e4..1ac0c6b4f 100644 --- a/rust/lancedb/src/table/add_columns.rs +++ b/rust/lancedb/src/table/add_columns.rs @@ -88,7 +88,7 @@ impl AddColumnsBuilder { } /// Declare every field of a named-struct Function result as one atomic - /// sibling group. Result-field aliases come from + /// binding. Result-field aliases come from /// [`FunctionApplication::columns`](crate::function::FunctionApplication::columns). /// /// ``` diff --git a/rust/lancedb/src/table/add_data.rs b/rust/lancedb/src/table/add_data.rs index 11ba43dd6..15ccf9f66 100644 --- a/rust/lancedb/src/table/add_data.rs +++ b/rust/lancedb/src/table/add_data.rs @@ -60,6 +60,7 @@ pub struct AddDataBuilder { pub(crate) embedding_registry: Option>, pub(crate) progress_callback: Option, pub(crate) write_parallelism: Option, + pub(crate) allow_external_blob_outside_bases: bool, } impl std::fmt::Debug for AddDataBuilder { @@ -87,6 +88,7 @@ impl AddDataBuilder { embedding_registry, progress_callback: None, write_parallelism: None, + allow_external_blob_outside_bases: false, } } @@ -141,6 +143,16 @@ impl AddDataBuilder { self } + /// Store blob URIs that sit outside registered blob bases. + /// + /// The row keeps a reference, so the object has to stay readable. + /// [`crate::table::Table::fetch_blobs`] reads from that location. + /// Defaults to `false`. Local tables only. + pub fn allow_external_blob_outside_bases(mut self, allow: bool) -> Self { + self.allow_external_blob_outside_bases = allow; + self + } + pub async fn execute(self) -> Result { if self.write_parallelism.map(|p| p == 0).unwrap_or(false) { return Err(Error::InvalidInput { @@ -199,6 +211,7 @@ impl AddDataBuilder { write_options: self.write_options, mode: self.mode, tracker, + allow_external_blob_outside_bases: self.allow_external_blob_outside_bases, }) } } @@ -212,6 +225,7 @@ pub struct PreprocessingOutput { pub write_options: WriteOptions, pub mode: AddDataMode, pub tracker: Option>, + pub allow_external_blob_outside_bases: bool, } /// Check that the input schema is valid for insert. diff --git a/rust/lancedb/src/table/computed_columns.rs b/rust/lancedb/src/table/computed_columns.rs index b1c0b8370..2b95cb34c 100644 --- a/rust/lancedb/src/table/computed_columns.rs +++ b/rust/lancedb/src/table/computed_columns.rs @@ -13,7 +13,7 @@ //! self-describing -- both are derived from the expression, so a caller writes //! neither -- while a kind resolved through a registry cannot be typed without //! consulting it. Registered Functions use an exact remote version plus a -//! schema-level grouped binding; unknown newer kinds remain readable and fail +//! schema-level Function binding; unknown newer kinds remain readable and fail //! closed before mutation. //! //! [`computed_columns`] and [`computed_column_from_field`] read declarations @@ -46,16 +46,16 @@ pub const EXPRESSION_META_KEY: &str = "computed_column.expression"; /// Field metadata key holding the column's inputs, as a JSON array of names. pub const INPUTS_META_KEY: &str = "computed_column.inputs"; -/// Field metadata key holding the grouped Function binding identity. +/// Field metadata key holding the Function binding identity. pub const FUNCTION_BINDING_ID_META_KEY: &str = "computed_column.function.binding_id"; /// Field metadata key holding this sibling's ordered Function output ordinal. pub const FUNCTION_OUTPUT_ORDINAL_META_KEY: &str = "computed_column.function.output_ordinal"; -/// Schema metadata key holding all immutable grouped Function bindings. +/// Schema metadata key holding all immutable Function bindings. pub const FUNCTION_BINDINGS_META_KEY: &str = "lancedb::function_bindings"; -/// Version of the schema-level grouped Function binding envelope. +/// Version of the schema-level Function binding envelope. pub const FUNCTION_BINDINGS_VERSION: u32 = 1; /// Value of [`KIND_META_KEY`] for a column defined by a SQL expression. @@ -81,7 +81,7 @@ pub enum ComputedColumnKind { /// The expression. expression: String, }, - /// One physical output in an immutable grouped registered-Function + /// One physical output in an immutable registered-Function /// binding. The full binding lives in schema metadata. Function { /// Shared immutable binding identity. @@ -159,7 +159,7 @@ struct FunctionBindingEnvelope { bindings: Vec, } -/// Encode immutable grouped bindings for schema-level persistence. +/// Encode immutable Function bindings for schema-level persistence. pub fn function_bindings_metadata(bindings: &[FunctionBinding]) -> Result { let bindings = bindings .iter() @@ -177,7 +177,7 @@ pub fn function_bindings_metadata(bindings: &[FunctionBinding]) -> Result Result> { let Some(envelope) = function_binding_envelope(schema)? else { @@ -238,21 +238,15 @@ pub(crate) fn ensure_supported_function_metadata(schema: &ArrowSchema) -> Result message: format!("duplicate Function binding '{}'", binding.binding_id()), }); } - if binding.revision() == 0 || binding.outputs().is_empty() { + if binding.outputs().is_empty() { return Err(Error::InvalidInput { - message: format!( - "Function binding '{}' has no immutable revision or outputs", - binding.binding_id() - ), + message: format!("Function binding '{}' has no outputs", binding.binding_id()), }); } - if binding.function().name.is_empty() - || binding.function().version.is_empty() - || binding.group_id().is_empty() - { + if binding.function().name.is_empty() || binding.function().version.is_empty() { return Err(Error::InvalidInput { message: format!( - "Function binding '{}' has no exact version or group identity", + "Function binding '{}' has no exact version", binding.binding_id() ), }); @@ -493,9 +487,7 @@ fn ensure_known_binding_shape(value: &Value) -> Result<()> { value, &[ "binding_id", - "revision", "function", - "group_id", "inputs", "outputs", "input_schema", @@ -586,6 +578,25 @@ fn canonical_input_arrow_type(field: &JsonArrowField) -> Result { } } +/// `fixed_size_list` -> (`item`, `size`); the comma must sit outside +/// any nested `<...>`. +fn split_fixed_size_list(raw: &str) -> Option<(&str, i32)> { + let inner = raw.strip_prefix("fixed_size_list<")?.strip_suffix('>')?; + let mut depth = 0_u32; + let mut separator = None; + for (index, byte) in inner.bytes().enumerate() { + match byte { + b'<' => depth += 1, + b'>' => depth = depth.checked_sub(1)?, + b',' if depth == 0 => separator = Some(index), + _ => {} + } + } + let (item, size) = inner.split_at(separator?); + let size: i32 = size[1..].trim().parse().ok()?; + (size > 0).then_some((item.trim(), size)) +} + fn parse_output_arrow_type(raw: &str) -> Result { fn parse(raw: &str) -> Result { let raw = raw.trim(); @@ -618,6 +629,16 @@ fn parse_output_arrow_type(raw: &str) -> Result { )]); return Ok(data_type); } + if let Some((inner, size)) = split_fixed_size_list(raw) { + let mut data_type = JsonArrowDataType::new("fixed_size_list".to_string()); + data_type.fields = Some(vec![JsonArrowField::new( + "item".to_string(), + false, + parse(inner)?, + )]); + data_type.length = Some(i64::from(size)); + return Ok(data_type); + } let normalized = match raw { "boolean" => "bool", "string" => "utf8", @@ -757,12 +778,9 @@ pub(crate) fn plan_function_application( message: "Function application contains fields from a newer contract".into(), }); } - if application.function().name.is_empty() - || application.function().version.is_empty() - || application.group_id().is_empty() - { + if application.function().name.is_empty() || application.function().version.is_empty() { return Err(invalid_function( - "Function application requires an exact version and group identity", + "Function application requires an exact version", )); } @@ -1322,6 +1340,32 @@ pub(super) async fn add_foreign_kind(table: &crate::Table, name: &str, kind: &st #[cfg(test)] mod tests { + #[test] + fn output_arrow_type_grammar_matches_the_shared_golden() { + let golden: serde_json::Value = serde_json::from_str(include_str!( + "../../tests/fixtures/first_class_functions/v1/arrow_types.json" + )) + .unwrap(); + let valid = golden["valid"].as_array().unwrap().iter(); + for case in valid.chain(golden["server_only"].as_array().unwrap()) { + let raw = case["arrow_type"].as_str().unwrap(); + let parsed = super::parse_output_arrow_type(raw) + .unwrap_or_else(|error| panic!("{raw}: {error}")); + assert_eq!( + serde_json::to_value(&parsed).unwrap(), + case["json"], + "{raw}" + ); + } + for raw in golden["invalid"].as_array().unwrap() { + let raw = raw.as_str().unwrap(); + assert!( + super::parse_output_arrow_type(raw).is_err(), + "{raw:?} should be rejected" + ); + } + } + use arrow_array::record_batch; use arrow_schema::DataType; use futures::TryStreamExt; @@ -2206,7 +2250,6 @@ mod tests { {{"name":"normalized_text","arrow_type":"utf8","nullable":false}}, {{"name":"token_count","arrow_type":"int64","nullable":false}} ]}}, - "group_id":"fg_exact", "columns":{columns} }}"# )) @@ -2387,8 +2430,7 @@ mod tests { r#"{ "function":{"name":"f","version":"fv"}, "inputs":[{"parameter":"title","kind":"future_source","value":{"path":"title"}}], - "output":{"kind":"scalar","arrow_type":"int64","nullable":false}, - "group_id":"fg" + "output":{"kind":"scalar","arrow_type":"int64","nullable":false} }"#, ) .unwrap(); @@ -2401,7 +2443,6 @@ mod tests { "function":{"name":"f","version":"fv"}, "inputs":[], "output":{"kind":"scalar","arrow_type":"int64","nullable":false}, - "group_id":"fg", "future_declaration":{"mode":"managed"} }"#, ) @@ -2415,8 +2456,7 @@ mod tests { r#"{ "function":{"name":"f","version":"fv"}, "inputs":[], - "output":{"kind":"scalar","arrow_type":"int64","nullable":false,"assignment":"cell_flag"}, - "group_id":"fg" + "output":{"kind":"scalar","arrow_type":"int64","nullable":false,"assignment":"cell_flag"} }"#, ) .unwrap(); diff --git a/rust/lancedb/src/table/datafusion/blob_coerce.rs b/rust/lancedb/src/table/datafusion/blob_coerce.rs index b29b2423b..cb984f7f4 100644 --- a/rust/lancedb/src/table/datafusion/blob_coerce.rs +++ b/rust/lancedb/src/table/datafusion/blob_coerce.rs @@ -7,7 +7,7 @@ use std::sync::Arc; -use arrow_schema::{DataType, Field, FieldRef}; +use arrow_schema::{DataType, Field, FieldRef, Fields}; use datafusion::functions::core::{get_field, named_struct}; use datafusion_common::ScalarValue; use datafusion_common::config::ConfigOptions; @@ -35,8 +35,9 @@ pub(super) fn coerce_blob_expr( }); }; - let input_struct_children = match input_field.data_type() { - DataType::Binary | DataType::LargeBinary | DataType::BinaryView => None, + let input_shape = match input_field.data_type() { + DataType::Binary | DataType::LargeBinary | DataType::BinaryView => BlobInputShape::Bytes, + DataType::Utf8 | DataType::LargeUtf8 | DataType::Utf8View => BlobInputShape::String, DataType::Struct(children) => { if !children .iter() @@ -49,13 +50,15 @@ pub(super) fn coerce_blob_expr( ), }); } - Some(children) + BlobInputShape::Struct(children) } other => { return Err(Error::InvalidInput { message: format!( "cannot coerce column '{}' with type {} into a blob v2 struct. \ - expected Binary, LargeBinary, BinaryView, or a Struct with a 'data' or 'uri' child", + expected binary bytes (Binary, LargeBinary, BinaryView), \ + strings (Utf8, LargeUtf8, Utf8View), \ + or a Struct with a 'data' or 'uri' child", table_field.name(), other, ), @@ -69,9 +72,8 @@ pub(super) fn coerce_blob_expr( declared.name().as_str(), )))); - let value: Arc = match input_struct_children { - // Raw binary lands in `data` and everything else is a typed null. - None => { + let value: Arc = match &input_shape { + BlobInputShape::Bytes => { if declared.name() == "data" { Arc::new(CastExpr::new( input_expr.clone(), @@ -82,30 +84,43 @@ pub(super) fn coerce_blob_expr( typed_null(declared.data_type())? } } - Some(children) => match children.iter().find(|c| c.name() == declared.name()) { - Some(child) => { - let field_expr: Arc = Arc::new(ScalarFunctionExpr::new( - &format!("get_field({})", declared.name()), - get_field(), - vec![ - input_expr.clone(), - Arc::new(Literal::new(ScalarValue::from(declared.name().as_str()))), - ], - Arc::new(child.as_ref().clone()), - config.clone(), - )); - if child.data_type() == declared.data_type() { - field_expr - } else { - Arc::new(CastExpr::new( - field_expr, - declared.data_type().clone(), - None, - )) - } + BlobInputShape::String => { + if declared.name() == "uri" { + Arc::new(CastExpr::new( + input_expr.clone(), + declared.data_type().clone(), + None, + )) + } else { + typed_null(declared.data_type())? } - None => typed_null(declared.data_type())?, - }, + } + BlobInputShape::Struct(children) => { + match children.iter().find(|c| c.name() == declared.name()) { + Some(child) => { + let field_expr: Arc = Arc::new(ScalarFunctionExpr::new( + &format!("get_field({})", declared.name()), + get_field(), + vec![ + input_expr.clone(), + Arc::new(Literal::new(ScalarValue::from(declared.name().as_str()))), + ], + Arc::new(child.as_ref().clone()), + config.clone(), + )); + if child.data_type() == declared.data_type() { + field_expr + } else { + Arc::new(CastExpr::new( + field_expr, + declared.data_type().clone(), + None, + )) + } + } + None => typed_null(declared.data_type())?, + } + } }; ns_args.push(value); } @@ -120,6 +135,12 @@ pub(super) fn coerce_blob_expr( Ok((expr, table_field.clone())) } +enum BlobInputShape<'a> { + Bytes, + String, + Struct(&'a Fields), +} + fn typed_null(data_type: &DataType) -> Result> { let scalar = ScalarValue::try_from(data_type).map_err(|e| Error::InvalidInput { message: format!("cannot build null literal for blob child type {data_type}: {e}"), @@ -134,7 +155,7 @@ mod tests { use crate::blob::blob; use arrow_array::{ Array, ArrayRef, BinaryArray, BinaryViewArray, Int32Array, Int64Array, LargeBinaryArray, - RecordBatch, StringArray, StructArray, UInt8Array, UInt64Array, + RecordBatch, StringArray, StringViewArray, StructArray, UInt8Array, UInt64Array, }; use arrow_schema::Schema; use datafusion::prelude::SessionContext; @@ -436,14 +457,78 @@ mod tests { #[tokio::test] async fn unsupported_input_type_is_rejected_with_column_name() { let batch = batch_with_image( - Field::new("image", DataType::Utf8, true), - Arc::new(StringArray::from(vec!["not bytes"])), + Field::new("image", DataType::Int64, true), + Arc::new(Int64Array::from(vec![42])), ); let err = coerce_err(batch, &blob_table_schema()).await; assert!(matches!(err, Error::InvalidInput { .. }), "got {err:?}"); assert!(err.to_string().contains("image")); } + #[tokio::test] + async fn utf8_string_coerces_to_uri_child() { + let batch = batch_with_image( + Field::new("image", DataType::Utf8, true), + Arc::new(StringArray::from(vec![Some("s3://bucket/key"), None])), + ); + let coerced = coerce(batch, &blob_table_schema()).await; + let image = image_struct(&coerced); + let uri: &StringArray = image + .column_by_name("uri") + .unwrap() + .as_any() + .downcast_ref() + .unwrap(); + assert_eq!(uri.value(0), "s3://bucket/key"); + assert!(image.column_by_name("data").unwrap().is_null(0)); + assert!(uri.is_null(1)); + } + + #[tokio::test] + async fn large_utf8_string_coerces_into_four_child_blob_layout() { + use arrow_array::LargeStringArray; + + let table_schema = Schema::new(vec![ + Field::new("id", DataType::Int64, false), + wide_blob_field("image"), + ]); + let batch = batch_with_image( + Field::new("image", DataType::LargeUtf8, true), + Arc::new(LargeStringArray::from(vec!["file:///tmp/blob.bin"])), + ); + let coerced = coerce(batch, &table_schema).await; + let image = image_struct(&coerced); + assert_eq!(image.num_columns(), 4); + let uri: &StringArray = image + .column_by_name("uri") + .unwrap() + .as_any() + .downcast_ref() + .unwrap(); + assert_eq!(uri.value(0), "file:///tmp/blob.bin"); + assert!(image.column_by_name("data").unwrap().is_null(0)); + assert!(image.column_by_name("position").unwrap().is_null(0)); + assert!(image.column_by_name("size").unwrap().is_null(0)); + } + + #[tokio::test] + async fn utf8_view_string_coerces_to_uri_child() { + let batch = batch_with_image( + Field::new("image", DataType::Utf8View, true), + Arc::new(StringViewArray::from(vec![Some("s3://bucket/view-key")])), + ); + let coerced = coerce(batch, &blob_table_schema()).await; + let image = image_struct(&coerced); + let uri: &StringArray = image + .column_by_name("uri") + .unwrap() + .as_any() + .downcast_ref() + .unwrap(); + assert_eq!(uri.value(0), "s3://bucket/view-key"); + assert!(image.column_by_name("data").unwrap().is_null(0)); + } + #[tokio::test] async fn blob_metadata_survives_cast_of_sibling_column() { let batch = RecordBatch::try_new( 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 b9dd5732b..3227e3edf 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 3151f8b56..2f3102095 100644 --- a/rust/lancedb/src/table/query.rs +++ b/rust/lancedb/src/table/query.rs @@ -70,6 +70,9 @@ async fn can_execute_namespace_query(table: &NativeTable, query: &AnyQuery) -> R .contains(&NamespaceClientPushdownOperation::QueryTable) && table.namespace_client.is_some() && table.dataset.current_branch().is_none() + // NsQueryTableRequest has no version field, so a pushed-down query would + // read latest and ignore the pin. + && table.dataset.time_travel_version().is_none() && !requires_local_namespace_execution(query)) { return Ok(false); @@ -701,6 +704,7 @@ mod tests { use super::*; use crate::query::{QueryExecutionOptions, QueryRequest}; + use crate::table::BaseTable; fn fixed_size_list_array(values: Vec, dimension: i32) -> FixedSizeListArray { FixedSizeListArray::try_new_from_values(Float32Array::from(values), dimension).unwrap() @@ -893,10 +897,56 @@ mod tests { async fn query_table(&self, _request: NsQueryTableRequest) -> lance::Result { self.query_table_calls.fetch_add(1, Ordering::SeqCst); - panic!("approx_mode queries must not be pushed down to namespace query_table"); + panic!("query must not be pushed down to namespace query_table"); } } + #[tokio::test] + async fn test_execute_query_pinned_snapshot_with_namespace_pushdown_runs_locally() { + use crate::connect; + 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, 4, 5]))], + ) + .unwrap(); + let table = conn + .create_table("test_pinned_namespace_fallback", vec![batch]) + .execute() + .await + .unwrap(); + + let namespace_client = Arc::new(CountingNamespaceClient::default()); + let mut native_table = table.as_native().unwrap().clone(); + native_table.namespace_client = Some(namespace_client.clone()); + native_table + .pushdown_operations + .insert(NamespaceClientPushdownOperation::QueryTable); + + let snapshot = native_table.checkout_current().await.unwrap(); + let snapshot = snapshot.as_any().downcast_ref::().unwrap(); + assert!(snapshot.dataset.time_travel_version().is_some()); + + let query = AnyQuery::Query(QueryRequest { + filter: Some(QueryFilter::Sql("id > 3".to_string())), + ..Default::default() + }); + let stream = execute_query(snapshot, &query, QueryExecutionOptions::default()) + .await + .unwrap(); + let batches = stream.try_collect::>().await.unwrap(); + + assert_eq!( + batches.iter().map(|batch| batch.num_rows()).sum::(), + 2 + ); + assert_eq!(namespace_client.query_table_calls.load(Ordering::SeqCst), 0); + } + #[tokio::test] async fn test_execute_query_approx_mode_with_namespace_pushdown_runs_locally() { use crate::connect; @@ -1013,6 +1063,37 @@ mod tests { assert_eq!(namespace_client.query_table_calls.load(Ordering::SeqCst), 0); } + #[tokio::test] + 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_multivector_structure() { use arrow_array::{Float32Array, RecordBatch}; diff --git a/rust/lancedb/src/table/query/lsm.rs b/rust/lancedb/src/table/query/lsm.rs index 255d649b1..5c340cc35 100644 --- a/rust/lancedb/src/table/query/lsm.rs +++ b/rust/lancedb/src/table/query/lsm.rs @@ -298,7 +298,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/tests/blob_integration.rs b/rust/lancedb/tests/blob_integration.rs index b92f961f4..7b709b645 100644 --- a/rust/lancedb/tests/blob_integration.rs +++ b/rust/lancedb/tests/blob_integration.rs @@ -5,12 +5,14 @@ use std::sync::Arc; use arrow_array::{ Array, ArrayRef, BinaryArray, Int64Array, LargeBinaryArray, RecordBatch, StringArray, - StructArray, UInt64Array, + StructArray, UInt64Array, new_null_array, }; use arrow_schema::{DataType, Field, Fields, Schema}; use futures::TryStreamExt; use lance::Dataset; +use lance::dataset::WriteParams; use lance_file::version::{ConcreteFileVersion, LanceFileVersion}; +use lance_table::format::BasePath; use lancedb::{ Connection, Error, Result, Table, blob::{BlobRangeRequest, blob}, @@ -19,7 +21,7 @@ use lancedb::{ ListingDatabaseOptions, NewTableConfig, OPT_NEW_TABLE_ENABLE_STABLE_ROW_IDS, }, query::{ExecutableQuery, QueryBase}, - table::{AddDataMode, CompactionOptions, OptimizeAction, OptimizeStats}, + table::{AddDataMode, CompactionOptions, OptimizeAction, OptimizeStats, WriteOptions}, }; use tempfile::tempdir; @@ -261,11 +263,11 @@ async fn add_rejects_uncoercible_blob_input() -> Result<()> { let batch = RecordBatch::try_new( Arc::new(Schema::new(vec![ Field::new("id", DataType::Int64, false), - Field::new("image", DataType::Utf8, true), + Field::new("image", DataType::Int64, true), ])), vec![ Arc::new(Int64Array::from(vec![1])), - Arc::new(StringArray::from(vec!["not bytes"])), + Arc::new(Int64Array::from(vec![42])), ], ) .unwrap(); @@ -1332,3 +1334,223 @@ async fn optimize_preserves_blob_v2_null_and_empty_distinction() -> Result<()> { ); Ok(()) } + +fn uri_struct_batch(id: i64, uri: &str) -> RecordBatch { + let image_field = blob("image", true); + let DataType::Struct(child_fields) = image_field.data_type().clone() else { + unreachable!("blob field is a struct"); + }; + let children: Vec = child_fields + .iter() + .map(|field| match field.name().as_str() { + "uri" => Arc::new(StringArray::from(vec![Some(uri)])) as ArrayRef, + _ => new_null_array(field.data_type(), 1), + }) + .collect(); + let image = StructArray::new(child_fields, children, None); + RecordBatch::try_new( + Arc::new(Schema::new(vec![ + Field::new("id", DataType::Int64, false), + image_field, + ])), + vec![Arc::new(Int64Array::from(vec![id])), Arc::new(image)], + ) + .unwrap() +} + +fn uri_string_batch(id: i64, uri: &str) -> RecordBatch { + RecordBatch::try_new( + Arc::new(Schema::new(vec![ + Field::new("id", DataType::Int64, false), + Field::new("image", DataType::Utf8, true), + ])), + vec![ + Arc::new(Int64Array::from(vec![id])), + Arc::new(StringArray::from(vec![Some(uri)])), + ], + ) + .unwrap() +} + +fn write_payload_file_uri(dir: &std::path::Path, name: &str, payload: &[u8]) -> String { + let path = dir.join(name); + std::fs::write(&path, payload).unwrap(); + url::Url::from_file_path(&path).unwrap().to_string() +} + +#[tokio::test] +async fn external_uri_struct_round_trips_with_flag() -> Result<()> { + let tmp = tempdir().unwrap(); + let db = connect(tmp.path().join("db").to_str().unwrap()) + .execute() + .await?; + let payload: &[u8] = b"external-struct-payload"; + let uri = write_payload_file_uri(tmp.path(), "payload.bin", payload); + let table = db + .create_empty_table("t", blob_table_schema()) + .execute() + .await?; + + table + .add(uri_struct_batch(1, &uri)) + .allow_external_blob_outside_bases(true) + .execute() + .await?; + + let ids = collect_row_ids(&table).await?; + let bytes = table.fetch_blobs("image", &ids).await?; + assert_eq!(bytes.value(0), payload); + Ok(()) +} + +#[tokio::test] +async fn external_uri_add_requires_opt_in() -> Result<()> { + let tmp = tempdir().unwrap(); + let db = connect(tmp.path().join("db").to_str().unwrap()) + .execute() + .await?; + let uri = write_payload_file_uri(tmp.path(), "payload.bin", b"unreachable"); + let table = db + .create_empty_table("t", blob_table_schema()) + .execute() + .await?; + + let err = table + .add(uri_struct_batch(1, &uri)) + .execute() + .await + .unwrap_err(); + + assert!( + err.to_string() + .contains("allow_external_blob_outside_bases"), + "got: {err}" + ); + assert_eq!(table.count_rows(None).await?, 0); + Ok(()) +} + +#[tokio::test] +async fn string_uri_input_round_trips_as_external_reference() -> Result<()> { + let tmp = tempdir().unwrap(); + let db = connect(tmp.path().join("db").to_str().unwrap()) + .execute() + .await?; + let payload: &[u8] = b"external-string-payload"; + let uri = write_payload_file_uri(tmp.path(), "payload.bin", payload); + let table = db + .create_empty_table("t", blob_table_schema()) + .execute() + .await?; + + table + .add(uri_string_batch(1, &uri)) + .allow_external_blob_outside_bases(true) + .execute() + .await?; + + let ids = collect_row_ids(&table).await?; + let bytes = table.fetch_blobs("image", &ids).await?; + assert_eq!(bytes.value(0), payload); + + let files = table.fetch_blob_files("image", &ids).await?; + let file = files[0].as_ref().expect("missing blob file"); + assert_eq!(file.uri(), Some(uri.as_str())); + Ok(()) +} + +#[tokio::test] +async fn string_uri_inside_registered_base_does_not_need_the_flag() -> Result<()> { + let tmp = tempdir().unwrap(); + let db_path = tmp.path().join("db"); + let external_base = tmp.path().join("external_base"); + let object_dir = external_base.join("objects"); + std::fs::create_dir_all(&object_dir).unwrap(); + let payload: &[u8] = b"mapped-in-base"; + let object_path = object_dir.join("mapped.bin"); + std::fs::write(&object_path, payload).unwrap(); + let object_uri = url::Url::from_file_path(&object_path).unwrap().to_string(); + let base_uri = url::Url::from_file_path(&external_base) + .unwrap() + .to_string(); + + let db = connect(db_path.to_str().unwrap()).execute().await?; + let table = db + .create_empty_table("t", blob_table_schema()) + .write_options(WriteOptions { + lance_write_params: Some(WriteParams { + initial_bases: Some(vec![BasePath { + id: 1, + name: Some("external".to_string()), + path: base_uri, + is_dataset_root: false, + }]), + ..Default::default() + }), + }) + .execute() + .await?; + + table + .add(uri_string_batch(1, &object_uri)) + .execute() + .await?; + + let ids = collect_row_ids(&table).await?; + let bytes = table.fetch_blobs("image", &ids).await?; + assert_eq!(bytes.value(0), payload); + Ok(()) +} + +#[tokio::test] +async fn external_uri_rows_mix_with_inline_rows() -> Result<()> { + let tmp = tempdir().unwrap(); + let db = connect(tmp.path().join("db").to_str().unwrap()) + .execute() + .await?; + let external_payload: &[u8] = b"external-bytes"; + let uri = write_payload_file_uri(tmp.path(), "payload.bin", external_payload); + let table = + create_inline_blob_table(&db, "t", &[1], &[Some(b"inline-bytes".as_slice())]).await?; + + table + .add(uri_string_batch(2, &uri)) + .allow_external_blob_outside_bases(true) + .execute() + .await?; + + let pairs = collect_id_rowid(&table).await?; + let row_ids: Vec = pairs.iter().map(|(_, r)| *r).collect(); + let bytes = table.fetch_blobs("image", &row_ids).await?; + for (i, (id, _)) in pairs.iter().enumerate() { + match id { + 1 => assert_eq!(bytes.value(i), b"inline-bytes"), + 2 => assert_eq!(bytes.value(i), external_payload), + _ => unreachable!(), + } + } + Ok(()) +} + +#[tokio::test] +async fn malformed_string_uri_is_rejected_at_write() -> Result<()> { + let tmp = tempdir().unwrap(); + let db = connect(tmp.path().join("db").to_str().unwrap()) + .execute() + .await?; + let table = db + .create_empty_table("t", blob_table_schema()) + .execute() + .await?; + + let err = table + .add(uri_string_batch(1, "not a uri")) + .allow_external_blob_outside_bases(true) + .execute() + .await + .unwrap_err(); + + assert!(err.to_string().contains("not a uri"), "got: {err}"); + assert_eq!(table.count_rows(None).await?, 0); + Ok(()) +} diff --git a/rust/lancedb/tests/first_class_function_slice1.rs b/rust/lancedb/tests/first_class_function_slice1.rs index dc565d309..ce020bd53 100644 --- a/rust/lancedb/tests/first_class_function_slice1.rs +++ b/rust/lancedb/tests/first_class_function_slice1.rs @@ -20,25 +20,6 @@ fn job_result(name: &str) -> Value { serde_json::from_str::(&fixture(name)).expect("remote Job fixture")["result"].clone() } -fn assert_no_secret_values(value: &Value) { - match value { - Value::Object(values) => { - for (key, value) in values { - assert!( - !matches!( - key.as_str(), - "secret_value" | "secret_values" | "resolved_secret" | "resolved_secrets" - ), - "client canonical value must not model resolved secret material" - ); - assert_no_secret_values(value); - } - } - Value::Array(values) => values.iter().for_each(assert_no_secret_values), - _ => {} - } -} - #[test] fn function_version_job_result_matches_shared_canonical_golden() { let result = job_result("remote_function_job.json"); @@ -47,7 +28,6 @@ fn function_version_job_result_matches_shared_canonical_golden() { assert_eq!(version.name(), "embed"); assert_eq!(version.version(), "fv_01K3EXACT"); assert_eq!(version.runtime_digest(), "sha256:runtime"); - assert_eq!(version.required_secrets(), &["HF_TOKEN"]); assert_eq!( version.to_canonical_json().expect("canonical JSON"), fixture("remote_function_version.canonical.json").trim() @@ -84,7 +64,6 @@ fn application_and_binding_match_shared_remote_goldens() { let binding = FunctionBinding::from_json(&fixture("remote_function_binding.json")) .expect("binding fixture"); - assert_eq!(binding.revision(), 3); assert_eq!(binding.function().version, "fv_01K3TEXT"); assert_eq!(binding.outputs()[0].output_ordinal, 0); assert_eq!(binding.outputs()[1].output_ordinal, 1); @@ -163,21 +142,3 @@ fn floating_point_application_literals_are_rejected_consistently() { .contains("floating-point Function literals") ); } - -#[test] -fn canonical_client_values_contain_secret_names_only() { - let result = job_result("remote_function_job.json"); - let version = FunctionVersion::from_json(&result.to_string()).expect("FunctionVersion result"); - let canonical: Value = serde_json::from_str( - &version - .to_canonical_json() - .expect("canonical FunctionVersion"), - ) - .expect("canonical JSON"); - - assert_eq!( - canonical["required_secrets"], - serde_json::json!(["HF_TOKEN"]) - ); - assert_no_secret_values(&canonical); -} diff --git a/rust/lancedb/tests/first_class_function_slice2.rs b/rust/lancedb/tests/first_class_function_slice2.rs index 3bae57122..93252dde4 100644 --- a/rust/lancedb/tests/first_class_function_slice2.rs +++ b/rust/lancedb/tests/first_class_function_slice2.rs @@ -6,7 +6,6 @@ use std::path::PathBuf; use lancedb::Error; use lancedb::function::FunctionRegistrationRequest; -use serde_json::Value; fn fixture(name: &str) -> String { let path = PathBuf::from(env!("CARGO_MANIFEST_DIR")) @@ -15,25 +14,6 @@ fn fixture(name: &str) -> String { fs::read_to_string(path).expect("fixture must be readable") } -fn assert_no_secret_values(value: &Value) { - match value { - Value::Object(values) => { - for (key, value) in values { - assert!( - !matches!( - key.as_str(), - "secret_value" | "secret_values" | "resolved_secret" | "resolved_secrets" - ), - "registration requests must not model resolved secret material" - ); - assert_no_secret_values(value); - } - } - Value::Array(values) => values.iter().for_each(assert_no_secret_values), - _ => {} - } -} - #[test] fn registration_request_matches_shared_canonical_golden() { let request = FunctionRegistrationRequest::from_json(&fixture( @@ -42,16 +22,10 @@ fn registration_request_matches_shared_canonical_golden() { .expect("registration request"); assert_eq!(request.name, "normalize_score"); assert_eq!(request.artifact.adapter.kind, "scalar_to_arrow_batch"); - assert_eq!(request.required_secrets, ["API_TOKEN"]); assert_eq!( request.to_canonical_json().expect("canonical request"), fixture("remote_function_registration_request.canonical.json").trim() ); - - let value: Value = - serde_json::from_str(&request.to_canonical_json().expect("canonical request")) - .expect("request JSON"); - assert_no_secret_values(&value); } #[tokio::test] diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json new file mode 100644 index 000000000..ff26e4c08 --- /dev/null +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json @@ -0,0 +1,333 @@ +{ + "valid": [ + { + "arrow_type": "bool", + "json": { + "type": "bool" + } + }, + { + "arrow_type": "int8", + "json": { + "type": "int8" + } + }, + { + "arrow_type": "int16", + "json": { + "type": "int16" + } + }, + { + "arrow_type": "int32", + "json": { + "type": "int32" + } + }, + { + "arrow_type": "int64", + "json": { + "type": "int64" + } + }, + { + "arrow_type": "uint8", + "json": { + "type": "uint8" + } + }, + { + "arrow_type": "uint16", + "json": { + "type": "uint16" + } + }, + { + "arrow_type": "uint32", + "json": { + "type": "uint32" + } + }, + { + "arrow_type": "uint64", + "json": { + "type": "uint64" + } + }, + { + "arrow_type": "float16", + "json": { + "type": "float16" + } + }, + { + "arrow_type": "float32", + "json": { + "type": "float32" + } + }, + { + "arrow_type": "float64", + "json": { + "type": "float64" + } + }, + { + "arrow_type": "utf8", + "json": { + "type": "utf8" + } + }, + { + "arrow_type": "binary", + "json": { + "type": "binary" + } + }, + { + "arrow_type": "date32", + "json": { + "type": "date32" + } + }, + { + "arrow_type": "date64", + "json": { + "type": "date64" + } + }, + { + "arrow_type": "list", + "json": { + "type": "list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "float32" + } + } + ] + } + }, + { + "arrow_type": "large_list", + "json": { + "type": "large_list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "float32" + } + } + ] + } + }, + { + "arrow_type": "list", + "json": { + "type": "list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "int64" + } + } + ] + } + }, + { + "arrow_type": "large_list", + "json": { + "type": "large_list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "int64" + } + } + ] + } + }, + { + "arrow_type": "list", + "json": { + "type": "list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "utf8" + } + } + ] + } + }, + { + "arrow_type": "large_list", + "json": { + "type": "large_list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "utf8" + } + } + ] + } + }, + { + "arrow_type": "fixed_size_list", + "json": { + "type": "fixed_size_list", + "length": 384, + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "float32" + } + } + ] + } + }, + { + "arrow_type": "fixed_size_list", + "json": { + "type": "fixed_size_list", + "length": 8, + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "float16" + } + } + ] + } + }, + { + "arrow_type": "fixed_size_list", + "json": { + "type": "fixed_size_list", + "length": 1, + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "uint8" + } + } + ] + } + }, + { + "arrow_type": "list>", + "json": { + "type": "list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "float32" + } + } + ] + } + } + ] + } + }, + { + "arrow_type": "fixed_size_list, 2>", + "json": { + "type": "fixed_size_list", + "length": 2, + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "int32" + } + } + ] + } + } + ] + } + }, + { + "arrow_type": "list>", + "json": { + "type": "list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "fixed_size_list", + "length": 3, + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "float32" + } + } + ] + } + } + ] + } + } + ], + "server_only": [ + { + "arrow_type": "null", + "json": { + "type": "null" + } + } + ], + "invalid": [ + "", + "list<>", + "list[3]", + "fixed_size_list", + "fixed_size_list", + "fixed_size_list", + "map", + "decimal128(10, 2)", + "timestamp[us]", + "struct" + ] +} \ No newline at end of file diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json new file mode 100644 index 000000000..fd3944412 --- /dev/null +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json @@ -0,0 +1,78 @@ +{ + "new_columns": [ + { + "name": "embedding", + "all_null": true + } + ], + "function": { + "application": { + "function": { + "name": "embed", + "version": "fv_01K3EXACT" + }, + "inputs": [ + { + "parameter": "text", + "kind": "column", + "value": { + "path": "description" + } + } + ], + "output": { + "kind": "scalar", + "arrow_type": "fixed_size_list", + "nullable": false + } + }, + "binding_metadata_version": 1, + "input_bindings": [ + { + "parameter": "text", + "field_path": "description", + "arrow_type": "utf8", + "nullable": true + } + ], + "input_schema": { + "fields": [ + { + "name": "text", + "nullable": true, + "type": { + "type": "utf8" + } + } + ] + }, + "output_schema": { + "fields": [ + { + "name": "embedding", + "nullable": true, + "type": { + "type": "fixed_size_list", + "length": 3, + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "float32" + } + } + ] + } + } + ] + }, + "outputs": [ + { + "result_field": "$value", + "output_name": "embedding", + "output_ordinal": 0 + } + ] + } +} diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.canonical.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.canonical.json index 05da9fe39..b91dd1061 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.canonical.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.canonical.json @@ -1 +1 @@ -{"columns":{"normalized_text":"search_text","token_count":"search_token_count"},"function":{"name":"text_features","version":"fv_01K3TEXT"},"group_id":"fg_01K3TEXT","inputs":[{"kind":"column","parameter":"title","value":{"path":"title"}},{"kind":"column","parameter":"body","value":{"path":"body"}}],"output":{"fields":[{"arrow_type":"utf8","name":"normalized_text","nullable":false},{"arrow_type":"int64","name":"token_count","nullable":false}],"kind":"named_struct"}} +{"columns":{"normalized_text":"search_text","token_count":"search_token_count"},"function":{"name":"text_features","version":"fv_01K3TEXT"},"inputs":[{"kind":"column","parameter":"title","value":{"path":"title"}},{"kind":"column","parameter":"body","value":{"path":"body"}}],"output":{"fields":[{"arrow_type":"utf8","name":"normalized_text","nullable":false},{"arrow_type":"int64","name":"token_count","nullable":false}],"kind":"named_struct"}} diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.json index 44aeff460..177821aaa 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.json @@ -11,7 +11,6 @@ {"name": "token_count", "arrow_type": "int64", "nullable": false} ] }, - "group_id": "fg_01K3TEXT", "columns": { "normalized_text": "search_text", "token_count": "search_token_count" diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application_float.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application_float.json index 47724eee0..c23c8f3ca 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application_float.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application_float.json @@ -3,6 +3,5 @@ "inputs": [ {"parameter": "threshold", "kind": "literal", "value": 1e-7} ], - "output": {"kind": "scalar", "arrow_type": "bool", "nullable": false}, - "group_id": "fg_01K3FLOAT" + "output": {"kind": "scalar", "arrow_type": "bool", "nullable": false} } diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.canonical.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.canonical.json index 7bf93b8a8..143190f78 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.canonical.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.canonical.json @@ -1 +1 @@ -{"binding_id":"fb_01K3TEXT","function":{"name":"text_features","version":"fv_01K3TEXT"},"group_id":"fg_01K3TEXT","input_schema":{"fields":[{"name":"title","nullable":true,"type":{"type":"utf8"}},{"name":"body","nullable":true,"type":{"type":"utf8"}}]},"inputs":[{"arrow_type":"utf8","field_id":11,"field_path":"title","nullable":true,"parameter":"title"},{"arrow_type":"utf8","field_id":12,"field_path":"body","nullable":true,"parameter":"body"}],"output_schema":{"fields":[{"name":"search_text","nullable":true,"type":{"type":"utf8"}},{"name":"search_token_count","nullable":true,"type":{"type":"int64"}}]},"outputs":[{"arrow_type":"utf8","nullable":false,"output_field_id":21,"output_name":"search_text","output_ordinal":0,"result_field":"normalized_text"},{"arrow_type":"int64","nullable":false,"output_field_id":22,"output_name":"search_token_count","output_ordinal":1,"result_field":"token_count"}],"revision":3} +{"binding_id":"fb_01K3TEXT","function":{"name":"text_features","version":"fv_01K3TEXT"},"input_schema":{"fields":[{"name":"title","nullable":true,"type":{"type":"utf8"}},{"name":"body","nullable":true,"type":{"type":"utf8"}}]},"inputs":[{"arrow_type":"utf8","field_id":11,"field_path":"title","nullable":true,"parameter":"title"},{"arrow_type":"utf8","field_id":12,"field_path":"body","nullable":true,"parameter":"body"}],"output_schema":{"fields":[{"name":"search_text","nullable":true,"type":{"type":"utf8"}},{"name":"search_token_count","nullable":true,"type":{"type":"int64"}}]},"outputs":[{"arrow_type":"utf8","nullable":false,"output_field_id":21,"output_name":"search_text","output_ordinal":0,"result_field":"normalized_text"},{"arrow_type":"int64","nullable":false,"output_field_id":22,"output_name":"search_token_count","output_ordinal":1,"result_field":"token_count"}]} diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.json index 1a2053e42..dec3fa3f8 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.json @@ -1,8 +1,6 @@ { "binding_id": "fb_01K3TEXT", - "revision": 3, "function": {"name": "text_features", "version": "fv_01K3TEXT"}, - "group_id": "fg_01K3TEXT", "inputs": [ {"parameter": "title", "field_id": 11, "field_path": "title", "arrow_type": "utf8", "nullable": true}, {"parameter": "body", "field_id": 12, "field_path": "body", "arrow_type": "utf8", "nullable": true} @@ -23,5 +21,5 @@ {"name": "search_token_count", "nullable": true, "type": {"type": "int64"}} ] }, - "future_binding": {"metadata_revision": 1} + "future_binding": {"mode": "managed"} } diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_job.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_job.json index 6ba4eb226..39a279692 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_job.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_job.json @@ -24,7 +24,6 @@ }, "runtime_digest": "sha256:runtime", "environment_digest": "sha256:environment", - "required_secrets": ["HF_TOKEN"], "created_at": "2026-08-21T00:00:00Z" }, "future_job": {"trace_id": "trace-1"} diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.canonical.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.canonical.json index 24fa2cf30..a2f2c4c21 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.canonical.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.canonical.json @@ -1 +1 @@ -{"artifact":{"adapter":{"kind":"scalar_to_arrow_batch","version":1},"content":{"data":"ZnJvbSBfX2Z1dHVyZV9fIGltcG9ydCBhbm5vdGF0aW9ucwoKZGVmIG5vcm1hbGl6ZV9zY29yZSh2YWx1ZTogZmxvYXQpIC0+IGZsb2F0OgogICAgcmV0dXJuIHZhbHVlIC8gMTAwLjAK","encoding":"base64"},"digest":"sha256:760784bdcef57b802f389b97804cc0b618aae39e86044733451bcd0b13089a7f","entrypoint":"normalize_score","kind":"python_callable"},"name":"normalize_score","required_secrets":["API_TOKEN"],"runtime":{"env":{"MODE":"test"},"environment":{"kind":"pip","packages":["numpy>=2"]},"kind":"python","python_version":"3.12"},"signature":{"inputs":[{"arrow_type":"float64","name":"value","nullable":false}],"output":{"arrow_type":"float64","kind":"scalar","nullable":false}}} +{"artifact":{"adapter":{"kind":"scalar_to_arrow_batch","version":1},"content":{"data":"ZnJvbSBfX2Z1dHVyZV9fIGltcG9ydCBhbm5vdGF0aW9ucwoKZGVmIG5vcm1hbGl6ZV9zY29yZSh2YWx1ZTogZmxvYXQpIC0+IGZsb2F0OgogICAgcmV0dXJuIHZhbHVlIC8gMTAwLjAK","encoding":"base64"},"digest":"sha256:760784bdcef57b802f389b97804cc0b618aae39e86044733451bcd0b13089a7f","entrypoint":"normalize_score","kind":"python_callable"},"name":"normalize_score","runtime":{"env":{"MODE":"test"},"environment":{"kind":"pip","packages":["numpy>=2"]},"kind":"python","python_version":"3.12"},"signature":{"inputs":[{"arrow_type":"float64","name":"value","nullable":false}],"output":{"arrow_type":"float64","kind":"scalar","nullable":false}}} diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.json index bbfec3169..092d76dc2 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.json @@ -39,8 +39,5 @@ "env": { "MODE": "test" } - }, - "required_secrets": [ - "API_TOKEN" - ] + } } diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_version.canonical.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_version.canonical.json index 7ab632a98..2670ad0b2 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_version.canonical.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_version.canonical.json @@ -1 +1 @@ -{"artifact":{"digest":"sha256:code","entrypoint":"embed","kind":"python_callable"},"created_at":"2026-08-21T00:00:00Z","environment_digest":"sha256:environment","name":"embed","required_secrets":["HF_TOKEN"],"runtime":{"env":{"TOKENIZERS_PARALLELISM":"false"},"environment":{"kind":"pip","packages":["sentence-transformers>=3"]},"kind":"python","python_version":"3.12"},"runtime_digest":"sha256:runtime","signature":{"inputs":[{"arrow_type":"utf8","name":"text","nullable":true}],"output":{"arrow_type":"list","kind":"scalar","nullable":false}},"version":"fv_01K3EXACT"} +{"artifact":{"digest":"sha256:code","entrypoint":"embed","kind":"python_callable"},"created_at":"2026-08-21T00:00:00Z","environment_digest":"sha256:environment","name":"embed","runtime":{"env":{"TOKENIZERS_PARALLELISM":"false"},"environment":{"kind":"pip","packages":["sentence-transformers>=3"]},"kind":"python","python_version":"3.12"},"runtime_digest":"sha256:runtime","signature":{"inputs":[{"arrow_type":"utf8","name":"text","nullable":true}],"output":{"arrow_type":"list","kind":"scalar","nullable":false}},"version":"fv_01K3EXACT"} diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_grouped_declaration_request.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_multi_output_declaration_request.json similarity index 97% rename from rust/lancedb/tests/fixtures/first_class_functions/v1/remote_grouped_declaration_request.json rename to rust/lancedb/tests/fixtures/first_class_functions/v1/remote_multi_output_declaration_request.json index d0b42cc99..c4ef5009f 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_grouped_declaration_request.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_multi_output_declaration_request.json @@ -17,7 +17,6 @@ {"name": "token_count", "arrow_type": "int64", "nullable": false} ] }, - "group_id": "fg_01K3TEXT", "columns": {"normalized_text": "search_text"} }, "binding_metadata_version": 1, diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_refresh_job.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_refresh_job.json index bb2490bc8..c2d49827d 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_refresh_job.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_refresh_job.json @@ -3,7 +3,7 @@ "job_type": "refresh_function_columns", "job_state": "DONE", "creation_ms": 1787270400001, - "spec": {"table": "documents", "binding_revision": 3}, + "spec": {"table": "documents", "binding_id": "fb_01K3TEXT"}, "result": { "rows_assigned": 999998800, "rows_failed": 0, diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_scalar_declaration_request.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_scalar_declaration_request.json index 0aaa0cf72..357834de0 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_scalar_declaration_request.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_scalar_declaration_request.json @@ -8,8 +8,7 @@ "inputs": [ {"parameter": "text", "kind": "column", "value": {"path": "description"}} ], - "output": {"kind": "scalar", "arrow_type": "list", "nullable": false}, - "group_id": "fg_scalar" + "output": {"kind": "scalar", "arrow_type": "list", "nullable": false} }, "binding_metadata_version": 1, "input_bindings": [