diff --git a/.bumpversion.toml b/.bumpversion.toml index 1c4aea809..3d57dc0fe 100644 --- a/.bumpversion.toml +++ b/.bumpversion.toml @@ -1,5 +1,5 @@ [tool.bumpversion] -current_version = "0.38.0-beta.10" +current_version = "0.38.0-beta.12" parse = """(?x) (?P0|[1-9]\\d*)\\. (?P0|[1-9]\\d*)\\. diff --git a/.github/workflows/npm-publish.yml b/.github/workflows/npm-publish.yml index 72ef5ad13..ee6906e00 100644 --- a/.github/workflows/npm-publish.yml +++ b/.github/workflows/npm-publish.yml @@ -40,40 +40,31 @@ jobs: - target: aarch64-apple-darwin host: macos-latest features: fp16kernels + # Fat LTO was ~111 of this job's ~113 minutes. + lto: thin + codegen_units: 16 pre_build: |- brew install protobuf - # Fat LTO (the workspace default in .cargo/config.toml) is - # single-threaded and is the peak-memory step of the build. On - # this runner it accounted for ~111 of the job's ~113 minutes, - # making it the critical path of the entire publish pipeline. - # ThinLTO parallelizes it across the runner's cores, for a few - # percent of runtime performance. - export CARGO_PROFILE_RELEASE_LTO=thin - export CARGO_PROFILE_RELEASE_CODEGEN_UNITS=16 - target: x86_64-pc-windows-msvc host: windows-2025 features: "," + # The lower peak also keeps this on the standard 4-core runner. + lto: thin + codegen_units: 16 pre_build: |- choco install --no-progress protoc ninja nasm tail -n 1000 /c/ProgramData/chocolatey/logs/chocolatey.log # There is an issue where choco doesn't add nasm to the path export PATH="$PATH:/c/Program Files/NASM" nasm -v - # See the ThinLTO note on aarch64-apple-darwin above. Keeping - # peak memory down is also what lets this run on the standard - # 4-core runner: the 8-core larger runner was only needed to - # stop fat LTO from OOMing rustc-LLVM. - export CARGO_PROFILE_RELEASE_LTO=thin - export CARGO_PROFILE_RELEASE_CODEGEN_UNITS=16 - target: aarch64-pc-windows-msvc host: windows-2025 features: "," + lto: thin + codegen_units: 16 pre_build: |- choco install --no-progress protoc rustup target add aarch64-pc-windows-msvc - # See the ThinLTO note on aarch64-apple-darwin above. - export CARGO_PROFILE_RELEASE_LTO=thin - export CARGO_PROFILE_RELEASE_CODEGEN_UNITS=16 - target: x86_64-unknown-linux-gnu host: ubuntu-latest features: fp16kernels @@ -103,6 +94,14 @@ jobs: # https://github.com/napi-rs/napi-rs/blob/main/debian-aarch64.Dockerfile docker: ghcr.io/napi-rs/napi-rs/nodejs-rust:lts-debian-aarch64 features: "fp16kernels" + # Fat LTO OOM-killed rustc every nightly; even with lld it peaked + # at 31391 MiB of the runner's 32 GiB. + lto: thin + codegen_units: 16 + # arm64 Linux links through GNU `ld` where x86_64 defaults to + # `rust-lld`, which is why only arm64 OOM'd. lld cut the largest + # linker process 7.0 -> 4.0 GiB (lancedb/sophon#7313). + linker: /tmp/aarch64-lld-clang pre_build: |- set -e && apt-get update && @@ -112,9 +111,30 @@ jobs: # AT_HWCAP2 (added in Linux 3.17). Define it for aws-lc-sys. export CFLAGS="$CFLAGS -DAT_HWCAP2=26" && rustup target add aarch64-unknown-linux-gnu + # Not `&&`-chained: in dash, errexit does not fire for a + # non-final command in an `&&` list, so failures were ignored. + # + # A wrapper rather than `-C link-arg` because the per-target + # rustflags variable does not reach every unit that links, while + # the linker variable does. `clang` because GCC silently ignores + # `-fuse-ld=lld` unless built with lld support. Two echoes + # because printf's newline escape gets rewritten to `;` between + # here and the container. + echo '#!/bin/sh' > /tmp/aarch64-lld-clang + echo 'exec clang --target=aarch64-unknown-linux-gnu --sysroot=/usr/aarch64-unknown-linux-gnu/aarch64-unknown-linux-gnu/sysroot --gcc-toolchain=/usr/aarch64-unknown-linux-gnu -fuse-ld=lld "$@"' >> /tmp/aarch64-lld-clang + chmod 0755 /tmp/aarch64-lld-clang + # Fail now, not at the cdylib link ~30 minutes later. Linking at + # all also proves lld resolved; clang errors out when it cannot. + echo 'int main(void){return 0;}' > /tmp/probe.c + /tmp/aarch64-lld-clang /tmp/probe.c -o /tmp/probe + readelf -h /tmp/probe | grep AArch64 - target: aarch64-unknown-linux-musl host: ubuntu-2404-8x-x64 features: "," + # Fat LTO took the whole runner down. lld cannot help: it died + # inside rustc's LLVM, before any linker was spawned. + lto: thin + codegen_units: 16 pre_build: |- set -e && sudo apt-get update && @@ -123,6 +143,19 @@ jobs: export EXTRA_ARGS="-x" name: build - ${{ matrix.settings.target }} runs-on: ${{ matrix.settings.host }} + # On the job, not exported from `pre_build`: `Swatinem/rust-cache` hashes + # `CARGO_*` into its cache key before any step runs, so a step-local export + # leaves the key unchanged while cargo still rebuilds cold. The ThinLTO + # legs had been doing that every run. + # + # Not `RUSTFLAGS`: setting it, even to "", discards every config-file + # rustflag, silently dropping .cargo/config.toml's `target-cpu` and + # `target-feature` from the published binaries. + env: + CARGO_PROFILE_RELEASE_LTO: ${{ matrix.settings.lto || 'fat' }} + CARGO_PROFILE_RELEASE_CODEGEN_UNITS: ${{ matrix.settings.codegen_units || '1' }} + # Empty elsewhere: a per-target variable is only read for that triple. + CARGO_TARGET_AARCH64_UNKNOWN_LINUX_GNU_LINKER: ${{ matrix.settings.linker }} defaults: run: working-directory: nodejs @@ -169,19 +202,15 @@ jobs: # creating ref). The nightly cadence also keeps entries inside # GitHub's 7-day eviction window, which a tag-only trigger would not. save-if: ${{ github.ref == 'refs/heads/main' }} - # Docker builds can use rust-cache too. `target/` already lives on the - # host because the whole workspace is bind-mounted into the container, and - # rust-cache's prune and save run host-side, so they can manage it -- which - # is what keeps the entry to dependency artifacts rather than a multi-GB - # copy of everything. + # Docker builds can use rust-cache too: the workspace is bind-mounted, so + # `target/` lives on the host and rust-cache's prune keeps the entry + # small. # # Two differences from the native builds. The container's CARGO_HOME is - # bind-mounted from `.cargo-cache` rather than the host's ~/.cargo, so that - # has to be cached explicitly. And the key is derived from the *host* rustc - # version, which is not the compiler that produced these artifacts; that is - # safe because cargo fingerprints the real compiler and rebuilds on a - # mismatch, it just means a base-image toolchain bump costs one cold build - # instead of invalidating the key. + # bind-mounted from `.cargo-cache` rather than ~/.cargo, so that is cached + # explicitly. And the key uses the *host* rustc version, not the compiler + # that built these artifacts -- safe, since cargo fingerprints the real + # one; a base-image bump just costs one cold build. - name: Cache cargo (docker builds) uses: Swatinem/rust-cache@v2 if: ${{ matrix.settings.docker }} @@ -210,9 +239,14 @@ jobs: # cache step above saves. Previously the registry mounts pointed at # `.cargo/...`, a path nothing cached, so the container re-downloaded # the whole crate registry on every run. + # + # `docker run` inherits nothing; `-e NAME` carries the job's `env:` in. options: "--user 0:0 -v ${{ github.workspace }}/.cargo-cache/git/db:/usr/local/cargo/git/db \ -v ${{ github.workspace }}/.cargo-cache/registry/cache:/usr/local/cargo/registry/cache \ -v ${{ github.workspace }}/.cargo-cache/registry/index:/usr/local/cargo/registry/index \ + -e CARGO_PROFILE_RELEASE_LTO \ + -e CARGO_PROFILE_RELEASE_CODEGEN_UNITS \ + -e CARGO_TARGET_AARCH64_UNKNOWN_LINUX_GNU_LINKER \ -v ${{ github.workspace }}:/build -w /build/nodejs" run: | set -e @@ -256,6 +290,18 @@ jobs: if: always() run: df -h shell: bash + - name: Report peak memory + if: always() && runner.os == 'Linux' + shell: bash + run: | + peak=$(find /sys/fs/cgroup -name memory.peak -readable \ + -exec cat {} + 2>/dev/null | sort -n | tail -1) + if [ -n "$peak" ]; then + echo "peak memory: $((peak / 1024 / 1024)) MiB" + else + echo "peak memory: unavailable (no readable cgroup v2 memory.peak)" + fi + free -g || true - name: Upload artifact uses: actions/upload-artifact@v7 with: diff --git a/Cargo.lock b/Cargo.lock index e27f9f271..f12988efa 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5402,7 +5402,7 @@ dependencies = [ [[package]] name = "lancedb" -version = "0.38.0-beta.10" +version = "0.38.0-beta.12" dependencies = [ "ahash", "anyhow", @@ -5490,7 +5490,7 @@ dependencies = [ [[package]] name = "lancedb-nodejs" -version = "0.38.0-beta.10" +version = "0.38.0-beta.12" dependencies = [ "arrow-array", "arrow-buffer", @@ -5515,7 +5515,7 @@ dependencies = [ [[package]] name = "lancedb-python" -version = "0.38.0-beta.10" +version = "0.38.0-beta.12" dependencies = [ "arrow", "async-trait", diff --git a/docs/src/java/java.md b/docs/src/java/java.md index 06dc267f3..e19880e29 100644 --- a/docs/src/java/java.md +++ b/docs/src/java/java.md @@ -14,7 +14,7 @@ Add the following dependency to your `pom.xml`: com.lancedb lancedb-core - 0.38.0-beta.10 + 0.38.0-beta.12 ``` diff --git a/docs/src/python/python.md b/docs/src/python/python.md index 3cb996a15..5f359f4b7 100644 --- a/docs/src/python/python.md +++ b/docs/src/python/python.md @@ -223,9 +223,13 @@ tokens = list( Blob columns store large binary values out of line so they can be read lazily instead of being materialized with the rest of the row. -::: lancedb.blob +`lancedb.BlobType` is `lance.blob.BlobType` when pylance is installed. Without +pylance, LanceDB uses a matching `lance.blob.v2` extension type so blob columns +still work. Queries return descriptors. Call +[`fetch_blob_files`][lancedb.table.Table.fetch_blob_files] for lazy reads or +[`fetch_blobs`][lancedb.table.Table.fetch_blobs] for eager bytes. -::: lancedb.BlobType +::: lancedb.blob ::: lancedb._blob.BlobFile options: diff --git a/java/lancedb-core/pom.xml b/java/lancedb-core/pom.xml index 25e3b10e3..3864ed127 100644 --- a/java/lancedb-core/pom.xml +++ b/java/lancedb-core/pom.xml @@ -8,7 +8,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.10 + 0.38.0-beta.12 ../pom.xml diff --git a/java/pom.xml b/java/pom.xml index a2cec19c0..b3521f9c6 100644 --- a/java/pom.xml +++ b/java/pom.xml @@ -6,7 +6,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.10 + 0.38.0-beta.12 pom ${project.artifactId} LanceDB Java SDK Parent POM diff --git a/nodejs/Cargo.toml b/nodejs/Cargo.toml index fd08e7a5e..c3b69424f 100644 --- a/nodejs/Cargo.toml +++ b/nodejs/Cargo.toml @@ -1,7 +1,7 @@ [package] name = "lancedb-nodejs" edition.workspace = true -version = "0.38.0-beta.10" +version = "0.38.0-beta.12" publish = false license.workspace = true description.workspace = true diff --git a/nodejs/__test__/arrow.test.ts b/nodejs/__test__/arrow.test.ts index c5bbbf169..83b4fae46 100644 --- a/nodejs/__test__/arrow.test.ts +++ b/nodejs/__test__/arrow.test.ts @@ -1,11 +1,16 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright The LanceDB Authors +import * as fs from "node:fs"; +import * as vm from "node:vm"; import * as arrow15 from "apache-arrow-15"; import * as arrow16 from "apache-arrow-16"; import * as arrow17 from "apache-arrow-17"; import * as arrow18 from "apache-arrow-18"; import { + Field as CurrentField, + LargeBinary as CurrentLargeBinary, + Schema as CurrentSchema, Vector as CurrentVector, convertToTable, tableFromIPC as currentTableFromIPC, @@ -36,6 +41,59 @@ function sampleRecords(): Array> { }, ]; } + +it("serializes an Arrow Table created in another JavaScript realm", async () => { + const context = vm.createContext({ + TextDecoder, + TextEncoder, + console, + setTimeout, + clearTimeout, + }); + vm.runInContext( + fs.readFileSync( + require.resolve("apache-arrow-15/Arrow.es2015.min"), + "utf8", + ), + context, + ); + const foreignTable: unknown = vm.runInContext( + "Arrow.tableFromArrays({ id: new Int32Array([1, 2, 3]), text: ['foo', 'bar', 'baz'] })", + context, + ); + + const foreignMetadata = ( + foreignTable as { schema: { metadata: Map } } + ).schema.metadata; + expect(foreignMetadata).not.toBeInstanceOf(Map); + + const buf = await fromDataToBuffer( + foreignTable as Parameters[0], + ); + const actual = currentTableFromIPC(buf); + + expect(actual.numRows).toBe(3); + expect(actual.getChild("id")?.toJSON()).toEqual([1, 2, 3]); + expect(actual.getChild("text")?.toJSON()).toEqual(["foo", "bar", "baz"]); +}); + +it("preserves field metadata from a provided schema", async function () { + const jsonMetadata = new Map([["ARROW:extension:name", "lance.json"]]); + const schema = new CurrentSchema([ + new CurrentField("meta", new CurrentLargeBinary(), true, jsonMetadata), + ]); + + const table = makeArrowTable( + [{ meta: Buffer.from(JSON.stringify({ source: "test" })) }], + { schema }, + ); + + expect(table.schema.fields[0].metadata).toEqual(jsonMetadata); + + const roundTripped = currentTableFromIPC(await fromTableToBuffer(table)); + expect(roundTripped.schema.fields[0].metadata).toEqual(jsonMetadata); +}); + describe.each([arrow15, arrow16, arrow17, arrow18])( "Arrow", ( diff --git a/nodejs/__test__/embedding.test.ts b/nodejs/__test__/embedding.test.ts index 2a8494e0f..45d171a3d 100644 --- a/nodejs/__test__/embedding.test.ts +++ b/nodejs/__test__/embedding.test.ts @@ -187,6 +187,58 @@ describe("embedding functions", () => { const vector0 = JSON.parse(JSON.stringify(arr[0].vector)); expect(vector0).toEqual([1, 2, 3]); }); + it("should append multiple Python embeddings with the same alias", async () => { + @register("python-mock") + // biome-ignore lint/correctness/noUnusedVariables: the decorator registers this class + class MockEmbeddingFunction extends EmbeddingFunction { + ndims() { + return 3; + } + embeddingDataType(): Float { + return new Float32(); + } + async computeQueryEmbeddings(_data: string) { + return [1, 2, 3]; + } + async computeSourceEmbeddings(data: string[]) { + return data.map((value) => + value === "hello world" ? [1, 2, 3] : [4, 5, 6], + ); + } + } + + const metadata = new Map([ + [ + "embedding_functions", + '[{"source_column":"text1","vector_column":"vector1","name":"python-mock","model":{}},{"source_column":"text2","vector_column":"vector2","name":"python-mock","model":{}}]', + ], + ]); + const schema = new Schema( + [ + new Field("text1", new Utf8(), true), + new Field("text2", new Utf8(), true), + new Field( + "vector1", + new FixedSizeList(3, new Field("item", new Float32(), true)), + true, + ), + new Field( + "vector2", + new FixedSizeList(3, new Field("item", new Float32(), true)), + true, + ), + ], + metadata, + ); + + const db = await connect(tmpDir.name); + const table = await db.createEmptyTable("test", schema); + await table.add([{ text1: "hello world", text2: "goodbye world" }]); + + const rows = await table.query().toArray(); + expect(JSON.parse(JSON.stringify(rows[0].vector1))).toEqual([1, 2, 3]); + expect(JSON.parse(JSON.stringify(rows[0].vector2))).toEqual([4, 5, 6]); + }); it("should append generated vectors to a non-nullable schema", async () => { @register("non_nullable_schema_test") diff --git a/nodejs/__test__/table.test.ts b/nodejs/__test__/table.test.ts index dae640850..554c7fcd3 100644 --- a/nodejs/__test__/table.test.ts +++ b/nodejs/__test__/table.test.ts @@ -3561,6 +3561,27 @@ describe("when creating an empty table", () => { expect((actualSchema.fields[1].type as Float64).precision).toBe(2); }); + it("can add and query JSON data", async () => { + const schema = new Schema([ + new Field("id", new Int32(), true), + new Field( + "meta", + new Utf8(), + true, + new Map([["ARROW:extension:name", "arrow.json"]]), + ), + ]); + const table = await con.createEmptyTable("json", schema); + const meta = JSON.stringify({ x: 1 }); + + await table.add([{ id: 1, meta }]); + + const rows = await table.query().toArray(); + expect(rows).toHaveLength(1); + expect(rows[0].id).toBe(1); + expect(rows[0].meta).toBe(meta); + }); + it("can create an empty table from schema that specifies field types by name", async () => { const schemaLike = { fields: [ diff --git a/nodejs/lancedb/arrow.ts b/nodejs/lancedb/arrow.ts index b52ab50ef..1b6b98cc9 100644 --- a/nodejs/lancedb/arrow.ts +++ b/nodejs/lancedb/arrow.ts @@ -72,8 +72,7 @@ export type FieldLike = }; export type DataLike = - // biome-ignore lint/suspicious/noExplicitAny: - | import("apache-arrow").Data> + | import("apache-arrow").Data | { // biome-ignore lint/suspicious/noExplicitAny: type: any; @@ -82,6 +81,7 @@ export type DataLike = stride: number; nullable: boolean; children: DataLike[]; + dictionary?: { data: readonly DataLike[] }; get nullCount(): number; // biome-ignore lint/suspicious/noExplicitAny: values: Buffers[BufferType.DATA]; diff --git a/nodejs/lancedb/sanitize.ts b/nodejs/lancedb/sanitize.ts index 8fb2f1a0a..454c82247 100644 --- a/nodejs/lancedb/sanitize.ts +++ b/nodejs/lancedb/sanitize.ts @@ -94,17 +94,24 @@ export function sanitizeMetadata( if (metadataLike === undefined || metadataLike === null) { return undefined; } - if (!(metadataLike instanceof Map)) { + + let entries: IterableIterator<[unknown, unknown]>; + try { + entries = Map.prototype.entries.call(metadataLike); + } catch { throw Error("Expected metadata, if present, to be a Map"); } - for (const item of metadataLike) { - if (typeof item[0] !== "string" || typeof item[1] !== "string") { + + const metadata = new Map(); + for (const [key, value] of entries) { + if (typeof key !== "string" || typeof value !== "string") { throw Error( "Expected metadata, if present, to be a Map but it had non-string keys or values", ); } + metadata.set(key, value); } - return metadataLike as Map; + return metadata; } export function sanitizeInt(typeLike: object) { diff --git a/nodejs/lancedb/schema.ts b/nodejs/lancedb/schema.ts index e4749ef37..3a3ee9316 100644 --- a/nodejs/lancedb/schema.ts +++ b/nodejs/lancedb/schema.ts @@ -406,10 +406,11 @@ function matchingFields(fields: Field[], tree: FieldTree): Field[] { field.name, new Struct(matchingFields(struct.children, value)), field.nullable, + field.metadata, ), ); } else { - matches.push(new Field(field.name, value as DataType, field.nullable)); + matches.push(field); } } return matches; diff --git a/nodejs/npm/darwin-arm64/package.json b/nodejs/npm/darwin-arm64/package.json index ff6347c4d..2b8d43c3d 100644 --- a/nodejs/npm/darwin-arm64/package.json +++ b/nodejs/npm/darwin-arm64/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-darwin-arm64", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.12", "os": ["darwin"], "cpu": ["arm64"], "main": "lancedb.darwin-arm64.node", diff --git a/nodejs/npm/linux-arm64-gnu/package.json b/nodejs/npm/linux-arm64-gnu/package.json index ed99fee05..6656fdd5c 100644 --- a/nodejs/npm/linux-arm64-gnu/package.json +++ b/nodejs/npm/linux-arm64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-gnu", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.12", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-gnu.node", diff --git a/nodejs/npm/linux-arm64-musl/package.json b/nodejs/npm/linux-arm64-musl/package.json index 5b215dcc0..f8f3e151f 100644 --- a/nodejs/npm/linux-arm64-musl/package.json +++ b/nodejs/npm/linux-arm64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-musl", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.12", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-musl.node", diff --git a/nodejs/npm/linux-x64-gnu/package.json b/nodejs/npm/linux-x64-gnu/package.json index e0f5a9f26..20efa860a 100644 --- a/nodejs/npm/linux-x64-gnu/package.json +++ b/nodejs/npm/linux-x64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-gnu", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.12", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-gnu.node", diff --git a/nodejs/npm/linux-x64-musl/package.json b/nodejs/npm/linux-x64-musl/package.json index d42541707..4b735a687 100644 --- a/nodejs/npm/linux-x64-musl/package.json +++ b/nodejs/npm/linux-x64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-musl", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.12", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-musl.node", diff --git a/nodejs/npm/win32-arm64-msvc/package.json b/nodejs/npm/win32-arm64-msvc/package.json index 496a40720..35fee5ee0 100644 --- a/nodejs/npm/win32-arm64-msvc/package.json +++ b/nodejs/npm/win32-arm64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-arm64-msvc", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.12", "os": [ "win32" ], diff --git a/nodejs/npm/win32-x64-msvc/package.json b/nodejs/npm/win32-x64-msvc/package.json index 734013343..925211fbe 100644 --- a/nodejs/npm/win32-x64-msvc/package.json +++ b/nodejs/npm/win32-x64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-x64-msvc", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.12", "os": ["win32"], "cpu": ["x64"], "main": "lancedb.win32-x64-msvc.node", diff --git a/nodejs/package-lock.json b/nodejs/package-lock.json index 732a7c01c..b996b0810 100644 --- a/nodejs/package-lock.json +++ b/nodejs/package-lock.json @@ -1,12 +1,12 @@ { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.12", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.12", "cpu": [ "x64", "arm64" diff --git a/nodejs/package.json b/nodejs/package.json index 262d4c4a7..19c7f4d32 100644 --- a/nodejs/package.json +++ b/nodejs/package.json @@ -11,7 +11,7 @@ "ann" ], "private": false, - "version": "0.38.0-beta.10", + "version": "0.38.0-beta.12", "main": "dist/index.js", "exports": { ".": "./dist/index.js", diff --git a/python/Cargo.toml b/python/Cargo.toml index 3a8a05522..0ee561977 100644 --- a/python/Cargo.toml +++ b/python/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb-python" -version = "0.38.0-beta.10" +version = "0.38.0-beta.12" publish = false edition.workspace = true description = "Python bindings for LanceDB" diff --git a/python/python/lancedb/__init__.py b/python/python/lancedb/__init__.py index 8cb85a3ed..21ffc8860 100644 --- a/python/python/lancedb/__init__.py +++ b/python/python/lancedb/__init__.py @@ -6,7 +6,7 @@ import importlib.metadata import os from concurrent.futures import ThreadPoolExecutor from datetime import timedelta -from typing import Dict, Optional, Union, Any, List, Iterable +from typing import Dict, Optional, Union, Any, List, Iterable, TYPE_CHECKING __version__ = importlib.metadata.version("lancedb") @@ -20,7 +20,7 @@ from .db import AsyncConnection, DBConnection, LanceDBConnection from .remote import ClientConfig from .remote.db import RemoteDBConnection from .expr import Expr, col, lit, func -from .schema import blob, vector, BlobType +from .schema import blob, vector from .job import AsyncJob, Job from .functions import ( FunctionArtifactRequest as FunctionArtifactRequest, @@ -49,6 +49,19 @@ from .namespace import ( ) +if TYPE_CHECKING: + from lance.blob import BlobType as BlobType + + +def __getattr__(name: str): + if name == "BlobType": + from .schema import BlobType + + globals()["BlobType"] = BlobType + return BlobType + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + def _check_s3_bucket_with_dots( uri: str, storage_options: Optional[Dict[str, str]] ) -> None: diff --git a/python/python/lancedb/_blob.py b/python/python/lancedb/_blob.py index 926769f48..dc4ed37df 100644 --- a/python/python/lancedb/_blob.py +++ b/python/python/lancedb/_blob.py @@ -12,7 +12,7 @@ from typing import TYPE_CHECKING, Optional, Union import pyarrow as pa from .expr import Expr -from .schema import blob_v2_column_paths +from .schema import row_addressable_blob_v2_paths from .types import BlobMode, QueryProjection, QueryProjectionSpec if TYPE_CHECKING: @@ -119,7 +119,7 @@ def blob_v2_projection_sources( schema: pa.Schema, projection: QueryProjection, ) -> dict[str, str]: - blob_columns = blob_v2_column_paths(schema) + blob_columns = row_addressable_blob_v2_paths(schema) if not blob_columns: return {} columns = set(blob_columns) @@ -140,7 +140,9 @@ def v2_projection_needs_row_id( ) -> bool: if with_row_id: return False - return projection_includes_blob_column(projection, blob_v2_column_paths(schema)) + return projection_includes_blob_column( + projection, row_addressable_blob_v2_paths(schema) + ) def blob_auto_row_id_for_scan( @@ -270,7 +272,8 @@ def _iter_projection_pairs( if isinstance(expr, str): yield name, expr elif isinstance(expr, Expr): - yield name, expr.to_sql() + source = expr._column_name() + yield name, source if source is not None else expr.to_sql() return for column in projection: if isinstance(column, str): @@ -280,7 +283,8 @@ def _iter_projection_pairs( if isinstance(expr, str): yield name, expr elif isinstance(expr, Expr): - yield name, expr.to_sql() + source = expr._column_name() + yield name, source if source is not None else expr.to_sql() def _set_blob_column(tbl: pa.Table, output_name: str, blobs: pa.Array) -> pa.Table: diff --git a/python/python/lancedb/_lancedb.pyi b/python/python/lancedb/_lancedb.pyi index 593bceffa..7d7ca7f2a 100644 --- a/python/python/lancedb/_lancedb.pyi +++ b/python/python/lancedb/_lancedb.pyi @@ -87,6 +87,7 @@ class PyExpr: def contains(self, substr: "PyExpr") -> "PyExpr": ... def isin(self, values: List["PyExpr"]) -> "PyExpr": ... def cast(self, data_type: pa.DataType) -> "PyExpr": ... + def column_name(self) -> Optional[str]: ... def to_sql(self) -> str: ... def expr_col(name: str) -> PyExpr: ... @@ -608,6 +609,7 @@ class PyQueryRequest: filter: Optional[Union[str, bytes]] full_text_search: Optional[FullTextQuery] select: Optional[Union[str, List[str]]] + select_source_columns: Optional[Dict[str, str]] fast_search: Optional[bool] with_row_id: Optional[bool] use_lsm: Optional[bool] diff --git a/python/python/lancedb/expr.py b/python/python/lancedb/expr.py index d16ba95d7..80d01e29a 100644 --- a/python/python/lancedb/expr.py +++ b/python/python/lancedb/expr.py @@ -249,6 +249,10 @@ class Expr: # ── utilities ──────────────────────────────────────────────────────────── + def _column_name(self) -> str | None: + """Return the source name when this is a bare column expression.""" + return self._inner.column_name() + def to_sql(self) -> str: """Render the expression as a SQL string (useful for debugging).""" return self._inner.to_sql() @@ -312,7 +316,7 @@ def func(name: str, *args: ExprLike) -> Expr: -------- >>> from lancedb.expr import col, func >>> func("lower", col("name")) - Expr(lower(name)) + Expr(lower(`name`)) """ inner_args = [_coerce(a)._inner for a in args] return Expr(expr_func(name, inner_args)) diff --git a/python/python/lancedb/query.py b/python/python/lancedb/query.py index 9301d7df8..451384ad1 100644 --- a/python/python/lancedb/query.py +++ b/python/python/lancedb/query.py @@ -167,6 +167,12 @@ def _projection_to_scanner_kwargs(columns: QueryProjection) -> Dict[str, Any]: return {"columns": projection} +def _query_request_projection(req: "PyQueryRequest") -> QueryProjection: + if req.select_source_columns is not None: + return req.select_source_columns + return req.select + + def _scanner_kwargs_for_query( query: Query, blob_mode: BlobMode, @@ -2799,15 +2805,16 @@ class AsyncQueryBase(object): req = self._inner.to_query_request() schema = await self._table.schema() + projection = _query_request_projection(req) self._blob_auto_row_id = blob_auto_row_id_for_scan( schema, - req.select, + projection, with_row_id=self._with_row_id, ) if not self._blob_auto_row_id: self._blob_paths = () return - self._blob_paths = tuple(blob_v2_projection_sources(schema, req.select).keys()) + self._blob_paths = tuple(blob_v2_projection_sources(schema, projection).keys()) self._inner.with_row_id() def select(self, columns: Union[List[str], dict[str, str]]) -> Self: @@ -3894,14 +3901,15 @@ class AsyncHybridQuery(AsyncStandardQuery, AsyncVectorQueryBase): blob_paths: tuple[str, ...] = () if self._table is not None: schema = await self._table.schema() + projection = _query_request_projection(req) blob_auto_row_id = blob_auto_row_id_for_scan( schema, - req.select, + projection, with_row_id=self._with_row_id, ) if blob_auto_row_id: blob_paths = tuple( - blob_v2_projection_sources(schema, req.select).keys() + blob_v2_projection_sources(schema, projection).keys() ) self._blob_auto_row_id = blob_auto_row_id self._blob_paths = blob_paths diff --git a/python/python/lancedb/remote/table.py b/python/python/lancedb/remote/table.py index d9139396b..02748b9bc 100644 --- a/python/python/lancedb/remote/table.py +++ b/python/python/lancedb/remote/table.py @@ -36,6 +36,7 @@ from lancedb._lancedb import ( UpdateResult, ) from lancedb.embeddings.base import EmbeddingFunctionConfig +from lancedb.expr import Expr from lancedb.index import ( FTS, BTree, @@ -863,7 +864,7 @@ class RemoteTable(Table): def update( self, - where: Optional[str] = None, + where: Optional[Union[str, Expr]] = None, values: Optional[dict] = None, *, values_sql: Optional[Dict[str, str]] = None, @@ -874,9 +875,11 @@ class RemoteTable(Table): Parameters ---------- - where: str, optional - The SQL where clause to use when updating rows. For example, 'x = 2' - or 'x IN (1, 2, 3)'. The filter must not be empty, or it will error. + where: str or [Expr][lancedb.expr.Expr], optional + The filter condition. Can be a SQL string or a type-safe + [Expr][lancedb.expr.Expr] built with [col][lancedb.expr.col] and + [lit][lancedb.expr.lit]. The filter must not be empty, or it will + error. values: dict, optional The values to update. The keys are the column names and the values are the values to set. diff --git a/python/python/lancedb/schema.py b/python/python/lancedb/schema.py index 33adbbae3..4dce09f0f 100644 --- a/python/python/lancedb/schema.py +++ b/python/python/lancedb/schema.py @@ -4,30 +4,34 @@ """Schema helpers for Lance blob columns.""" +import importlib +from typing import TYPE_CHECKING + import pyarrow as pa +import pyarrow.ipc + +if TYPE_CHECKING: + from lance.blob import BlobType as BlobType _BLOB_EXTENSION_NAME = "lance.blob.v2" _BLOB_V1_KEY = "lance-encoding:blob" _ARROW_EXT_NAME_KEY = "ARROW:extension:name" +_BLOB_V2_STORAGE_TYPE = pa.struct( + [ + pa.field("data", pa.large_binary(), nullable=True), + pa.field("uri", pa.utf8(), nullable=True), + pa.field("position", pa.uint64(), nullable=True), + pa.field("size", pa.uint64(), nullable=True), + ] +) +_resolved_blob_type = None -class BlobType(pa.ExtensionType): - """PyArrow extension type for a Lance blob v2 column. - - Queries return descriptors; call :meth:`~lancedb.table.Table.fetch_blob_files` - for lazy reads or :meth:`~lancedb.table.Table.fetch_blobs` for eager bytes. - """ +class _FallbackBlobType(pa.ExtensionType): + """lance.blob.v2 extension type used when pylance is not installed.""" def __init__(self) -> None: - storage_type = pa.struct( - [ - pa.field("data", pa.large_binary(), nullable=True), - pa.field("uri", pa.utf8(), nullable=True), - pa.field("position", pa.uint64(), nullable=True), - pa.field("size", pa.uint64(), nullable=True), - ] - ) - super().__init__(storage_type, _BLOB_EXTENSION_NAME) + pa.ExtensionType.__init__(self, _BLOB_V2_STORAGE_TYPE, _BLOB_EXTENSION_NAME) def __arrow_ext_serialize__(self) -> bytes: return b"" @@ -35,23 +39,16 @@ class BlobType(pa.ExtensionType): @classmethod def __arrow_ext_deserialize__( cls, storage_type: pa.DataType, serialized: bytes - ) -> "BlobType": + ) -> "_FallbackBlobType": return cls() def __reduce__(self): - # Ensure pickle round-trips on older pyarrow (apache/arrow#35599). return type(self).__arrow_ext_deserialize__, ( self.storage_type, self.__arrow_ext_serialize__(), ) -try: - pa.register_extension_type(BlobType()) # type: ignore[arg-type] -except pa.ArrowKeyError: - pass - - def _metadata_value(metadata: dict, key: str): return metadata.get(key.encode()) or metadata.get(key) @@ -92,43 +89,105 @@ def is_blob_like_field(field: pa.Field) -> bool: return is_blob_v2_field(field) or _metadata_marks_legacy_blob(field.metadata or {}) -def _collect_blob_paths(schema: pa.Schema, is_blob) -> list[str]: - paths: list[str] = [] +def _collect_blob_paths(schema: pa.Schema, is_blob) -> list[tuple[str, bool]]: + """Walk the schema and return (path, has_list_ancestor) for each blob field.""" + paths: list[tuple[str, bool]] = [] - def walk(fields, prefix: str) -> None: + def walk(fields, prefix: str, has_list_ancestor: bool) -> None: for field in fields: path = f"{prefix}.{field.name}" if prefix else field.name if is_blob(field): - paths.append(path) + paths.append((path, has_list_ancestor)) elif pa.types.is_struct(field.type): - walk(field.type, path) + walk(field.type, path, has_list_ancestor) elif ( pa.types.is_list(field.type) or pa.types.is_large_list(field.type) or pa.types.is_fixed_size_list(field.type) ): - walk([field.type.value_field], path) + walk([field.type.value_field], path, True) - walk(schema, "") + walk(schema, "", False) return paths def blob_column_paths(schema: pa.Schema) -> list[str]: """Dotted paths of blob-like columns (v2 extension or legacy metadata).""" - return _collect_blob_paths(schema, is_blob_like_field) + return [path for path, _ in _collect_blob_paths(schema, is_blob_like_field)] def blob_v2_column_paths(schema: pa.Schema) -> list[str]: - return _collect_blob_paths(schema, is_blob_v2_field) + return [path for path, _ in _collect_blob_paths(schema, is_blob_v2_field)] + + +def row_addressable_blob_v2_paths(schema: pa.Schema) -> list[str]: + """Blob v2 paths with one blob addressable by table row id. + + ``fetch_blobs`` and the descriptor row-id ride-along address one blob per + row, so a blob inside a list container has no row-id slot and no fetch + path. Those columns still store and query as raw descriptors. + """ + return [ + path + for path, has_list_ancestor in _collect_blob_paths(schema, is_blob_v2_field) + if not has_list_ancestor + ] def schema_has_blob_field(schema: pa.Schema) -> bool: return bool(blob_column_paths(schema)) +def _deserialize_registered_type(extension_type: pa.ExtensionType) -> pa.DataType: + """Return the type Arrow reconstructs for this extension name.""" + schema = pa.schema([pa.field("value", extension_type)]) + restored = pa.ipc.read_schema(schema.serialize()) + return restored.field("value").type + + +def _resolve_blob_type(): + """Return the BlobType class this process should use. + + pylance's class when it owns the lance.blob.v2 registry entry, + otherwise LanceDB's fallback. A different registered class is an error. + """ + global _resolved_blob_type + if _resolved_blob_type is not None: + return _resolved_blob_type + try: + blob_module = importlib.import_module("lance.blob") + except ModuleNotFoundError as err: + if err.name not in ("lance", "lance.blob"): + raise + else: + blob_type = getattr(blob_module, "BlobType", None) + if blob_type is not None: + registered_type = _deserialize_registered_type(blob_type()) + if type(registered_type) is not blob_type: + registered_cls = type(registered_type) + raise ValueError( + "lance.blob.v2 is already registered by " + f"{registered_cls.__module__}.{registered_cls.__qualname__}" + ) + _resolved_blob_type = blob_type + return blob_type + try: + pa.register_extension_type(_FallbackBlobType()) # type: ignore[arg-type] + except pa.ArrowKeyError as err: + raise ValueError( + "lance.blob.v2 is already registered by another extension class" + ) from err + _resolved_blob_type = _FallbackBlobType + return _resolved_blob_type + + def blob(name: str, nullable: bool = True) -> pa.Field: - """Create a Lance blob v2 column field.""" - return pa.field(name, BlobType(), nullable=nullable) + """Create a Lance blob v2 column field. + + When pylance is installed this is ``lance.blob.BlobType``. + """ + blob_type = _resolve_blob_type() + return pa.field(name, blob_type(), nullable=nullable) def vector(dimension: int, value_type: pa.DataType = pa.float32()) -> pa.DataType: @@ -155,3 +214,11 @@ def vector(dimension: int, value_type: pa.DataType = pa.float32()) -> pa.DataTyp ... ]) """ return pa.list_(value_type, dimension) + + +def __getattr__(name: str): + if name == "BlobType": + blob_type = _resolve_blob_type() + globals()["BlobType"] = blob_type + return blob_type + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index 78a9e12a2..b96fc64f1 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -104,7 +104,12 @@ from .util import ( value_to_sql, ) from .index import lang_mapping -from .schema import blob_v2_column_paths, schema_has_blob_field +from .schema import ( + blob_v2_column_paths, + is_blob_v2_field, + row_addressable_blob_v2_paths, + schema_has_blob_field, +) def _should_push_down_query_table( @@ -426,6 +431,7 @@ def _cast_to_target_schema( def gen(): for batch in reader: + batch = _coerce_blob_write_columns(batch, reordered_schema) # Table but not RecordBatch has cast. cast_batches = ( pa.Table.from_batches([batch]).cast(reordered_schema).to_batches() @@ -438,6 +444,185 @@ def _cast_to_target_schema( return pa.RecordBatchReader.from_batches(reordered_schema, gen()) +def _coerce_blob_write_columns( + batch: pa.RecordBatch, target_schema: pa.Schema +) -> pa.RecordBatch: + """Materialize blob storage structs before the stream leaves Python. + + merge_insert requires its source reader to already match the table's + physical schema. Unlike add and insert, it does not pass through + LanceDB's Rust blob coercion, so preserving binary input here would + reach Lance as binary and fail the schema check. + """ + columns = [] + fields = [] + changed = False + for field, column in zip(batch.schema, batch.columns): + target_field = target_schema.field(field.name) + coerced = _coerce_blob_value(column, target_field) + if coerced is not column: + column = coerced + field = pa.field( + field.name, + coerced.type, + field.nullable, + target_field.metadata, + ) + changed = True + columns.append(column) + fields.append(field) + if not changed: + return batch + return pa.RecordBatch.from_arrays( + columns, schema=pa.schema(fields, metadata=batch.schema.metadata) + ) + + +def _coerce_blob_value(column: pa.Array, target_field: pa.Field) -> pa.Array: + if is_blob_v2_field(target_field) and _can_coerce_to_blob(column.type): + return _coerce_value_to_blob(column, target_field) + + target_type = target_field.type + if pa.types.is_struct(target_type) and pa.types.is_struct(column.type): + children = [] + fields = [] + changed = False + for source_field in column.type: + source_column = column.field(source_field.name) + nested_target = next( + (field for field in target_type if field.name == source_field.name), + None, + ) + if nested_target is None: + children.append(source_column) + fields.append(source_field) + continue + coerced = _coerce_blob_value(source_column, nested_target) + if coerced is not source_column: + changed = True + child_array, child_type = _physical_array_and_type(coerced) + children.append(child_array) + fields.append( + pa.field( + source_field.name, + child_type, + source_field.nullable, + nested_target.metadata, + ) + ) + if not changed: + return column + return pa.StructArray.from_arrays( + children, + fields=fields, + mask=column.is_null() if column.null_count else None, + ) + + if _is_list_like(target_type) and _is_list_like(column.type): + return _coerce_blob_list_values(column, target_type.value_field) + + return column + + +def _coerce_blob_list_values( + column: pa.Array, target_value_field: pa.Field +) -> pa.Array: + """Coerce blob values inside a list column, preserving offsets and nulls. + + Works on the raw child values window instead of ``pc.list_flatten`` because + flatten drops values spanned by null slots, which would misalign offsets. + """ + mask = column.is_null() if column.null_count else None + if pa.types.is_fixed_size_list(column.type): + list_size = column.type.list_size + values = column.values.slice(column.offset * list_size, len(column) * list_size) + coerced = _coerce_blob_value(values, target_value_field) + if coerced is values: + return column + physical_values, _ = _physical_array_and_type(coerced) + return pa.FixedSizeListArray.from_arrays(physical_values, list_size, mask=mask) + offsets = column.offsets + first_offset = offsets[0].as_py() + values = column.values.slice( + first_offset, + offsets[-1].as_py() - first_offset, + ) + coerced = _coerce_blob_value(values, target_value_field) + if coerced is values: + return column + physical_values, _ = _physical_array_and_type(coerced) + if first_offset: + offsets = pc.subtract(offsets, pa.scalar(first_offset, offsets.type)) + if pa.types.is_large_list(column.type): + return pa.LargeListArray.from_arrays(offsets, physical_values, mask=mask) + return pa.ListArray.from_arrays(offsets, physical_values, mask=mask) + + +def _coerce_value_to_blob(values: pa.Array, target_field: pa.Field) -> pa.Array: + if _is_string_like(values.type): + carrier_name = "uri" + carrier = values + elif pa.types.is_null(values.type): + carrier_name = None + carrier = None + elif pa.types.is_large_binary(values.type): + carrier_name = "data" + carrier = values + else: + carrier_name = "data" + carrier = values.cast(pa.large_binary()) + length = len(values) + storage_type = target_field.type + if isinstance(storage_type, pa.ExtensionType): + storage_type = storage_type.storage_type + storage_fields = list(storage_type) + children = [] + for storage_field in storage_fields: + if storage_field.name == carrier_name: + children.append(carrier.cast(storage_field.type)) + else: + children.append(pa.nulls(length, type=storage_field.type)) + storage = pa.StructArray.from_arrays( + children, + fields=storage_fields, + mask=values.is_null() if values.null_count else None, + ) + if isinstance(target_field.type, pa.ExtensionType): + return pa.ExtensionArray.from_storage(target_field.type, storage) + return storage + + +def _physical_array_and_type(array: pa.Array) -> tuple[pa.Array, pa.DataType]: + if isinstance(array.type, pa.ExtensionType): + return array.storage, array.type.storage_type + return array, array.type + + +def _can_coerce_to_blob(data_type: pa.DataType) -> bool: + return ( + _is_binary_like(data_type) + or _is_string_like(data_type) + or pa.types.is_null(data_type) + ) + + +def _is_binary_like(data_type: pa.DataType) -> bool: + return ( + pa.types.is_binary(data_type) + or pa.types.is_large_binary(data_type) + or pa.types.is_binary_view(data_type) + ) + + +def _is_string_like(data_type: pa.DataType) -> bool: + predicates = ("is_string", "is_large_string", "is_string_view") + return any( + predicate(data_type) + for name in predicates + if (predicate := getattr(pa.types, name, None)) is not None + ) + + def _field_extension_name(field: pa.Field) -> Optional[str]: extension_name = getattr(field.type, "extension_name", None) if extension_name is not None: @@ -633,29 +818,6 @@ def _prepare_extension_list(data: DATA, target_schema: pa.Schema) -> DATA: return pa.Table.from_pylist(prepared_data, schema=insert_schema) -def _is_blob_source_field(field: pa.Field) -> bool: - if _field_extension_name(field) == _BLOB_EXTENSION_NAME: - return True - - predicates = ( - "is_binary", - "is_large_binary", - "is_binary_view", - "is_string", - "is_large_string", - "is_string_view", - ) - if any( - predicate(field.type) - for name in predicates - if (predicate := getattr(pa.types, name, None)) is not None - ): - return True - return pa.types.is_struct(field.type) and any( - child.name in {"data", "uri"} for child in field.type - ) - - def _align_field_types( fields: List[pa.Field], target_fields: List[pa.Field], @@ -668,71 +830,71 @@ def _align_field_types( target_field = next((f for f in target_fields if f.name == field.name), None) if target_field is None: raise ValueError(f"Field '{field.name}' not found in target schema") - target_extension_name = _field_extension_name(target_field) - # Preserve accepted blob carriers so Lance can construct the declared - # blob struct after optional Python preprocessing. - if target_extension_name == _BLOB_EXTENSION_NAME and _is_blob_source_field( - field - ): - new_fields.append(field) - continue - # Preserve arrow.json input until it reaches Lance. LanceDB exposes stored - # JSON columns as lance.json (JSONB-backed LargeBinary), but casting the - # input to that storage type here merely relabels the raw JSON bytes as - # JSONB. Lance must see arrow.json so it can perform the JSONB encoding. - if ( - _field_extension_name(field) == "arrow.json" - and target_extension_name in _JSON_EXTENSION_NAMES - ): - new_fields.append(field) - continue - if pa.types.is_struct(target_field.type): - if pa.types.is_struct(field.type): - new_type = pa.struct( - _align_field_types( - field.type.fields, - target_field.type.fields, - ) + new_fields.append(_align_field(field, target_field)) + return new_fields + + +def _align_list_value_field( + value_field: pa.Field, target_value_field: pa.Field +) -> pa.Field: + # A list has exactly one child, so the inferred child name ("item") aligns + # positionally and adopts the table's child name; pa.Table.cast renames it. + return _align_field(value_field, target_value_field).with_name( + target_value_field.name + ) + + +def _align_field(field: pa.Field, target_field: pa.Field) -> pa.Field: + # Preserve arrow.json input until it reaches Lance. LanceDB exposes stored + # JSON columns as lance.json (JSONB-backed LargeBinary), but casting the + # input to that storage type here merely relabels the raw JSON bytes as + # JSONB. Lance must see arrow.json so it can perform the JSONB encoding. + if ( + _field_extension_name(field) == "arrow.json" + and _field_extension_name(target_field) == "lance.json" + ): + return field + if pa.types.is_struct(target_field.type): + if pa.types.is_struct(field.type): + new_type = pa.struct( + _align_field_types( + field.type.fields, + target_field.type.fields, ) - else: - new_type = target_field.type - elif pa.types.is_list(target_field.type): - if _is_list_like(field.type): - new_type = pa.list_( - _align_field_types( - [field.type.value_field], - [target_field.type.value_field], - )[0] - ) - else: - new_type = target_field.type - elif pa.types.is_large_list(target_field.type): - if _is_list_like(field.type): - new_type = pa.large_list( - _align_field_types( - [field.type.value_field], - [target_field.type.value_field], - )[0] - ) - else: - new_type = target_field.type - elif pa.types.is_fixed_size_list(target_field.type): - if _is_list_like(field.type): - new_type = pa.list_( - _align_field_types( - [field.type.value_field], - [target_field.type.value_field], - )[0], - target_field.type.list_size, - ) - else: - new_type = target_field.type + ) else: new_type = target_field.type - new_fields.append( - pa.field(field.name, new_type, field.nullable, target_field.metadata) - ) - return new_fields + elif pa.types.is_list(target_field.type): + if _is_list_like(field.type): + new_type = pa.list_( + _align_list_value_field( + field.type.value_field, target_field.type.value_field + ) + ) + else: + new_type = target_field.type + elif pa.types.is_large_list(target_field.type): + if _is_list_like(field.type): + new_type = pa.large_list( + _align_list_value_field( + field.type.value_field, target_field.type.value_field + ) + ) + else: + new_type = target_field.type + elif pa.types.is_fixed_size_list(target_field.type): + if _is_list_like(field.type): + new_type = pa.list_( + _align_list_value_field( + field.type.value_field, target_field.type.value_field + ), + target_field.type.list_size, + ) + else: + new_type = target_field.type + else: + new_type = target_field.type + return pa.field(field.name, new_type, field.nullable, target_field.metadata) def _infer_subschema( @@ -801,7 +963,7 @@ def sanitize_create_table( schema = data.schema else: if schema is not None: - data = pa.Table.from_pylist([], schema) + data = pa.Table.from_batches([], schema=schema) if schema is None: if data is None: raise ValueError("Either data or schema must be provided") @@ -1956,7 +2118,7 @@ class Table(ABC): @abstractmethod def update( self, - where: Optional[str] = None, + where: Optional[Union[str, Expr]] = None, values: Optional[dict] = None, *, values_sql: Optional[Dict[str, str]] = None, @@ -1971,9 +2133,11 @@ class Table(ABC): Parameters ---------- - where: str, optional - The SQL where clause to use when updating rows. For example, 'x = 2' - or 'x IN (1, 2, 3)'. The filter must not be empty, or it will error. + where: str or [Expr][lancedb.expr.Expr], optional + The filter condition. Can be a SQL string or a type-safe + [Expr][lancedb.expr.Expr] built with [col][lancedb.expr.col] and + [lit][lancedb.expr.lit]. The filter must not be empty, or it will + error. values: dict, optional The values to update. The keys are the column names and the values are the values to set. @@ -1991,6 +2155,7 @@ class Table(ABC): Examples -------- >>> import lancedb + >>> from lancedb.expr import col >>> import pandas as pd >>> data = pd.DataFrame({"x": [1, 2, 3], "vector": [[1.0, 2], [3, 4], [5, 6]]}) >>> db = lancedb.connect("./.lancedb") @@ -2000,7 +2165,7 @@ class Table(ABC): 0 1 [1.0, 2.0] 1 2 [3.0, 4.0] 2 3 [5.0, 6.0] - >>> table.update(where="x = 2", values={"vector": [10.0, 10]}) + >>> table.update(where=col("x") == 2, values={"vector": [10.0, 10]}) UpdateResult(rows_updated=1, version=2) >>> table.to_pandas() x vector @@ -2907,7 +3072,7 @@ class LanceTable(Table): arrow_tbl = self.to_arrow() if blob_mode == "descriptions": arrow_tbl = strip_auto_row_ids( - arrow_tbl, blob_v2_column_paths(self.schema) + arrow_tbl, row_addressable_blob_v2_paths(self.schema) ) return arrow_tbl.to_pandas(**kwargs) @@ -4053,7 +4218,7 @@ class LanceTable(Table): def update( self, - where: Optional[str] = None, + where: Optional[Union[str, Expr]] = None, values: Optional[dict] = None, *, values_sql: Optional[Dict[str, str]] = None, @@ -4064,9 +4229,11 @@ class LanceTable(Table): Parameters ---------- - where: str, optional - The SQL where clause to use when updating rows. For example, 'x = 2' - or 'x IN (1, 2, 3)'. The filter must not be empty, or it will error. + where: str or [Expr][lancedb.expr.Expr], optional + The filter condition. Can be a SQL string or a type-safe + [Expr][lancedb.expr.Expr] built with [col][lancedb.expr.col] and + [lit][lancedb.expr.lit]. The filter must not be empty, or it will + error. values: dict, optional The values to update. The keys are the column names and the values are the values to set. @@ -4084,6 +4251,7 @@ class LanceTable(Table): Examples -------- >>> import lancedb + >>> from lancedb.expr import col >>> import pandas as pd >>> data = pd.DataFrame({"x": [1, 2, 3], "vector": [[1.0, 2], [3, 4], [5, 6]]}) >>> db = lancedb.connect("./.lancedb") @@ -4093,7 +4261,7 @@ class LanceTable(Table): 0 1 [1.0, 2.0] 1 2 [3.0, 4.0] 2 3 [5.0, 6.0] - >>> table.update(where="x = 2", values={"vector": [10.0, 10]}) + >>> table.update(where=col("x") == 2, values={"vector": [10.0, 10]}) UpdateResult(rows_updated=1, version=2) >>> table.to_pandas() x vector @@ -5308,7 +5476,9 @@ class AsyncTable: if blob_mode == "descriptions" or not schema_has_blob_field(schema): arrow_tbl = await self.to_arrow() if blob_mode == "descriptions": - arrow_tbl = strip_auto_row_ids(arrow_tbl, blob_v2_column_paths(schema)) + arrow_tbl = strip_auto_row_ids( + arrow_tbl, row_addressable_blob_v2_paths(schema) + ) return arrow_tbl.to_pandas(**kwargs) if blob_mode == "lazy" and get_uri_scheme(await self.uri()) == "memory": @@ -6210,7 +6380,7 @@ class AsyncTable: self, updates: Optional[Dict[str, Any]] = None, *, - where: Optional[str] = None, + where: Optional[Union[str, Expr]] = None, updates_sql: Optional[Dict[str, str]] = None, ) -> UpdateResult: """ @@ -6225,9 +6395,11 @@ class AsyncTable: The updates to apply. The keys should be the name of the column to update. The values should be the new values to assign. This is required unless updates_sql is supplied. - where: str, optional - An SQL filter that controls which rows are updated. For example, 'x = 2' - or 'x IN (1, 2, 3)'. Only rows that satisfy this filter will be udpated. + where: str or [Expr][lancedb.expr.Expr], optional + The filter condition. Can be a SQL string or a type-safe + [Expr][lancedb.expr.Expr] built with [col][lancedb.expr.col] and + [lit][lancedb.expr.lit]. Only rows that satisfy this filter will + be updated. updates_sql: dict, optional The updates to apply, expressed as SQL expression strings. The keys should be column names. The values should be SQL expressions. These can be SQL @@ -6245,13 +6417,14 @@ class AsyncTable: -------- >>> import asyncio >>> import lancedb + >>> from lancedb.expr import col >>> import pandas as pd >>> async def demo_update(): ... data = pd.DataFrame({"x": [1, 2], "vector": [[1, 2], [3, 4]]}) ... db = await lancedb.connect_async("./.lancedb") ... table = await db.create_table("my_table", data) ... # x is [1, 2], vector is [[1, 2], [3, 4]] - ... await table.update({"vector": [10, 10]}, where="x = 2") + ... await table.update({"vector": [10, 10]}, where=col("x") == 2) ... # x is [1, 2], vector is [[1, 2], [10, 10]] ... await table.update(updates_sql={"x": "x + 1"}) ... # x is [2, 3], vector is [[1, 2], [10, 10]] @@ -6265,7 +6438,8 @@ class AsyncTable: if updates is not None: updates_sql = {k: value_to_sql(v) for k, v in updates.items()} - return await self._inner.update(updates_sql, where) + predicate = where.to_sql() if isinstance(where, Expr) else where + return await self._inner.update(updates_sql, predicate) async def add_columns( self, diff --git a/python/python/tests/test_blob.py b/python/python/tests/test_blob.py index edff6d255..13c6f5531 100644 --- a/python/python/tests/test_blob.py +++ b/python/python/tests/test_blob.py @@ -2,17 +2,41 @@ # SPDX-FileCopyrightText: Copyright The LanceDB Authors import io +import subprocess +import sys +import textwrap +import lance import pyarrow as pa import pyarrow.compute as pc import pytest +from lance.blob import BlobType as LanceBlobType import lancedb -from lancedb._blob import read_row_ids_from_hits, stash_auto_row_ids +from lancedb._blob import ( + blob_v2_projection_sources, + read_row_ids_from_hits, + stash_auto_row_ids, +) +from lancedb.expr import col from lancedb.index import FTS from lancedb.schema import blob_column_paths, blob_v2_column_paths +_HIDE_LANCE_BLOB = """\ +import importlib.abc +import sys + +class _MissingLanceBlob(importlib.abc.MetaPathFinder): + def find_spec(self, fullname, path, target=None): + if fullname == "lance.blob" or fullname.startswith("lance.blob."): + raise ModuleNotFoundError(fullname, name="lance.blob") + +sys.modules.pop("lance.blob", None) +sys.meta_path.insert(0, _MissingLanceBlob()) +""" + + def _blob_table(name, rows): db = lancedb.connect("memory:///") schema = pa.schema([pa.field("id", pa.int64()), lancedb.blob("image")]) @@ -46,6 +70,181 @@ def test_blob_factory_declares_v2_field(): field = lancedb.blob("image") assert isinstance(field.type, pa.ExtensionType) assert field.type.extension_name == "lance.blob.v2" + assert lancedb.BlobType is LanceBlobType + assert type(field.type) is LanceBlobType + + +def test_blob_type_works_without_pylance(): + script = _HIDE_LANCE_BLOB + textwrap.dedent( + """\ + import lancedb + import pyarrow as pa + + field = lancedb.blob("image") + if not isinstance(field.type, pa.ExtensionType): + raise SystemExit("expected an extension type") + if field.type.extension_name != "lance.blob.v2": + raise SystemExit(field.type.extension_name) + if lancedb.BlobType is not type(field.type): + raise SystemExit("BlobType is not the field type class") + if lancedb.BlobType.__module__ != "lancedb.schema": + raise SystemExit(lancedb.BlobType.__module__) + + db = lancedb.connect("memory:///") + table = db.create_table( + "images", + schema=pa.schema([pa.field("id", pa.int64()), field]), + ) + table.add([{"id": 1, "image": b"hello"}]) + result = ( + table.merge_insert("id") + .when_matched_update_all() + .when_not_matched_insert_all() + .execute([{"id": 1, "image": b"updated"}, {"id": 2, "image": b"inserted"}]) + ) + if result.num_updated_rows != 1 or result.num_inserted_rows != 1: + raise SystemExit( + f"merge_insert rows updated={result.num_updated_rows} " + f"inserted={result.num_inserted_rows}" + ) + """ + ) + result = subprocess.run( + [sys.executable, "-c", script], + capture_output=True, + text=True, + check=False, + ) + assert result.returncode == 0, result.stderr + + +def test_blob_resolves_pylance_type_without_eager_import(): + script = textwrap.dedent( + """\ + import sys + import lancedb + + if "lance.blob" in sys.modules: + raise SystemExit("import lancedb imported lance.blob") + field = lancedb.blob("image") + from lance.blob import BlobType + + if type(field.type) is not BlobType: + raise SystemExit(f"{type(field.type)} is not {BlobType}") + import lance + + image = lance.blob_array([b"x"]) + if type(image.type) is not BlobType: + raise SystemExit("blob_array used a different class") + if type(image.type) is not type(field.type): + raise SystemExit("field and array classes differ") + """ + ) + result = subprocess.run( + [sys.executable, "-c", script], + capture_output=True, + text=True, + check=False, + ) + assert result.returncode == 0, result.stderr + + +def test_blob_fallback_fails_if_name_already_registered(): + script = _HIDE_LANCE_BLOB + textwrap.dedent( + """\ + import pyarrow as pa + + class OtherBlobType(pa.ExtensionType): + def __init__(self): + super().__init__( + pa.struct([pa.field("data", pa.large_binary())]), + "lance.blob.v2", + ) + + def __arrow_ext_serialize__(self): + return b"" + + @classmethod + def __arrow_ext_deserialize__(cls, storage_type, serialized): + return cls() + + pa.register_extension_type(OtherBlobType()) + import lancedb + + try: + lancedb.blob("image") + except ValueError as err: + if "already registered" not in str(err): + raise SystemExit(err) + else: + raise SystemExit("expected ValueError") + """ + ) + result = subprocess.run( + [sys.executable, "-c", script], + capture_output=True, + text=True, + check=False, + ) + assert result.returncode == 0, result.stderr + + +def test_blob_type_rejects_competing_registration_with_pylance(): + script = textwrap.dedent( + """\ + import pyarrow as pa + import pyarrow.ipc + + class OtherBlobType(pa.ExtensionType): + def __init__(self): + super().__init__( + pa.struct( + [ + pa.field("data", pa.large_binary()), + pa.field("uri", pa.utf8()), + pa.field("position", pa.uint64()), + pa.field("size", pa.uint64()), + ] + ), + "lance.blob.v2", + ) + + def __arrow_ext_serialize__(self): + return b"" + + @classmethod + def __arrow_ext_deserialize__(cls, storage_type, serialized): + return cls() + + pa.register_extension_type(OtherBlobType()) + + from lance.blob import BlobType + + if BlobType is OtherBlobType: + raise SystemExit("pylance BlobType was replaced") + schema = pa.schema([pa.field("value", BlobType())]) + restored = pa.ipc.read_schema(schema.serialize()) + if type(restored.field("value").type) is not OtherBlobType: + raise SystemExit(type(restored.field("value").type)) + + import lancedb + + try: + lancedb.blob("image") + except ValueError as err: + if "__main__.OtherBlobType" not in str(err): + raise SystemExit(err) + else: + raise SystemExit("expected ValueError") + """ + ) + result = subprocess.run( + [sys.executable, "-c", script], + capture_output=True, + text=True, + check=False, + ) + assert result.returncode == 0, result.stderr def test_blob_v2_column_paths_include_list_children(): @@ -70,6 +269,14 @@ def test_blob_v2_column_paths_include_list_children(): ] +def test_blob_v2_projection_sources_use_typed_column_name(): + schema = pa.schema([lancedb.blob("blob")]) + + assert blob_v2_projection_sources(schema, {"blob_alias": col("blob")}) == { + "blob_alias": "blob" + } + + def _legacy_v1_table(name): db = lancedb.connect("memory:///") schema = pa.schema( @@ -166,6 +373,20 @@ async def test_async_table_to_pandas_descriptions_mode_omits_row_id(): assert set(descriptor.keys()) == {"kind", "position", "size", "blob_id", "blob_uri"} +@pytest.mark.asyncio +async def test_async_typed_blob_projection_preserves_source_column(): + db = await lancedb.connect_async("memory:///typed_blob_projection") + schema = pa.schema([pa.field("id", pa.int64()), lancedb.blob("blob")]) + table = await db.create_table("typed_blob_projection", schema=schema) + await table.add([{"id": 1, "blob": b"alpha"}]) + + hits = await table.query().select({"blob_alias": col("blob")}).to_arrow() + + assert "_lance_row_id" in hits.schema.field("blob_alias").type.names + blobs = await table.fetch_blobs("blob", hits) + assert blobs.to_pylist() == [b"alpha"] + + def test_fetch_blobs_round_trip(): table = _blob_table( "round_trip", @@ -176,6 +397,292 @@ def test_fetch_blobs_round_trip(): assert [blobs[0].as_py(), blobs[1].as_py()] == [b"alpha", b"beta"] +def test_merge_insert_writes_python_bytes(): + table = _blob_table("merge_bytes", [{"id": 1, "image": b"before"}]) + result = ( + table.merge_insert("id") + .when_matched_update_all() + .when_not_matched_insert_all() + .execute([{"id": 1, "image": b"updated"}, {"id": 2, "image": b"inserted"}]) + ) + assert result.num_updated_rows == 1 + assert result.num_inserted_rows == 1 + by_id = _row_ids_by_id(table) + blobs = table.fetch_blobs("image", [by_id[1], by_id[2]]) + assert blobs.to_pylist() == [b"updated", b"inserted"] + + +def test_merge_insert_bytes_after_reopen_without_touching_blob_type(tmp_path): + db = lancedb.connect(tmp_path) + schema = pa.schema([pa.field("id", pa.int64()), lancedb.blob("image")]) + table = db.create_table("images", schema=schema) + table.add([{"id": 1, "image": b"hello"}]) + + script = textwrap.dedent( + f"""\ + import lancedb + + db = lancedb.connect({str(tmp_path)!r}) + table = db.open_table("images") + image_type = table.schema.field("image").type + if type(image_type).__name__ != "StructType": + raise SystemExit(f"expected StructType, got {{type(image_type)}}") + result = ( + table.merge_insert("id") + .when_matched_update_all() + .when_not_matched_insert_all() + .execute( + [{{"id": 1, "image": b"updated"}}, {{"id": 2, "image": b"inserted"}}] + ) + ) + if result.num_updated_rows != 1 or result.num_inserted_rows != 1: + raise SystemExit( + f"rows updated={{result.num_updated_rows}} " + f"inserted={{result.num_inserted_rows}}" + ) + hits = table.search().with_row_id(True).limit(10).to_arrow() + by_id = dict(zip(hits["id"].to_pylist(), hits["_rowid"].to_pylist())) + blobs = table.fetch_blobs("image", [by_id[1], by_id[2]]) + if blobs.to_pylist() != [b"updated", b"inserted"]: + raise SystemExit(blobs.to_pylist()) + """ + ) + result = subprocess.run( + [sys.executable, "-c", script], + capture_output=True, + text=True, + check=False, + ) + assert result.returncode == 0, result.stderr + + +def test_merge_insert_bytes_after_reopen_without_pylance(tmp_path): + db = lancedb.connect(tmp_path) + schema = pa.schema([pa.field("id", pa.int64()), lancedb.blob("image")]) + table = db.create_table("images", schema=schema) + table.add([{"id": 1, "image": b"hello"}]) + + script = _HIDE_LANCE_BLOB + textwrap.dedent( + f"""\ + import lancedb + + db = lancedb.connect({str(tmp_path)!r}) + table = db.open_table("images") + image_type = table.schema.field("image").type + if type(image_type).__name__ != "StructType": + raise SystemExit(f"expected StructType, got {{type(image_type)}}") + result = ( + table.merge_insert("id") + .when_matched_update_all() + .when_not_matched_insert_all() + .execute( + [{{"id": 1, "image": b"updated"}}, {{"id": 2, "image": b"inserted"}}] + ) + ) + if result.num_updated_rows != 1 or result.num_inserted_rows != 1: + raise SystemExit( + f"rows updated={{result.num_updated_rows}} " + f"inserted={{result.num_inserted_rows}}" + ) + hits = table.search().with_row_id(True).limit(10).to_arrow() + by_id = dict(zip(hits["id"].to_pylist(), hits["_rowid"].to_pylist())) + blobs = table.fetch_blobs("image", [by_id[1], by_id[2]]) + if blobs.to_pylist() != [b"updated", b"inserted"]: + raise SystemExit(blobs.to_pylist()) + """ + ) + result = subprocess.run( + [sys.executable, "-c", script], + capture_output=True, + text=True, + check=False, + ) + assert result.returncode == 0, result.stderr + + +def test_merge_insert_blob_array_into_reopened_unregistered_table(tmp_path): + db = lancedb.connect(tmp_path) + schema = pa.schema([pa.field("id", pa.int64()), lancedb.blob("image")]) + table = db.create_table("images", schema=schema) + table.add([{"id": 1, "image": b"before"}]) + + script = textwrap.dedent( + f"""\ + import pyarrow as pa + import lancedb + + db = lancedb.connect({str(tmp_path)!r}) + table = db.open_table("images") + image_type = table.schema.field("image").type + if type(image_type).__name__ != "StructType": + raise SystemExit( + f"expected StructType before lance import, got {{type(image_type)}}" + ) + + import lance + + updates = pa.Table.from_arrays( + [ + pa.array([1, 2], type=pa.int64()), + lance.blob_array([b"updated", b"inserted"]), + ], + names=["id", "image"], + ) + result = ( + table.merge_insert("id") + .when_matched_update_all() + .when_not_matched_insert_all() + .execute(updates) + ) + if result.num_updated_rows != 1 or result.num_inserted_rows != 1: + raise SystemExit( + f"rows updated={{result.num_updated_rows}} " + f"inserted={{result.num_inserted_rows}}" + ) + hits = table.search().with_row_id(True).limit(10).to_arrow() + by_id = dict(zip(hits["id"].to_pylist(), hits["_rowid"].to_pylist())) + blobs = table.fetch_blobs("image", [by_id[1], by_id[2]]) + if blobs.to_pylist() != [b"updated", b"inserted"]: + raise SystemExit(blobs.to_pylist()) + """ + ) + result = subprocess.run( + [sys.executable, "-c", script], + capture_output=True, + text=True, + check=False, + ) + assert result.returncode == 0, result.stderr + + +def test_add_all_null_blob_column(): + db = lancedb.connect("memory:///") + schema = pa.schema([pa.field("id", pa.int64()), lancedb.blob("image")]) + table = db.create_table("all_null", schema=schema) + table.add([{"id": 1, "image": None}, {"id": 2, "image": None}]) + by_id = _row_ids_by_id(table) + blobs = table.fetch_blobs("image", [by_id[1], by_id[2]]) + assert blobs.to_pylist() == [None, None] + + +def test_create_table_nested_blob_schema_without_rows(): + db = lancedb.connect("memory:///") + schema = pa.schema( + [ + pa.field("id", pa.int64()), + pa.field("info", pa.struct([lancedb.blob("blob")])), + pa.field("images", pa.list_(lancedb.blob("image"))), + ] + ) + table = db.create_table("nested_empty", schema=schema) + assert table.count_rows() == 0 + + +def test_merge_insert_nested_blob_dicts(): + db = lancedb.connect("memory:///") + info = pa.StructArray.from_arrays( + [ + pa.array(["first"], type=pa.string()), + _blob_array("blob", [b"before"]), + ], + names=["name", "blob"], + ) + data = pa.Table.from_arrays( + [pa.array([1], type=pa.int64()), info], + names=["id", "info"], + ) + table = db.create_table("nested_merge", data=data) + result = ( + table.merge_insert("id") + .when_matched_update_all() + .execute([{"id": 1, "info": {"name": "first", "blob": b"after"}}]) + ) + assert result.num_updated_rows == 1 + by_id = _row_ids_by_id(table) + blobs = table.fetch_blobs("info.blob", [by_id[1]]) + assert blobs.to_pylist() == [b"after"] + + +def _list_blob_table(name): + db = lancedb.connect("memory:///") + blob_field = lancedb.blob("image") + images = pa.ListArray.from_arrays( + pa.array([0, 1], type=pa.int32()), _blob_array("image", [b"before"]) + ) + data = pa.Table.from_arrays( + [pa.array([1], type=pa.int64()), images], + schema=pa.schema( + [pa.field("id", pa.int64()), pa.field("images", pa.list_(blob_field))] + ), + ) + return db.create_table(name, data=data) + + +def test_merge_insert_list_blob_dicts(): + table = _list_blob_table("list_merge") + result = ( + table.merge_insert("id") + .when_matched_update_all() + .when_not_matched_insert_all() + .execute([{"id": 1, "images": [b"one", b"two"]}, {"id": 2, "images": None}]) + ) + assert result.num_updated_rows == 1 + assert result.num_inserted_rows == 1 + hits = table.search().limit(10).to_arrow() + sizes = { + row["id"]: None if row["images"] is None else [d["size"] for d in row["images"]] + for row in hits.to_pylist() + } + assert sizes == {1: [3, 3], 2: None} + + +def test_list_blob_column_queries_as_raw_descriptors(): + table = _list_blob_table("list_query") + hits = table.search().limit(10).to_arrow() + element = hits.schema.field("images").type.value_type + assert pa.types.is_struct(element) + assert "_lance_row_id" not in element.names + with pytest.raises(ValueError, match="expected struct before segment"): + table.fetch_blobs("images.image", [0]) + + +def test_row_addressable_paths_exclude_list_children(): + from lancedb.schema import row_addressable_blob_v2_paths + + schema = pa.schema( + [ + pa.field("id", pa.int64()), + pa.field("info", pa.struct([lancedb.blob("blob")])), + pa.field("images", pa.list_(lancedb.blob("image"))), + ] + ) + assert blob_v2_column_paths(schema) == ["info.blob", "images.image"] + assert row_addressable_blob_v2_paths(schema) == ["info.blob"] + + +def test_merge_insert_writes_pylance_blob_array(): + table = _blob_table("merge_pylance", [{"id": 1, "image": b"before"}]) + image = lance.blob_array([b"updated", b"inserted"]) + assert type(image.type) is LanceBlobType + assert type(image.type) is type(lancedb.BlobType()) + updates = pa.Table.from_arrays( + [pa.array([1, 2], type=pa.int64()), image], names=["id", "image"] + ) + + result = ( + table.merge_insert("id") + .when_matched_update_all() + .when_not_matched_insert_all() + .execute(updates) + ) + + assert result.num_updated_rows == 1 + assert result.num_inserted_rows == 1 + by_id = _row_ids_by_id(table) + blobs = table.fetch_blobs("image", [by_id[1], by_id[2]]) + assert blobs.to_pylist() == [b"updated", b"inserted"] + + def test_fetch_blobs_accepts_query_result(): table = _blob_table("from_result", [{"id": 1, "image": b"gamma"}]) hits = table.search().limit(10).to_arrow() @@ -477,6 +984,50 @@ async def test_blob_v2_hybrid_fetch_blobs_async(): assert {blobs[i].as_py() for i in range(len(blobs))} == {b"alpha", b"beta"} +@pytest.mark.asyncio +async def test_async_hybrid_typed_blob_projection_preserves_source_column(): + db = await lancedb.connect_async("memory:///hybrid_typed_blob") + schema = pa.schema( + [ + pa.field("id", pa.int64()), + pa.field("text", pa.utf8()), + pa.field("vector", pa.list_(pa.float32(), list_size=2)), + lancedb.blob("blob"), + ] + ) + table = await db.create_table("hybrid_typed_blob", schema=schema) + await table.add( + [ + { + "id": 1, + "text": "hello alpha", + "vector": [1.0, 0.0], + "blob": b"alpha", + }, + { + "id": 2, + "text": "hello beta", + "vector": [0.9, 0.1], + "blob": b"beta", + }, + ] + ) + await table.create_index("text", config=FTS(with_position=False)) + + hits = await ( + table.query() + .nearest_to([1.0, 0.0]) + .nearest_to_text("hello") + .select({"blob_alias": col("blob")}) + .limit(2) + .to_arrow() + ) + + assert "_lance_row_id" in hits.schema.field("blob_alias").type.names + blobs = await table.fetch_blobs("blob", hits) + assert {blobs[i].as_py() for i in range(len(blobs))} == {b"alpha", b"beta"} + + def test_blob_file_seek_read_and_read_range(): payload = _identifiable_payload(1024) table = _blob_table("seek_read", [{"id": 1, "image": payload}]) diff --git a/python/python/tests/test_expr.py b/python/python/tests/test_expr.py index 0eb6f8929..0f49231f1 100644 --- a/python/python/tests/test_expr.py +++ b/python/python/tests/test_expr.py @@ -52,7 +52,7 @@ class TestExprConstruction: def test_func(self): e = func("lower", col("name")) assert isinstance(e, Expr) - assert e.to_sql() == "lower(name)" + assert e.to_sql() == "lower(`name`)" def test_func_unknown_raises(self): with pytest.raises(Exception): @@ -115,7 +115,7 @@ class TestExprOperators: def test_and_operator(self): e = (col("age") > lit(18)) & (col("status") == lit("active")) assert isinstance(e, Expr) - assert e.to_sql() == "((age > 18) AND (status = 'active'))" + assert e.to_sql() == "((age > 18) AND (`status` = 'active'))" def test_or_operator(self): e = (col("a") == lit(1)) | (col("b") == lit(2)) @@ -166,7 +166,7 @@ class TestExprOperators: def test_coerce_plain_str(self): e = col("name") == "alice" assert isinstance(e, Expr) - assert e.to_sql() == "(name = 'alice')" + assert e.to_sql() == "(`name` = 'alice')" def test_reflexive_comparisons(self): # 10 < col("age") swaps to col("age") > 10 @@ -198,85 +198,85 @@ class TestExprBytesLiteral: def test_bytes_equality_expr_sql(self): e = col("data") == lit(b"\xca\xfe") - assert e.to_sql() == "(data = X'CAFE')" + assert e.to_sql() == "(`data` = X'CAFE')" def test_bytes_ne_expr_sql(self): e = col("data") != lit(b"\xff") - assert e.to_sql() == "(data <> X'FF')" + assert e.to_sql() == "(`data` <> X'FF')" def test_bytes_compound_expr_sql(self): e = (col("data") == lit(b"\x01")) & (col("id") > lit(5)) - assert e.to_sql() == "((data = X'01') AND (id > 5))" + assert e.to_sql() == "((`data` = X'01') AND (id > 5))" def test_bytes_in_function_call(self): # Regression test: binary literals inside scalar function calls # used to fail because DataFusion's unparser does not support Binary # scalars. Now handled via a placeholder-substitution rewrite. e = func("contains", col("data"), lit(b"\xff")) - assert e.to_sql() == "contains(data, X'FF')" + assert e.to_sql() == "contains(`data`, X'FF')" def test_bytes_in_not(self): e = ~(col("data") == lit(b"\xff")) - assert e.to_sql() == "NOT (data = X'FF')" + assert e.to_sql() == "NOT (`data` = X'FF')" class TestExprStringMethods: def test_lower(self): e = col("name").lower() assert isinstance(e, Expr) - assert e.to_sql() == "lower(name)" + assert e.to_sql() == "lower(`name`)" def test_upper(self): e = col("name").upper() assert isinstance(e, Expr) - assert e.to_sql() == "upper(name)" + assert e.to_sql() == "upper(`name`)" def test_contains(self): e = col("text").contains(lit("hello")) assert isinstance(e, Expr) - assert e.to_sql() == "contains(text, 'hello')" + assert e.to_sql() == "contains(`text`, 'hello')" def test_contains_with_str_coerce(self): e = col("text").contains("hello") assert isinstance(e, Expr) - assert e.to_sql() == "contains(text, 'hello')" + assert e.to_sql() == "contains(`text`, 'hello')" def test_chained_lower_eq(self): e = col("name").lower() == lit("alice") assert isinstance(e, Expr) - assert e.to_sql() == "(lower(name) = 'alice')" + assert e.to_sql() == "(lower(`name`) = 'alice')" class TestExprCast: def test_cast_string(self): e = col("id").cast("string") assert isinstance(e, Expr) - assert e.to_sql() == "CAST(id AS VARCHAR)" + assert e.to_sql() == "arrow_cast(id, 'Utf8')" def test_cast_int32(self): e = col("score").cast("int32") assert isinstance(e, Expr) - assert e.to_sql() == "CAST(score AS INTEGER)" + assert e.to_sql() == "arrow_cast(score, 'Int32')" def test_cast_float64(self): e = col("val").cast("float64") assert isinstance(e, Expr) - assert e.to_sql() == "CAST(val AS DOUBLE)" + assert e.to_sql() == "arrow_cast(val, 'Float64')" def test_cast_pyarrow_type(self): e = col("score").cast(pa.int32()) assert isinstance(e, Expr) - assert e.to_sql() == "CAST(score AS INTEGER)" + assert e.to_sql() == "arrow_cast(score, 'Int32')" def test_cast_pyarrow_float64(self): e = col("val").cast(pa.float64()) assert isinstance(e, Expr) - assert e.to_sql() == "CAST(val AS DOUBLE)" + assert e.to_sql() == "arrow_cast(val, 'Float64')" def test_cast_pyarrow_string(self): e = col("id").cast(pa.string()) assert isinstance(e, Expr) - assert e.to_sql() == "CAST(id AS VARCHAR)" + assert e.to_sql() == "arrow_cast(id, 'Utf8')" def test_cast_pyarrow_and_string_equivalent(self): # pa.int32() and "int32" should produce equivalent SQL @@ -597,14 +597,14 @@ class TestExprIsin: def test_isin_strs(self): assert ( col("status").isin(["active", "pending"]).to_sql() - == "status IN ('active', 'pending')" + == "`status` IN ('active', 'pending')" ) def test_isin_coerces_and_mixes(self): assert col("id").isin([lit(1), 2]).to_sql() == "id IN (1, 2)" def test_isin_empty(self): - assert col("id").isin([]).to_sql() == "id IN ()" + assert col("id").isin([]).to_sql() == "false" def test_isin_filter(self, simple_table): result = simple_table.search().where(col("id").isin([1, 3, 5])).to_arrow() diff --git a/python/python/tests/test_query.py b/python/python/tests/test_query.py index 4758f0e2d..ff62b2b51 100644 --- a/python/python/tests/test_query.py +++ b/python/python/tests/test_query.py @@ -675,6 +675,21 @@ def test_distance_range(table: lancedb.table.Table): assert res["_distance"].to_pylist() == [min_dist, max_dist] +@pytest.mark.parametrize("expression", ["1 - _distance", "1.0 - _distance"]) +def test_select_arithmetic_with_distance(table, expression): + result = ( + table.search([10, 10]) + .select({"similarity": expression, "_distance": "_distance"}) + .distance_type("cosine") + .to_arrow() + ) + + assert result.schema.field("similarity").type == pa.float32() + assert result["similarity"].to_pylist() == pytest.approx( + [1 - distance for distance in result["_distance"].to_pylist()] + ) + + @pytest.mark.asyncio async def test_distance_range_async(table_async: AsyncTable): q = [0, 0] diff --git a/python/python/tests/test_table.py b/python/python/tests/test_table.py index f51434f2e..c4e768236 100644 --- a/python/python/tests/test_table.py +++ b/python/python/tests/test_table.py @@ -11,6 +11,7 @@ import warnings import weakref from concurrent.futures import ThreadPoolExecutor from datetime import date, datetime, timedelta +from decimal import Decimal from time import sleep from typing import List from unittest.mock import patch @@ -336,6 +337,21 @@ async def test_update_async(mem_db_async: AsyncConnection): assert await table.count_rows("id == 10") == 1 +@pytest.mark.asyncio +async def test_update_expr_filter_literals_async(mem_db_async: AsyncConnection): + values = ["5", "4.66e-84", "it's"] + table = await mem_db_async.create_table( + "update_expr_literals", + data=[{"field": value, "result": "original"} for value in values], + ) + + for value in values: + update_res = await table.update({"result": value}, where=col("field") == value) + assert update_res.rows_updated == 1 + + assert (await table.to_arrow())["result"].to_pylist() == values + + def test_create_table(mem_db: DBConnection): schema = pa.schema( { @@ -2392,6 +2408,148 @@ def test_update(mem_db: DBConnection): assert np.allclose(v, np.array([[1.2, 1.9], [1.1, 1.1]])) +def test_update_expr_filter_literals(mem_db: DBConnection): + values = ["5", "4.66e-84", "it's"] + table = mem_db.create_table( + "update_expr_literals", + data=[{"field": value, "result": "original"} for value in values], + ) + + for value in values: + update_res = table.update(where=col("field") == value, values={"result": value}) + assert update_res.rows_updated == 1 + + assert table.to_arrow()["result"].to_pylist() == values + + +def test_update_expr_filter_preserves_typed_semantics(mem_db: DBConnection): + low = Decimal("1.234567890123456789") + high = Decimal("1.234567890123456790") + decimal_schema = pa.schema( + [("val", pa.decimal128(19, 18)), ("result", pa.string())] + ) + decimal_table = mem_db.create_table( + "update_expr_decimal", + pa.table( + {"val": [low, high], "result": ["old", "old"]}, + schema=decimal_schema, + ), + ) + predicate = col("val") < lit(high) + assert decimal_table.search().where(predicate).to_arrow().num_rows == 1 + result = decimal_table.update(where=predicate, values={"result": "new"}) + assert result.rows_updated == 1 + + keyword_table = mem_db.create_table( + "update_expr_keyword", [{"null": 1, "result": "old"}] + ) + predicate = col("null") == 1 + assert keyword_table.search().where(predicate).to_arrow().num_rows == 1 + result = keyword_table.update(where=predicate, values={"result": "new"}) + assert result.rows_updated == 1 + + empty_in_table = mem_db.create_table( + "update_expr_empty_in", [{"id": 1, "result": "old"}] + ) + predicate = col("id").isin([]) + assert empty_in_table.search().where(predicate).to_arrow().num_rows == 0 + result = empty_in_table.update(where=predicate, values={"result": "new"}) + assert result.rows_updated == 0 + + marker = "__lancedb_binary_placeholder_0__" + binary_schema = pa.schema( + [("payload", pa.binary()), ("text", pa.string()), ("result", pa.string())] + ) + binary_table = mem_db.create_table( + "update_expr_binary", + pa.table( + { + "payload": [b"\x01", b"\x02"], + "text": ["other", marker], + "result": ["old", "old"], + }, + schema=binary_schema, + ), + ) + predicate = (col("payload") == lit(b"\x01")) | (col("text") == marker) + assert binary_table.search().where(predicate).to_arrow().num_rows == 2 + result = binary_table.update(where=predicate, values={"result": "new"}) + assert result.rows_updated == 2 + + nonfinite_table = mem_db.create_table( + "update_expr_nonfinite", + [{"x": 1.0, "result": "old"}, {"x": 2.0, "result": "old"}], + ) + predicate = col("x") < float("inf") + assert nonfinite_table.search().where(predicate).to_arrow().num_rows == 2 + result = nonfinite_table.update(where=predicate, values={"result": "new"}) + assert result.rows_updated == 2 + + float16_table = mem_db.create_table( + "update_expr_float16", + [{"x": 1.0, "result": "old"}, {"x": 3.0, "result": "old"}], + ) + predicate = col("x").cast(pa.float16()) < 2.0 + assert float16_table.search().where(predicate).to_arrow().num_rows == 1 + result = float16_table.update(where=predicate, values={"result": "new"}) + assert result.rows_updated == 1 + + string_cast_table = mem_db.create_table( + "update_expr_string_cast", + [{"x": 1, "result": "old"}, {"x": 2, "result": "old"}], + ) + predicate = col("x").cast("string") == "1" + assert string_cast_table.search().where(predicate).to_arrow().num_rows == 1 + result = string_cast_table.update(where=predicate, values={"result": "new"}) + assert result.rows_updated == 1 + + quoted_identifier_schema = pa.schema( + [("payload", pa.binary()), ("odd'name", pa.int64()), ("result", pa.string())] + ) + quoted_identifier_table = mem_db.create_table( + "update_expr_quoted_identifier", + pa.table( + {"payload": [b"\x01"], "odd'name": [1], "result": ["old"]}, + schema=quoted_identifier_schema, + ), + ) + predicate = (col("payload") == lit(b"\x01")) & (col("odd'name") == 1) + assert quoted_identifier_table.search().where(predicate).to_arrow().num_rows == 1 + result = quoted_identifier_table.update(where=predicate, values={"result": "new"}) + assert result.rows_updated == 1 + + decimal256_schema = pa.schema( + [("val", pa.decimal256(40, 2)), ("result", pa.string())] + ) + decimal256_table = mem_db.create_table( + "update_expr_decimal256", + pa.table( + { + "val": [Decimal("1.00"), Decimal("3.00")], + "result": ["old", "old"], + }, + schema=decimal256_schema, + ), + ) + predicate = col("val") < lit(Decimal("2.00")).cast(pa.decimal256(40, 2)) + assert decimal256_table.search().where(predicate).to_arrow().num_rows == 1 + result = decimal256_table.update(where=predicate, values={"result": "new"}) + assert result.rows_updated == 1 + + binary_empty_table = mem_db.create_table( + "update_expr_binary_empty", + pa.table( + {"payload": [b"\x01", b"\x02"], "result": ["old", "old"]}, + schema=pa.schema([("payload", pa.binary()), ("result", pa.string())]), + ), + ) + predicate = (col("payload") == lit(b"\x01")).isin([]) + assert binary_empty_table.search().where(predicate).to_arrow().num_rows == 0 + assert predicate.to_sql() == "false" + result = binary_empty_table.update(where=predicate, values={"result": "new"}) + assert result.rows_updated == 0 + + def test_update_with_arrow_scalar(mem_db: DBConnection): schema = pa.schema({"id": pa.int64(), "vector": pa.list_(pa.float32(), 4)}) table = mem_db.create_table("my_table", schema=schema) diff --git a/python/python/tests/test_util.py b/python/python/tests/test_util.py index a9b66b2dd..acdef8eac 100644 --- a/python/python/tests/test_util.py +++ b/python/python/tests/test_util.py @@ -7,6 +7,7 @@ import pathlib from typing import Optional import lance +from lance.blob import BlobType as LanceBlobType from lancedb.conftest import MockTextEmbeddingFunction from lancedb.embeddings.base import EmbeddingFunctionConfig from lancedb.embeddings.registry import EmbeddingFunctionRegistry @@ -907,6 +908,165 @@ def test_cast_to_target_schema(): assert output == expected +def test_cast_to_target_schema_coerces_binary_to_blob_v2(): + data = pa.table({"image": pa.array([b"hello", None], type=pa.binary())}) + target = pa.schema([lancedb.blob("image")]) + + output = _cast_to_target_schema(data.to_reader(), target).read_all() + + image = output["image"].chunk(0) + assert type(image.type) is lancedb.BlobType + assert image.storage.to_pylist() == [ + {"data": b"hello", "uri": None, "position": None, "size": None}, + None, + ] + + +def test_cast_to_target_schema_coerces_binary_to_metadata_blob_struct(): + storage = lancedb.blob("image").type.storage_type + target = pa.schema( + [ + pa.field( + "image", + storage, + metadata={ + b"ARROW:extension:name": b"lance.blob.v2", + b"ARROW:extension:metadata": b"", + }, + ) + ] + ) + data = pa.table({"image": pa.array([b"hello", None], type=pa.binary())}) + + output = _cast_to_target_schema(data.to_reader(), target).read_all() + + image = output["image"].chunk(0) + assert not isinstance(image.type, pa.ExtensionType) + assert image.to_pylist() == [ + {"data": b"hello", "uri": None, "position": None, "size": None}, + None, + ] + + +def test_cast_to_target_schema_coerces_nested_binary_blob(): + data = pa.table( + { + "info": pa.array( + [{"blob": b"hello"}, {"blob": None}], + type=pa.struct([pa.field("blob", pa.binary())]), + ) + } + ) + target = pa.schema([pa.field("info", pa.struct([lancedb.blob("blob")]))]) + + output = _cast_to_target_schema(data.to_reader(), target).read_all() + + blob = output["info"].chunk(0).field("blob") + assert type(blob.type) is lancedb.BlobType + assert blob.storage.to_pylist() == [ + {"data": b"hello", "uri": None, "position": None, "size": None}, + None, + ] + + +def test_cast_to_target_schema_coerces_list_binary_blob_with_inferred_child_name(): + data = pa.table( + {"images": pa.array([[b"a", b"b"], None], type=pa.list_(pa.binary()))} + ) + target = pa.schema([pa.field("images", pa.list_(lancedb.blob("image")))]) + + output = _cast_to_target_schema(data.to_reader(), target).read_all() + + images = output["images"].chunk(0) + assert images.type.value_field.name == "image" + assert type(images.type.value_type) is lancedb.BlobType + assert images.to_pylist()[1] is None + assert images.values.storage.to_pylist() == [ + {"data": b"a", "uri": None, "position": None, "size": None}, + {"data": b"b", "uri": None, "position": None, "size": None}, + ] + + +def test_list_blob_coercion_preserves_null_slots_with_nonzero_extent(): + child = pa.field("image", pa.binary()) + source = pa.ListArray.from_arrays( + pa.array([0, 2, 4], type=pa.int32()), + pa.array([b"a", b"b", b"dead", b"beef"], type=pa.binary()), + mask=pa.array([False, True]), + ).cast(pa.list_(child)) + target = pa.schema([pa.field("images", pa.list_(lancedb.blob("image")))]) + + output = _cast_to_target_schema( + pa.table({"images": source}).to_reader(), target + ).read_all() + + images = output["images"].chunk(0) + assert images.to_pylist()[1] is None + assert [b["data"] for b in images.to_pylist()[0]] == [b"a", b"b"] + + +def test_fixed_size_list_blob_coercion_keeps_null_rows(): + child = pa.field("frame", pa.binary()) + source = ( + pa.FixedSizeListArray.from_arrays( + pa.array([b"a", b"b", b"c", b"d"], type=pa.binary()), 2 + ) + .take(pa.array([0, None], type=pa.int32())) + .cast(pa.list_(child, 2)) + ) + target = pa.schema([pa.field("frames", pa.list_(lancedb.blob("frame"), 2))]) + + output = _cast_to_target_schema( + pa.table({"frames": source}).to_reader(), target + ).read_all() + + frames = output["frames"].chunk(0) + assert frames.to_pylist()[1] is None + assert [b["data"] for b in frames.to_pylist()[0]] == [b"a", b"b"] + + +def test_cast_to_target_schema_accepts_pylance_blob_v2(): + target_type = lancedb.BlobType() + source = lance.blob_array([b"hello", None]) + assert type(source.type) is LanceBlobType + assert type(source.type) is type(target_type) + data = pa.table({"image": source}) + target = pa.schema([pa.field("image", target_type)]) + + output = _cast_to_target_schema(data.to_reader(), target).read_all() + + image = output["image"].chunk(0) + assert type(image.type) is LanceBlobType + assert image.type == target_type + assert image.storage.to_pylist() == [ + {"data": b"hello", "uri": None, "position": None, "size": None}, + None, + ] + + +def test_cast_to_target_schema_rejects_different_blob_v2_class(): + class OtherBlobType(pa.ExtensionType): + def __init__(self): + super().__init__(lancedb.BlobType().storage_type, "lance.blob.v2") + + def __arrow_ext_serialize__(self) -> bytes: + return b"" + + @classmethod + def __arrow_ext_deserialize__( + cls, storage_type: pa.DataType, serialized: bytes + ) -> "OtherBlobType": + return cls() + + storage = lance.blob_array([b"hello"]).storage + source = pa.ExtensionArray.from_storage(OtherBlobType(), storage) + data = pa.table({"image": source}) + target = pa.schema([lancedb.blob("image")]) + + with pytest.raises(pa.ArrowTypeError, match="different extension type"): + _cast_to_target_schema(data.to_reader(), target).read_all() + + def test_sanitize_data_stream(): # Make sure we don't collect the whole stream when running sanitize_data schema = pa.schema({"a": pa.int32()}) diff --git a/python/src/expr.rs b/python/src/expr.rs index eae1d96ec..79b448fdf 100644 --- a/python/src/expr.rs +++ b/python/src/expr.rs @@ -130,6 +130,14 @@ impl PyExpr { // ── utilities ──────────────────────────────────────────────────────────── + /// Return the referenced column name for a bare column expression. + fn column_name(&self) -> Option { + match &self.0 { + DfExpr::Column(column) if column.relation.is_none() => Some(column.name.clone()), + _ => None, + } + } + /// Render the expression as a SQL string (useful for debugging). fn to_sql(&self) -> PyResult { lancedb::expr::expr_to_sql_string(&self.0).map_err(|e| PyValueError::new_err(e.to_string())) diff --git a/python/src/query.rs b/python/src/query.rs index 014e79e2d..38153729f 100644 --- a/python/src/query.rs +++ b/python/src/query.rs @@ -1,6 +1,7 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright The LanceDB Authors +use std::collections::HashMap; use std::sync::Arc; use std::time::Duration; @@ -325,6 +326,7 @@ pub struct PyQueryRequest { pub filter: Option, pub full_text_search: Option>, pub select: PySelect, + pub select_source_columns: Option>, pub fast_search: Option, pub with_row_id: Option, pub use_lsm: Option, @@ -355,6 +357,7 @@ impl From for PyQueryRequest { full_text_search: query_request .full_text_search .map(|fts| PyLanceDB(fts.query)), + select_source_columns: PySelect::source_columns(&query_request.select), select: PySelect(query_request.select), fast_search: Some(query_request.fast_search), with_row_id: Some(query_request.with_row_id), @@ -380,6 +383,7 @@ impl From for PyQueryRequest { offset: vector_query.base.offset, filter: vector_query.base.filter.map(PyQueryFilter), full_text_search: None, + select_source_columns: PySelect::source_columns(&vector_query.base.select), select: PySelect(vector_query.base.select), fast_search: Some(vector_query.base.fast_search), with_row_id: Some(vector_query.base.with_row_id), @@ -412,6 +416,25 @@ impl From for PyQueryRequest { #[derive(Clone)] pub struct PySelect(Select); +impl PySelect { + fn source_columns(select: &Select) -> Option> { + match select { + Select::Expr(pairs) => Some( + pairs + .iter() + .filter_map(|(output, expr)| match expr { + lancedb::expr::DfExpr::Column(column) if column.relation.is_none() => { + Some((output.clone(), column.name.clone())) + } + _ => None, + }) + .collect(), + ), + _ => None, + } + } +} + impl<'py> IntoPyObject<'py> for PySelect { type Target = PyAny; type Output = Bound<'py, Self::Target>; diff --git a/rust/lancedb/Cargo.toml b/rust/lancedb/Cargo.toml index 8276e5bb3..881a5017e 100644 --- a/rust/lancedb/Cargo.toml +++ b/rust/lancedb/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb" -version = "0.38.0-beta.10" +version = "0.38.0-beta.12" edition.workspace = true description = "LanceDB: A serverless, low-latency vector database for AI applications" license.workspace = true diff --git a/rust/lancedb/src/database/listing.rs b/rust/lancedb/src/database/listing.rs index 064b5d28f..c22b73dd7 100644 --- a/rust/lancedb/src/database/listing.rs +++ b/rust/lancedb/src/database/listing.rs @@ -13,7 +13,7 @@ use lance::dataset::{ReadParams, WriteMode, builder::DatasetBuilder}; use lance::io::{ObjectStore, ObjectStoreParams, WrappingObjectStore}; use lance_datafusion::utils::StreamingWriteSource; use lance_file::version::LanceFileVersion; -use lance_io::object_store::{StorageOptionsAccessor, StorageOptionsProvider}; +use lance_io::object_store::{ReadDirOptions, StorageOptionsAccessor, StorageOptionsProvider}; use lance_table::io::commit::commit_handler_from_url; use object_store::local::LocalFileSystem; use snafu::ResultExt; @@ -281,6 +281,22 @@ impl std::fmt::Display for ListingDatabase { } const LANCE_EXTENSION: &str = "lance"; + +/// The table a listed child of the database names, or `None` if the child is not a table. +/// +/// A table is the directory `.lance`; a loose file or any other directory under the +/// database prefix belongs to something else. `dir_suffix` is `.lance`, built once by the +/// caller rather than per child. +/// The table a listed child directory holds, or `None` if it is not a table at all. +/// +/// Only directories are considered, so a loose object named like a table is not one. +fn table_name(location: &object_store::path::Path, dir_suffix: &str) -> Option { + location + .filename()? + .strip_suffix(dir_suffix) + .map(String::from) + .filter(|name| !name.is_empty()) +} const ENGINE: &str = "engine"; const MIRRORED_STORE: &str = "mirroredStore"; @@ -944,51 +960,72 @@ impl Database for ListingDatabase { Ok(f) } + /// List the tables in the database, a page at a time. + /// + /// The page_token is opaque, unlike the `start_after` parameter of [`Self::table_names()`]. + /// + /// When there are no more results, the returned page_token will be None. + /// + /// `limit` is the maximum number of tables to return in the response. But it is possible + /// for the response to contain fewer than `limit` tables, even when there are more tables + /// to return. Clients should check the returned page_token to determine if there are + /// more results, rather than relying on the number of tables returned. + /// + /// The order that results are returned in not guaranteed to be stable across calls, + /// so clients should not rely on it. async fn list_tables(&self, request: ListTablesRequest) -> Result { if request.id.as_ref().map(|v| !v.is_empty()).unwrap_or(false) { return self.namespace_database().list_tables(request).await; } - let mut f = self - .object_store - .read_dir(self.base_path.clone()) - .await? - .iter() - .map(Path::new) - .filter(|path| { - let is_lance = path - .extension() - .and_then(|e| e.to_str()) - .map(|e| e == LANCE_EXTENSION); - is_lance.unwrap_or(false) - }) - .filter_map(|p| p.file_stem().and_then(|s| s.to_str().map(String::from))) - .collect::>(); - f.sort(); + let limit = request.limit.map(|limit| limit.max(0) as usize); + let dir_suffix = format!(".{LANCE_EXTENSION}"); + let mut tables = Vec::new(); + let mut page_token = request.page_token.filter(|token| !token.is_empty()); - // Handle pagination with page_token - if let Some(ref page_token) = request.page_token { - let index = f - .iter() - .position(|name| name.as_str() > page_token.as_str()) - .unwrap_or(f.len()); - f.drain(0..index); + // A page of nothing: the store rejects a limit of zero, and no table was handed over + // for a token to resume after. + if limit == Some(0) { + return Ok(ListTablesResponse { + context: None, + tables, + page_token: None, + }); } - // Determine if there's a next page. The token is the last name of this page, - // not the first of the next one: the next page resumes strictly after the - // token, so naming the next page's first entry would skip it. - let next_page_token = match request.limit { - Some(limit) if f.len() > limit as usize => { - f.truncate(limit as usize); - f.last().cloned() + loop { + // Ask only for what the page still has room for, so a database holding more + // than one page costs one request per page rather than one per table. + let listing = self + .object_store + .read_dir_page( + self.base_path.clone(), + ReadDirOptions { + page_token: page_token.take(), + limit: limit.map(|limit| limit - tables.len()), + }, + ) + .await?; + page_token = listing.page_token; + // Only child directories can be tables, and the store already separates them + // out, so the objects in the page are not looked at. + tables.extend( + listing + .result + .common_prefixes + .iter() + .filter_map(|location| table_name(location, &dir_suffix)), + ); + // Children that are not tables leave the page short of the limit, so keep + // going until the page is full or the database runs out. + if page_token.is_none() || limit.is_none_or(|limit| tables.len() >= limit) { + break; } - _ => None, - }; + } Ok(ListTablesResponse { context: None, - tables: f, - page_token: next_page_token, + tables, + page_token, }) } @@ -1484,6 +1521,182 @@ mod tests { use tokio::sync::Barrier; use tokio::time::timeout; + async fn create_tables(db: &ListingDatabase, names: &[&str]) { + let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)])); + for name in names { + db.create_table(CreateTableRequest { + name: name.to_string(), + namespace_path: vec![], + data: Box::new(RecordBatch::new_empty(schema.clone())) as Box, + mode: CreateTableMode::Create, + write_options: Default::default(), + location: None, + namespace_client: None, + }) + .await + .unwrap(); + } + } + + /// Every table in the database, taken `limit` at a time, which is how a caller walks a + /// listing: the token ends the walk, never a short page. + async fn walk(db: &ListingDatabase, limit: Option) -> Vec { + let mut seen = Vec::new(); + let mut page_token = None; + loop { + let page = db + .list_tables(ListTablesRequest { + limit, + page_token, + ..Default::default() + }) + .await + .unwrap(); + seen.extend(page.tables); + page_token = page.page_token; + if page_token.is_none() { + return seen; + } + assert!( + seen.len() < 100, + "the walk is serving tables more than once" + ); + } + } + + /// Paging with the returned token has to visit every table exactly once, whatever the + /// page size, with nothing lost or repeated at a boundary. + #[rstest::rstest] + #[tokio::test] + async fn test_list_tables_pages_over_every_table_once(#[values(1, 2, 3, 5, 10)] limit: i32) { + let (_tempdir, db) = setup_database().await; + create_tables(&db, &["a", "b", "c", "d", "e"]).await; + + assert_eq!(walk(&db, Some(limit)).await, vec!["a", "b", "c", "d", "e"]); + } + + /// The token is opaque: it is whatever resumes the store the database sits on, not a + /// table name. Callers hand it back and nothing else. + /// + /// Nothing validates a token, so one invented by a caller is read as a position rather + /// than refused — which is why the token has to come back from a previous page. + #[tokio::test] + async fn test_the_page_token_is_not_a_table_name() { + let (_tempdir, db) = setup_database().await; + create_tables(&db, &["a", "b", "c"]).await; + + let page = db + .list_tables(ListTablesRequest { + limit: Some(1), + ..Default::default() + }) + .await + .unwrap(); + + assert_eq!(page.tables, vec!["a"]); + let token = page.page_token.expect("two tables are still to come"); + assert_ne!(token, "a"); + + // Handing it back is the only thing a caller does with it, and it resumes. + let rest = db + .list_tables(ListTablesRequest { + page_token: Some(token), + ..Default::default() + }) + .await + .unwrap(); + assert_eq!(rest.tables, vec!["b", "c"]); + } + + /// A limit the listing does not fill leaves no token behind, so a caller paging by token + /// stops without asking for an empty page. + #[tokio::test] + async fn test_a_listing_that_runs_out_has_no_token() { + let (_tempdir, db) = setup_database().await; + create_tables(&db, &["a", "b"]).await; + + let page = db + .list_tables(ListTablesRequest { + limit: Some(10), + ..Default::default() + }) + .await + .unwrap(); + + assert_eq!(page.tables, vec!["a", "b"]); + assert_eq!(page.page_token, None); + } + + /// An empty page token means "from the start", which is how a client looping on a token + /// spells its first request. + #[tokio::test] + async fn test_an_empty_page_token_lists_from_the_start() { + let (_tempdir, db) = setup_database().await; + create_tables(&db, &["a", "b"]).await; + + let page = db + .list_tables(ListTablesRequest { + page_token: Some(String::new()), + ..Default::default() + }) + .await + .unwrap(); + + assert_eq!(page.tables, vec!["a", "b"]); + } + + /// Listing follows the order the object store lists directories in, so a name that + /// extends another comes first: the `-` of `users-archive.lance` sorts below the `.` of + /// `users.lance`. Pagination pushes its cursor into the list request, so it cannot report + /// an order other than the one it resumes in. + #[tokio::test] + async fn test_listing_order_follows_the_store_not_the_table_name() { + let (_tempdir, db) = setup_database().await; + create_tables(&db, &["users", "users-archive", "users.old"]).await; + + assert_eq!( + walk(&db, None).await, + vec!["users-archive", "users", "users.old"] + ); + // And paging reports the same order, so a walk sees each table once. + assert_eq!( + walk(&db, Some(1)).await, + vec!["users-archive", "users", "users.old"] + ); + } + + /// Only directories named `.lance` are tables; loose files and other directories + /// under the database prefix are not. A page spent on them is filled from the next one, + /// so a page holding only non-tables does not read as an empty database. + #[tokio::test] + async fn test_listing_ignores_non_table_children() { + let (tempdir, db) = setup_database().await; + create_tables(&db, &["real"]).await; + std::fs::write(tempdir.path().join("aaa-loose.lance"), b"not a table").unwrap(); + create_dir_all(tempdir.path().join("aaa-scratch")).unwrap(); + + let page = db + .list_tables(ListTablesRequest { + limit: Some(1), + ..Default::default() + }) + .await + .unwrap(); + + assert_eq!(page.tables, vec!["real"]); + } + + #[tokio::test] + async fn listing_ignores_empty_table_name() { + let (tempdir, db) = setup_database().await; + create_dir_all(tempdir.path().join(".lance")).unwrap(); + let page = db.list_tables(ListTablesRequest::default()).await.unwrap(); + assert!( + page.tables.is_empty(), + "invalid empty table name was listed" + ); + } + async fn setup_database() -> (tempfile::TempDir, ListingDatabase) { let tempdir = tempdir().unwrap(); let uri = tempdir.path().to_str().unwrap(); diff --git a/rust/lancedb/src/expr.rs b/rust/lancedb/src/expr.rs index da69914e3..1625d9632 100644 --- a/rust/lancedb/src/expr.rs +++ b/rust/lancedb/src/expr.rs @@ -157,7 +157,7 @@ mod tests { use datafusion_common::ScalarValue; let expr = col("data").eq(lit(ScalarValue::Binary(Some(vec![0xca, 0xfe])))); let sql = expr_to_sql_string(&expr).unwrap(); - assert_eq!(sql, "(data = X'CAFE')"); + assert_eq!(sql, "(`data` = X'CAFE')"); } #[test] @@ -167,7 +167,7 @@ mod tests { let int_expr = col("id").gt(lit(5i64)); let combined = bin_expr.and(int_expr); let sql = expr_to_sql_string(&combined).unwrap(); - assert_eq!(sql, "((data = X'01') AND (id > 5))"); + assert_eq!(sql, "((`data` = X'01') AND (id > 5))"); } #[test] @@ -185,7 +185,7 @@ mod tests { // serialized correctly (regression test for placeholder rewrite path). let expr = contains(col("data"), lit(ScalarValue::Binary(Some(vec![0xff])))); let sql = expr_to_sql_string(&expr).unwrap(); - assert_eq!(sql, "contains(data, X'FF')"); + assert_eq!(sql, "contains(`data`, X'FF')"); } #[test] @@ -196,7 +196,7 @@ mod tests { .eq(lit(ScalarValue::Binary(Some(vec![0xab, 0xcd])))) .not(); let sql = expr_to_sql_string(&expr).unwrap(); - assert_eq!(sql, "NOT (data = X'ABCD')"); + assert_eq!(sql, "NOT (`data` = X'ABCD')"); } #[test] @@ -206,6 +206,122 @@ mod tests { assert!(sql.contains("IN"), "expected IN in: {}", sql); } + #[test] + fn test_empty_is_in() { + let expr = is_in(col("id"), vec![]); + assert_eq!(expr_to_sql_string(&expr).unwrap(), "false"); + } + + #[test] + fn test_empty_is_in_discards_binary_children() { + use datafusion_common::ScalarValue; + + let expr = is_in( + col("payload").eq(lit(ScalarValue::Binary(Some(vec![0x01])))), + vec![], + ); + assert_eq!(expr_to_sql_string(&expr).unwrap(), "false"); + } + + #[test] + fn test_keyword_identifier() { + let expr = col("null").eq(lit(1i64)); + assert_eq!(expr_to_sql_string(&expr).unwrap(), "(`null` = 1)"); + } + + #[test] + fn test_decimal_literal_preserves_type() { + use datafusion_common::ScalarValue; + + let expr = col("val").lt(lit(ScalarValue::Decimal128( + Some(1_234_567_890_123_456_790), + 19, + 18, + ))); + let sql = expr_to_sql_string(&expr).unwrap(); + assert_eq!( + sql, + "(val < arrow_cast('1.234567890123456790', 'Decimal128(19, 18)'))" + ); + } + + #[test] + fn test_non_finite_float_literal_preserves_type() { + let expr = col("x").lt(lit(f64::INFINITY)); + assert_eq!( + expr_to_sql_string(&expr).unwrap(), + "(x < arrow_cast('inf', 'Float64'))" + ); + } + + #[test] + fn test_cast_uses_arrow_type_name() { + let string = expr_cast(col("x"), DataType::Utf8); + assert_eq!( + expr_to_sql_string(&string).unwrap(), + "arrow_cast(x, 'Utf8')" + ); + + let int32 = expr_cast(col("x"), DataType::Int32); + assert_eq!( + expr_to_sql_string(&int32).unwrap(), + "arrow_cast(x, 'Int32')" + ); + + let expr = expr_cast(col("x"), DataType::Float16).lt(lit(2.0)); + assert_eq!( + expr_to_sql_string(&expr).unwrap(), + "(arrow_cast(x, 'Float16') < 2.0)" + ); + + let decimal = expr_cast(lit("2.00"), DataType::Decimal256(40, 2)); + assert_eq!( + expr_to_sql_string(&decimal).unwrap(), + "arrow_cast('2.00', 'Decimal256(40, 2)')" + ); + } + + #[test] + fn test_binary_placeholder_does_not_rewrite_user_string() { + use datafusion_common::ScalarValue; + + let marker = "__lancedb_binary_placeholder_0__"; + let expr = col("payload") + .eq(lit(ScalarValue::Binary(Some(vec![0x01])))) + .or(col("text").eq(lit(marker))); + assert_eq!( + expr_to_sql_string(&expr).unwrap(), + "((payload = X'01') OR (`text` = '__lancedb_binary_placeholder_0__'))" + ); + } + + #[test] + fn test_binary_binding_skips_quoted_identifiers() { + use datafusion_common::ScalarValue; + + let expr = col("payload") + .eq(lit(ScalarValue::Binary(Some(vec![0x01])))) + .and(col("odd'name").eq(lit(1i64))) + .and(col("odd`'name").eq(lit(2i64))); + assert_eq!( + expr_to_sql_string(&expr).unwrap(), + "(((payload = X'01') AND (`odd'name` = 1)) AND (`odd``'name` = 2))" + ); + } + + #[test] + fn test_binary_placeholder_collision_search_is_linear() { + use datafusion_common::ScalarValue; + + let collision_shaped = format!("__lancedb_binary_placeholder_0__{}", "_".repeat(64_000)); + let expr = col("payload") + .eq(lit(ScalarValue::Binary(Some(vec![0x01])))) + .and(col("text").eq(lit(collision_shaped.clone()))); + let sql = expr_to_sql_string(&expr).unwrap(); + assert!(sql.contains("X'01'")); + assert!(sql.contains(&format!("'{collision_shaped}'"))); + } + #[test] fn test_multiple_binary_literals() { use datafusion_common::ScalarValue; diff --git a/rust/lancedb/src/expr/sql.rs b/rust/lancedb/src/expr/sql.rs index 24a676485..2a1ca201d 100644 --- a/rust/lancedb/src/expr/sql.rs +++ b/rust/lancedb/src/expr/sql.rs @@ -1,13 +1,24 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright The LanceDB Authors -use std::any::TypeId; +use std::{ + any::TypeId, + collections::{HashMap, HashSet}, +}; +use arrow_array::types::{ + Decimal32Type, Decimal64Type, Decimal128Type, Decimal256Type, DecimalType, +}; +use arrow_schema::DataType; use datafusion_common::ScalarValue; use datafusion_common::tree_node::{Transformed, TreeNode, TreeNodeRecursion}; use datafusion_expr::Expr; +use datafusion_functions::core::expr_fn::{ + arrow_cast as datafusion_arrow_cast, arrow_try_cast as datafusion_arrow_try_cast, +}; use datafusion_sql::sqlparser::{ dialect::{Dialect as SqlParserDialect, GenericDialect}, + keywords::ALL_KEYWORDS, tokenizer::{Token, Tokenizer}, }; use datafusion_sql::unparser::{self, dialect::Dialect as UnparserDialect}; @@ -27,11 +38,13 @@ struct LanceSqlDialect; impl UnparserDialect for LanceSqlDialect { fn identifier_quote_style(&self, identifier: &str) -> Option { - let needs_quote = identifier.chars().any(|c| c.is_ascii_uppercase()) - || !identifier - .chars() - .enumerate() - .all(|(i, c)| c == '_' || c.is_ascii_alphabetic() || (i > 0 && c.is_ascii_digit())); + let identifier_upper = identifier.to_ascii_uppercase(); + let needs_quote = + (identifier_upper != "ID" && ALL_KEYWORDS.contains(&identifier_upper.as_str())) + || identifier.chars().any(|c| c.is_ascii_uppercase()) + || !identifier.chars().enumerate().all(|(i, c)| { + c == '_' || c.is_ascii_alphabetic() || (i > 0 && c.is_ascii_digit()) + }); if needs_quote { Some('`') } else { None } } } @@ -100,24 +113,128 @@ fn bytes_to_hex_sql(bytes: &[u8]) -> String { format!("X'{hex}'") } -/// Returns true if *expr* contains a `Binary` or `LargeBinary` scalar literal -/// anywhere in its subtree. DataFusion's SQL unparser cannot serialize those -/// variants, so we route such expressions through a placeholder-substitution -/// path that emits SQL `X'...'` byte-string literals. -fn has_binary_literal(expr: &Expr) -> bool { - let mut found = false; +fn string_literals(expr: &Expr) -> HashSet { + let mut literals = HashSet::new(); let _ = expr.apply(&mut |e: &Expr| { - if matches!( - e, - Expr::Literal(ScalarValue::Binary(_) | ScalarValue::LargeBinary(_), _) - ) { - found = true; - Ok(TreeNodeRecursion::Stop) - } else { - Ok(TreeNodeRecursion::Continue) + if let Expr::Literal( + ScalarValue::Utf8(Some(value)) + | ScalarValue::LargeUtf8(Some(value)) + | ScalarValue::Utf8View(Some(value)), + _, + ) = e + { + literals.insert(value.clone()); } + Ok(TreeNodeRecursion::Continue) }); - found + literals +} + +fn typed_string_literal(value: String, data_type: DataType) -> Expr { + datafusion_arrow_cast( + Expr::Literal(ScalarValue::Utf8(Some(value)), None), + Expr::Literal(ScalarValue::Utf8(Some(data_type.to_string())), None), + ) +} + +fn next_binary_placeholder(user_strings: &HashSet, next_id: &mut usize) -> String { + loop { + let placeholder = format!("{BINARY_PLACEHOLDER_PREFIX}{}__", *next_id); + *next_id += 1; + if !user_strings.contains(&placeholder) { + return placeholder; + } + } +} + +fn bind_binary_literals( + sql: &str, + mut bindings: HashMap>, +) -> crate::Result { + let bytes = sql.as_bytes(); + let mut output = Vec::with_capacity(bytes.len()); + let mut index = 0; + + // Walk SQL string tokens once. Placeholders are plain, unescaped string + // literals, so this remains linear even when user strings are large or + // deliberately resemble the placeholder prefix. + while index < bytes.len() { + if bytes[index] == b'`' { + let identifier_start = index; + index += 1; + let mut identifier_end = None; + while index < bytes.len() { + if bytes[index] == b'`' { + if index + 1 < bytes.len() && bytes[index + 1] == b'`' { + index += 2; + } else { + index += 1; + identifier_end = Some(index); + break; + } + } else { + index += 1; + } + } + + let Some(identifier_end) = identifier_end else { + return Err(crate::Error::InvalidInput { + message: "unterminated identifier while binding binary literal".to_string(), + }); + }; + output.extend_from_slice(&bytes[identifier_start..identifier_end]); + continue; + } + + if bytes[index] != b'\'' { + output.push(bytes[index]); + index += 1; + continue; + } + + let literal_start = index; + index += 1; + let content_start = index; + let mut escaped = false; + let mut content_end = None; + while index < bytes.len() { + if bytes[index] == b'\'' { + if index + 1 < bytes.len() && bytes[index + 1] == b'\'' { + escaped = true; + index += 2; + } else { + content_end = Some(index); + index += 1; + break; + } + } else { + index += 1; + } + } + + let Some(content_end) = content_end else { + return Err(crate::Error::InvalidInput { + message: "unterminated string while binding binary literal".to_string(), + }); + }; + + let placeholder = &sql[content_start..content_end]; + if !escaped && let Some(value) = bindings.remove(placeholder) { + output.extend_from_slice(bytes_to_hex_sql(&value).as_bytes()); + } else { + output.extend_from_slice(&bytes[literal_start..index]); + } + } + + if !bindings.is_empty() { + return Err(crate::Error::InvalidInput { + message: "failed to bind binary literal while serializing expression".to_string(), + }); + } + + String::from_utf8(output).map_err(|e| crate::Error::InvalidInput { + message: format!("failed to bind binary literal: {e}"), + }) } fn run_unparser(expr: &Expr) -> crate::Result { @@ -130,25 +247,37 @@ fn run_unparser(expr: &Expr) -> crate::Result { } pub fn expr_to_sql_string(expr: &Expr) -> crate::Result { - // Fast path: no binary literals — DataFusion's unparser handles everything. - if !has_binary_literal(expr) { - return run_unparser(expr); - } - - // Slow path: DataFusion's unparser cannot serialize `Binary`/`LargeBinary` - // scalars, so we rewrite each one to a unique string-literal placeholder, - // let the unparser do the rest of the work, then substitute the SQL - // `X'...'` byte-string literal back in. This keeps the operator/function - // serialization logic centralized in DataFusion and works for every - // expression node type the unparser supports. - let mut bindings: Vec> = Vec::new(); + // DataFusion's unparser needs a few adaptations before its SQL can be + // reparsed by Lance without changing the typed expression's semantics: + // + // * decimal literals need an explicit cast to preserve precision and scale; + // * casts need exact Arrow type names rather than SQL type aliases; + // * an empty IN list is valid in DataFusion but invalid SQL; + // * binary literals are unsupported by the unparser and need placeholders. + // Eliminate empty membership expressions before visiting their children. + // Otherwise a discarded binary child could leave behind a stale binding. let rewritten = expr .clone() + .transform(|e: Expr| match e { + Expr::InList(in_list) if in_list.list.is_empty() => Ok(Transformed::yes( + Expr::Literal(ScalarValue::Boolean(Some(in_list.negated)), None), + )), + other => Ok(Transformed::no(other)), + }) + .map_err(|e| crate::Error::InvalidInput { + message: format!("failed to rewrite expression: {e}"), + })? + .data; + + let user_strings = string_literals(&rewritten); + let mut next_placeholder_id = 0; + let mut binary_bindings = HashMap::new(); + let rewritten = rewritten .transform(|e: Expr| match e { Expr::Literal(ScalarValue::Binary(Some(bytes)), m) | Expr::Literal(ScalarValue::LargeBinary(Some(bytes)), m) => { - let placeholder = format!("{}{}__", BINARY_PLACEHOLDER_PREFIX, bindings.len()); - bindings.push(bytes); + let placeholder = next_binary_placeholder(&user_strings, &mut next_placeholder_id); + binary_bindings.insert(placeholder.clone(), bytes); Ok(Transformed::yes(Expr::Literal( ScalarValue::Utf8(Some(placeholder)), m, @@ -158,6 +287,57 @@ pub fn expr_to_sql_string(expr: &Expr) -> crate::Result { | Expr::Literal(ScalarValue::LargeBinary(None), m) => { Ok(Transformed::yes(Expr::Literal(ScalarValue::Null, m))) } + Expr::Literal(ScalarValue::Decimal32(Some(value), precision, scale), _m) => { + let value = Decimal32Type::format_decimal(value, precision, scale); + Ok(Transformed::yes(typed_string_literal( + value, + DataType::Decimal32(precision, scale), + ))) + } + Expr::Literal(ScalarValue::Decimal64(Some(value), precision, scale), _m) => { + let value = Decimal64Type::format_decimal(value, precision, scale); + Ok(Transformed::yes(typed_string_literal( + value, + DataType::Decimal64(precision, scale), + ))) + } + Expr::Literal(ScalarValue::Decimal128(Some(value), precision, scale), _m) => { + let value = Decimal128Type::format_decimal(value, precision, scale); + Ok(Transformed::yes(typed_string_literal( + value, + DataType::Decimal128(precision, scale), + ))) + } + Expr::Literal(ScalarValue::Decimal256(Some(value), precision, scale), _m) => { + let value = Decimal256Type::format_decimal(value, precision, scale); + Ok(Transformed::yes(typed_string_literal( + value, + DataType::Decimal256(precision, scale), + ))) + } + Expr::Literal(ScalarValue::Float16(Some(value)), _m) if !value.is_finite() => Ok( + Transformed::yes(typed_string_literal(value.to_string(), DataType::Float16)), + ), + Expr::Literal(ScalarValue::Float32(Some(value)), _m) if !value.is_finite() => Ok( + Transformed::yes(typed_string_literal(value.to_string(), DataType::Float32)), + ), + Expr::Literal(ScalarValue::Float64(Some(value)), _m) if !value.is_finite() => Ok( + Transformed::yes(typed_string_literal(value.to_string(), DataType::Float64)), + ), + Expr::Cast(cast) => Ok(Transformed::yes(datafusion_arrow_cast( + *cast.expr, + Expr::Literal( + ScalarValue::Utf8(Some(cast.field.data_type().to_string())), + None, + ), + ))), + Expr::TryCast(cast) => Ok(Transformed::yes(datafusion_arrow_try_cast( + *cast.expr, + Expr::Literal( + ScalarValue::Utf8(Some(cast.field.data_type().to_string())), + None, + ), + ))), other => Ok(Transformed::no(other)), }) .map_err(|e| crate::Error::InvalidInput { @@ -165,14 +345,12 @@ pub fn expr_to_sql_string(expr: &Expr) -> crate::Result { })? .data; - let mut sql = run_unparser(&rewritten)?; - for (i, bytes) in bindings.iter().enumerate() { - // The unparser quotes string literals with single quotes, so the - // placeholder appears as `'__lancedb_binary_placeholder___'`. - let quoted = format!("'{}{}__'", BINARY_PLACEHOLDER_PREFIX, i); - sql = sql.replace("ed, &bytes_to_hex_sql(bytes)); + let sql = run_unparser(&rewritten)?; + if binary_bindings.is_empty() { + Ok(sql) + } else { + bind_binary_literals(&sql, binary_bindings) } - Ok(sql) } #[cfg(test)] diff --git a/rust/lancedb/src/table/computed_columns.rs b/rust/lancedb/src/table/computed_columns.rs index 2b95cb34c..0f89ca612 100644 --- a/rust/lancedb/src/table/computed_columns.rs +++ b/rust/lancedb/src/table/computed_columns.rs @@ -22,7 +22,7 @@ use std::collections::{BTreeSet, HashMap}; use std::sync::Arc; -use arrow_schema::{DataType, Field as ArrowField, Schema as ArrowSchema, SchemaRef}; +use arrow_schema::{DataType, Field as ArrowField, Fields, Schema as ArrowSchema, SchemaRef}; use datafusion_common::tree_node::TreeNode; use datafusion_physical_plan::PhysicalExpr; use lance::dataset::NewColumnTransform; @@ -1273,6 +1273,11 @@ pub(crate) fn bind(schema: SchemaRef, column: &str, expression: &str) -> Result< /// refresh time: that the expression parses, that every column it reads /// exists, and that the target name is free. A declaration that survives this /// is one a refresh can always act on. +/// +/// Each accepted column joins the schema the next one resolves against, so a +/// batch may declare `a` and then `b = a + 1` in one commit. Refresh order +/// then matters, and refresh enforces it: `b` is refused while `a` still has +/// unfilled rows. pub(crate) fn plan(schema: SchemaRef, columns: &[(String, String)]) -> Result> { if columns.is_empty() { return Err(Error::InvalidInput { @@ -1280,11 +1285,11 @@ pub(crate) fn plan(schema: SchemaRef, columns: &[(String, String)]) -> Result = Vec::with_capacity(columns.len()); for (name, expression) in columns { - if schema.field_with_name(name).is_ok() || declared.contains(&name.as_str()) { + if schema.field_with_name(name).is_ok() { return Err(Error::ColumnAlreadyExists { name: name.clone() }); } @@ -1292,16 +1297,50 @@ pub(crate) fn plan(schema: SchemaRef, columns: &[(String, String)]) -> Result(), + schema.metadata().clone(), + )); + fields.push(field); } Ok(fields) } +/// Run the schema-level checks of +/// [`AddColumnsBuilder::computed`](super::AddColumnsBuilder::computed) against +/// `schema` without committing: the Function-binding guard and the planning of +/// every declaration. For callers that stage declarations behind other work +/// and need those rejections before any of it lands. +/// +/// Only the schema is consulted. Declaring also refuses a table with an LSM +/// write spec or retained SSTables; that is table state, checked at commit. +/// +/// ``` +/// # use std::sync::Arc; +/// # use arrow_schema::{DataType, Field, Schema}; +/// use lancedb::table::computed_columns::validate_declarations; +/// +/// let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int32, false)])); +/// let declarations = vec![ +/// ("a".to_string(), "x + 1".to_string()), +/// ("b".to_string(), "a * 2".to_string()), +/// ]; +/// assert!(validate_declarations(schema.clone(), &declarations).is_ok()); +/// assert!(validate_declarations(schema, &[("c".into(), "random()".into())]).is_err()); +/// ``` +pub fn validate_declarations(schema: SchemaRef, columns: &[(String, String)]) -> Result<()> { + ensure_no_function_bindings_for_mutation(schema.as_ref(), "schema evolution")?; + plan(schema, columns).map(drop) +} + /// Build the transform that declares `columns` against `schema`. /// /// An all-null column is how a binding with no values yet is carried into a @@ -1340,6 +1379,22 @@ pub(super) async fn add_foreign_kind(table: &crate::Table, name: &str, kind: &st #[cfg(test)] mod tests { + /// The gate's reproducer: the validator applies the same schema-level + /// guard declaring does, so a staging caller is refused before it commits + /// anything else. + #[test] + fn test_validate_declarations_matches_schema_admission_barriers() { + let schema = Arc::new(ArrowSchema::new_with_metadata( + vec![ArrowField::new("x", DataType::Int32, true)], + HashMap::from([( + FUNCTION_BINDINGS_META_KEY.to_string(), + "not valid binding metadata".to_string(), + )]), + )); + let declarations = vec![("a".to_string(), "x + 1".to_string())]; + assert!(super::validate_declarations(schema, &declarations).is_err()); + } + #[test] fn output_arrow_type_grammar_matches_the_shared_golden() { let golden: serde_json::Value = serde_json::from_str(include_str!( @@ -1582,6 +1637,40 @@ mod tests { assert!(declared(&table).await.is_empty()); } + /// A batch may build on itself: one commit, and the later entry's inputs + /// name the earlier one. + #[tokio::test] + async fn test_a_declaration_may_read_one_declared_before_it() { + let table = table_with_ints("chain").await; + let before = table.version().await.unwrap(); + add_computed( + &table, + &[("a".into(), "x + 1".into()), ("b".into(), "a * 2".into())], + ) + .await + .unwrap(); + assert_eq!(table.version().await.unwrap(), before + 1); + let declared = declared(&table).await; + assert_eq!(declared[1].name, "b"); + assert_eq!(declared[1].inputs, vec!["a".to_string()]); + + // Order is the dependency order; reading ahead is still unknown. + let err = add_computed( + &table, + &[("c".into(), "d + 1".into()), ("d".into(), "x + 1".into())], + ) + .await + .unwrap_err(); + assert!(matches!(err, Error::InvalidExpression { column, .. } if column == "c")); + assert!( + validate_declarations( + table.schema().await.unwrap(), + &[("e".into(), "random()".into())] + ) + .is_err() + ); + } + /// A column added by an ordinary transform is materialized, not bound, so /// it carries no declaration to report. #[tokio::test] diff --git a/rust/lancedb/src/table/datafusion/blob_coerce.rs b/rust/lancedb/src/table/datafusion/blob_coerce.rs index cb984f7f4..0596e7a2d 100644 --- a/rust/lancedb/src/table/datafusion/blob_coerce.rs +++ b/rust/lancedb/src/table/datafusion/blob_coerce.rs @@ -36,6 +36,14 @@ pub(super) fn coerce_blob_expr( }; let input_shape = match input_field.data_type() { + DataType::Null => { + let expr: Arc = Arc::new(CastExpr::new( + input_expr, + table_field.data_type().clone(), + None, + )); + return Ok((expr, table_field.clone())); + } DataType::Binary | DataType::LargeBinary | DataType::BinaryView => BlobInputShape::Bytes, DataType::Utf8 | DataType::LargeUtf8 | DataType::Utf8View => BlobInputShape::String, DataType::Struct(children) => { @@ -155,7 +163,7 @@ mod tests { use crate::blob::blob; use arrow_array::{ Array, ArrayRef, BinaryArray, BinaryViewArray, Int32Array, Int64Array, LargeBinaryArray, - RecordBatch, StringArray, StringViewArray, StructArray, UInt8Array, UInt64Array, + NullArray, RecordBatch, StringArray, StringViewArray, StructArray, UInt8Array, UInt64Array, }; use arrow_schema::Schema; use datafusion::prelude::SessionContext; @@ -279,6 +287,18 @@ mod tests { assert_eq!(data.value(0), b"view"); } + #[tokio::test] + async fn null_column_coerces_to_all_null_blob_struct() { + let batch = batch_with_image( + Field::new("image", DataType::Null, true), + Arc::new(NullArray::new(2)), + ); + let coerced = coerce(batch, &blob_table_schema()).await; + let image = image_struct(&coerced); + assert!(image.is_null(0)); + assert!(image.is_null(1)); + } + #[tokio::test] async fn binary_nulls_stay_null_after_coercion() { let batch = batch_with_image( diff --git a/rust/lancedb/src/table/query.rs b/rust/lancedb/src/table/query.rs index b413b2e17..6f6bbf372 100644 --- a/rust/lancedb/src/table/query.rs +++ b/rust/lancedb/src/table/query.rs @@ -17,6 +17,7 @@ use arrow::array::{AsArray, FixedSizeListBuilder, Float32Builder}; use arrow::datatypes::{Float32Type, UInt8Type}; use arrow_array::Array; use arrow_schema::{DataType, Schema}; +use datafusion_common::{Column, DataFusionError, SchemaError}; use datafusion_physical_plan::ExecutionPlan; use datafusion_physical_plan::projection::ProjectionExec; use datafusion_physical_plan::repartition::RepartitionExec; @@ -191,7 +192,7 @@ pub async fn create_plan( if query.query_vector.len() > 1 { if column.is_none() { // Infer a vector column with the same dimension of the query vector. - let arrow_schema = Schema::from(ds_ref.schema()); + let arrow_schema = Schema::from(schema); column = Some(default_vector_column( &arrow_schema, Some(query.query_vector[0].len() as i32), @@ -268,7 +269,7 @@ pub async fn create_plan( let column = if let Some(col) = column { col } else { - let arrow_schema = Schema::from(ds_ref.schema()); + let arrow_schema = Schema::from(schema); default_vector_column(&arrow_schema, Some(query_vector.len() as i32))? }; @@ -374,7 +375,97 @@ pub async fn create_plan( scanner.order_by(Some(order_by.clone()))?; } - Ok(scanner.create_plan().await?) + scanner + .create_plan() + .await + .map_err(|error| enrich_lance_field_not_found(error, schema)) +} + +/// Replace DataFusion's top-level field candidates with qualified leaf paths. +/// +/// DataFusion resolves nested fields but its `FieldNotFound` error only lists the +/// top-level Arrow fields. This makes a missing leaf look unavailable even when it +/// exists below a struct. Keep every other Lance/DataFusion error unchanged and +/// enrich only this one schema error at the LanceDB query boundary. +fn enrich_lance_field_not_found( + error: lance::Error, + schema: &lance_core::datatypes::Schema, +) -> Error { + let Some(field) = find_missing_field(&error) else { + return error.into(); + }; + field_not_found_error(field, &Schema::from(schema)) +} + +fn field_not_found_diagnostic( + error: &(dyn std::error::Error + 'static), + schema: &Schema, +) -> Option { + let field = find_missing_field(error)?; + Some(field_not_found_error(field, schema)) +} + +fn field_not_found_error(field: &Column, schema: &Schema) -> Error { + let valid_fields = leaf_field_paths(schema); + let mut message = format!("Schema error: No field named {}", field.quoted_flat_name()); + if !valid_fields.is_empty() { + message.push_str(". Valid fields are "); + message.push_str(&valid_fields.join(", ")); + } + message.push('.'); + + Error::InvalidInput { message } +} + +fn find_missing_field<'a>(error: &'a (dyn std::error::Error + 'static)) -> Option<&'a Column> { + if let Some(DataFusionError::SchemaError(schema_error, _)) = + error.downcast_ref::() + && let SchemaError::FieldNotFound { field, .. } = schema_error.as_ref() + { + return Some(field); + } + + error.source().and_then(find_missing_field) +} + +fn leaf_field_paths(schema: &Schema) -> Vec { + fn format_segment(segment: &str) -> String { + // Quote every segment instead of maintaining a SQL keyword list. Bare + // lowercase names such as `true` can be parsed as expressions rather + // than identifiers, while backticks preserve all field names in both + // local SQL parsers. + format!("`{}`", segment.replace('`', "``")) + } + + fn visit(fields: &arrow_schema::Fields, path: &mut Vec, paths: &mut Vec) { + for field in fields { + // Neither local planner can address an empty field-path segment, + // even when it is backtick-quoted. Do not advertise leaves beneath + // such a segment as valid filter fields. + if field.name().is_empty() { + continue; + } + path.push(field.name().clone()); + match field.data_type() { + DataType::Struct(children) if !children.is_empty() => { + visit(children, path, paths); + } + _ => { + paths.push( + path.iter() + .map(|segment| format_segment(segment)) + .collect::>() + .join("."), + ); + } + } + path.pop(); + } + } + + let mut paths = Vec::new(); + visit(schema.fields(), &mut Vec::new(), &mut paths); + paths } //Helper functions below @@ -734,7 +825,10 @@ async fn parse_arrow_ipc_response(bytes: bytes::Bytes) -> Result, dimension: i32) -> FixedSizeListArray { @@ -884,7 +978,6 @@ mod tests { async fn test_execute_query_local_routing() { use crate::connect; use crate::table::query::execute_query; - use arrow_array::{Int32Array, RecordBatch}; use arrow_schema::{DataType, Field, Schema}; let conn = connect("memory://").execute().await.unwrap(); @@ -924,6 +1017,164 @@ mod tests { assert_eq!(count, 2); // 4 and 5 } + #[tokio::test] + async fn test_missing_filter_field_lists_nested_fields_in_local_planners() { + use crate::connect; + use arrow_schema::{DataType, Field, Schema}; + + let conn = connect("memory://").execute().await.unwrap(); + let metadata = Arc::new(StructArray::from(vec![ + ( + Arc::new(Field::new("year", DataType::Int32, false)), + Arc::new(Int32Array::from(vec![2024])) as ArrayRef, + ), + ( + Arc::new(Field::new("genre", DataType::Utf8, false)), + Arc::new(StringArray::from(vec!["fiction"])) as ArrayRef, + ), + ( + Arc::new(Field::new("Title", DataType::Int32, false)), + Arc::new(Int32Array::from(vec![7])) as ArrayRef, + ), + ( + Arc::new(Field::new("true", DataType::Int32, false)), + Arc::new(Int32Array::from(vec![8])) as ArrayRef, + ), + ( + Arc::new(Field::new("", DataType::Int32, false)), + Arc::new(Int32Array::from(vec![10])) as ArrayRef, + ), + ])); + let vector = Arc::new(fixed_size_list_array(vec![0.0, 1.0], 2)); + let schema = Arc::new(Schema::new(vec![ + Field::new("id", DataType::Int32, false), + Field::new("vector", vector.data_type().clone(), false), + Field::new("content", DataType::Utf8, false), + Field::new("metadata", metadata.data_type().clone(), false), + ])); + let batch = RecordBatch::try_new( + schema, + vec![ + Arc::new(Int32Array::from(vec![1])), + vector, + Arc::new(StringArray::from(vec!["example"])), + metadata, + ], + ) + .unwrap(); + let table = conn + .create_table("nested_error", batch) + .execute() + .await + .unwrap(); + + let error = table + .query() + .only_if("year = 2024") + .execute() + .await + .err() + .expect("query should reject the unqualified nested field"); + let case_sensitive_path = "`metadata`.`Title`"; + let keyword_path = "`metadata`.`true`"; + let expected = format!( + "No field named year. Valid fields are `id`, `vector`, `content`, `metadata`.`year`, `metadata`.`genre`, {case_sensitive_path}, {keyword_path}." + ); + + assert!( + error.to_string().contains(&expected), + "unexpected error: {error}" + ); + for (path, value) in [(case_sensitive_path, 7), (keyword_path, 8)] { + table + .query() + .only_if(format!("{path} = {value}")) + .execute() + .await + .expect("the path advertised by the diagnostic should be reusable"); + } + + table.set_unenforced_primary_key(["id"]).await.unwrap(); + table + .set_lsm_write_spec(crate::table::LsmWriteSpec::unsharded()) + .await + .unwrap(); + let lsm_error = table + .query() + .only_if("year = 2024") + .execute() + .await + .err() + .expect("LSM query should reject the unqualified nested field"); + + assert!( + lsm_error.to_string().contains(&expected), + "unexpected LSM error: {lsm_error}" + ); + for (path, value) in [(case_sensitive_path, 7), (keyword_path, 8)] { + table + .query() + .only_if(format!("{path} = {value}")) + .execute() + .await + .expect("the path advertised by the diagnostic should be reusable in LSM queries"); + } + } + + #[test] + fn test_leaf_field_paths_preserve_arbitrary_depth() { + use arrow_schema::{DataType, Field, Schema}; + + fn nested_field(path: &[&str]) -> Field { + let mut segments = path.iter().rev(); + let mut field = Field::new( + *segments.next().expect("path must have a leaf"), + DataType::Int32, + false, + ); + for segment in segments { + field = Field::new(*segment, DataType::Struct(vec![field].into()), false); + } + field + } + + let schema = Schema::new(vec![ + nested_field(&["a", "b", "c", "d", "e"]), + nested_field(&["metadata", "child.with.dot"]), + nested_field(&["metadata", "Title"]), + nested_field(&["metadata", "123child"]), + nested_field(&["metadata", "child`tick"]), + nested_field(&["metadata", ""]), + nested_field(&["", "child"]), + ]); + + assert_eq!( + leaf_field_paths(&schema), + vec![ + "`a`.`b`.`c`.`d`.`e`", + "`metadata`.`child.with.dot`", + "`metadata`.`Title`", + "`metadata`.`123child`", + "`metadata`.`child``tick`", + ] + ); + + let source = DataFusionError::SchemaError( + Box::new(SchemaError::FieldNotFound { + field: Box::new(Column::from_name("missing")), + valid_fields: Vec::new(), + }), + Box::new(None), + ); + let error = field_not_found_diagnostic(&source, &schema).unwrap(); + assert!( + error.to_string().contains( + "Valid fields are `a`.`b`.`c`.`d`.`e`, `metadata`.`child.with.dot`, `metadata`.`Title`, `metadata`.`123child`, `metadata`.`child``tick`" + ), + "unexpected error: {error}" + ); + } + #[derive(Debug, Default)] struct CountingNamespaceClient { query_table_calls: AtomicUsize, diff --git a/rust/lancedb/src/table/query/lsm.rs b/rust/lancedb/src/table/query/lsm.rs index 5c340cc35..86c1fe5f2 100644 --- a/rust/lancedb/src/table/query/lsm.rs +++ b/rust/lancedb/src/table/query/lsm.rs @@ -27,6 +27,8 @@ use std::sync::Arc; use arrow_array::Array; use arrow_schema::{DataType, Schema as ArrowSchema}; +use datafusion::common::{DataFusionError, ToDFSchema}; +use datafusion::prelude::SessionContext; use datafusion_physical_plan::expressions::Column; use datafusion_physical_plan::projection::ProjectionExec; use datafusion_physical_plan::{ExecutionPlan, PhysicalExpr}; @@ -391,7 +393,21 @@ fn base_scanner( } if let Some(filter) = &query.base.filter { scanner = match filter { - QueryFilter::Sql(sql) => scanner.filter(sql)?, + QueryFilter::Sql(sql) => { + // Parse here instead of inside `LsmScanner::filter` so the typed + // DataFusion `FieldNotFound` error is still available for the + // same nested-field enrichment used by the ordinary scanner. + let schema = ArrowSchema::from(dataset.schema()); + let df_schema = schema.clone().to_dfschema().map_err(|error| { + enrich_filter_error(error, &schema, "Failed to create DFSchema") + })?; + let expr = SessionContext::new() + .parse_sql_expr(sql, &df_schema) + .map_err(|error| { + enrich_filter_error(error, &schema, "Failed to parse filter expression") + })?; + scanner.filter_expr(expr) + } QueryFilter::Datafusion(expr) => scanner.filter_expr(expr.clone()), QueryFilter::Substrait(_) => { return Err(Error::NotSupported { @@ -403,6 +419,12 @@ fn base_scanner( Ok(scanner) } +fn enrich_filter_error(error: DataFusionError, schema: &ArrowSchema, context: &str) -> Error { + super::field_not_found_diagnostic(&error, schema).unwrap_or_else(|| Error::InvalidInput { + message: format!("{context}: {error}"), + }) +} + /// Plain scan: filter / projection / limit over base ∪ SSTables ∪ in-memory. /// The plain scan applies limit and offset inside the planner. async fn plain_plan( diff --git a/rust/lancedb/src/table/refresh.rs b/rust/lancedb/src/table/refresh.rs index 35f883411..bc2cc38d1 100644 --- a/rust/lancedb/src/table/refresh.rs +++ b/rust/lancedb/src/table/refresh.rs @@ -7,6 +7,16 @@ //! therefore idempotent and does not observe input mutation -- once a row is //! filled, changing what the expression reads leaves the stored result alone. //! +//! A column's computed inputs are filled first -- the dependency graph is +//! walked once, each reachable column filled once in dependency order, each +//! fill its own commit. Every fill in the pass, the requested column's +//! included, covers only the fragments of the snapshot the pass started +//! from: a commit may rebase over a concurrent append, and the fragment that +//! admits carries placeholder nulls no earlier fill covered, so it waits for +//! a later refresh rather than being read as values. Two concurrent fills of +//! one input collide on its field in lance's conflict check, so a dependent +//! fill can only commit over inputs that were durable when it read them. +//! //! Two passes per fragment. The first scans only the unfilled live rows and //! evaluates the expression over them, which yields the exact fill count and //! decides whether the fragment is staged at all -- a fragment where nothing @@ -41,7 +51,8 @@ use crate::{Error, Result}; /// The result of refreshing a computed column. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] pub struct RefreshColumnResult { - /// Rows that had a value computed. + /// Rows that had a value computed, in the requested column only; inputs + /// filled on its behalf are not counted. #[serde(default)] pub rows_filled: u64, /// The commit version associated with the operation. @@ -52,6 +63,7 @@ pub struct RefreshColumnResult { struct RefreshExecution { result: RefreshColumnResult, source_version: u64, + published_version: Option, } /// Internal implementation of the refresh logic. @@ -74,7 +86,12 @@ async fn execute_refresh_column_with_source( let expression = declared_expression(&dataset, column)?; let schema = Arc::new(ArrowSchema::from(dataset.schema())); - let bound = Arc::new(super::computed_columns::bind(schema, column, &expression)?); + let bound = Arc::new(super::computed_columns::bind( + schema.clone(), + column, + &expression, + )?); + ensure_inputs_filled(&dataset, &schema, column, &bound).await?; let field = dataset .schema() .field(column) @@ -100,25 +117,25 @@ async fn execute_refresh_column_with_source( replacements.push(fragment.write_columns(values, &column_schema).await?); } + let source_version = dataset.version().version; if replacements.is_empty() { - let source_version = dataset.version().version; return Ok(RefreshExecution { result: RefreshColumnResult { rows_filled: 0, version: source_version, }, source_version, + published_version: None, }); } - let read_version = dataset.version().version; // The dataset's own session, so registrations and caches survive the // commit being installed on the handle. let session = dataset.session(); let new_dataset = Dataset::commit( WriteDestination::Dataset(dataset.clone()), Operation::DataReplacement { replacements }, - Some(read_version), + Some(source_version), None, None, session, @@ -133,10 +150,52 @@ async fn execute_refresh_column_with_source( rows_filled, version, }, - source_version: read_version, + source_version, + published_version: Some(version), }) } +/// Refuse while a computed input still has rows a refresh of it would fill: +/// read now, its placeholder null would be evaluated as a value and kept. +async fn ensure_inputs_filled( + dataset: &Dataset, + schema: &Arc, + column: &str, + bound: &BoundExpression, +) -> Result<()> { + for input in &bound.roots { + let Some(declaration) = schema + .field_with_name(input) + .ok() + .and_then(computed_column_from_field) + else { + continue; + }; + let ComputedColumnKind::Sql { expression } = &declaration.kind else { + return Err(Error::NotSupported { + message: format!( + "computed column '{column}' reads '{input}', whose fill state this \ + refresh cannot check; refresh '{input}' first" + ), + }); + }; + let input_bound = super::computed_columns::bind(schema.clone(), input, expression)?; + let mut unfilled = 0u64; + for fragment in dataset.get_fragments() { + unfilled += count_fragment_gains(dataset, &fragment, &input_bound, input).await?; + } + if unfilled > 0 { + return Err(Error::InvalidInput { + message: format!( + "computed column '{column}' reads '{input}', which has {unfilled} unfilled \ + rows; refresh '{input}' first" + ), + }); + } + } + Ok(()) +} + /// Run the refresh as a [`Job`] in this process. pub(crate) async fn execute_refresh_column_async( table: &NativeTable, @@ -160,8 +219,7 @@ pub(crate) async fn execute_refresh_column_async( rows_failed: 0, rows_remaining: 0, source_version: execution.source_version, - published_version: (execution.result.rows_filled > 0) - .then_some(execution.result.version), + published_version: execution.published_version, }) }))) } @@ -384,7 +442,8 @@ mod tests { .version) } - async fn read(table: &Table, column: &str) -> Vec> { + async fn read(table: &Table, column: &str) -> Vec> { + use arrow_array::{Array, Int64Array}; let batches = table .query() .select(Select::columns(&[column])) @@ -394,15 +453,19 @@ mod tests { .try_collect::>() .await .unwrap(); - let mut values: Vec> = batches + let mut values: Vec> = batches .iter() .flat_map(|batch| { - batch[column] - .as_any() - .downcast_ref::() - .unwrap() - .iter() - .collect::>() + let array = &batch[column]; + match array.as_any().downcast_ref::() { + Some(ints) => ints.iter().map(|v| v.map(i64::from)).collect::>(), + None => array + .as_any() + .downcast_ref::() + .unwrap() + .iter() + .collect::>(), + } }) .collect(); values.sort(); @@ -414,6 +477,98 @@ mod tests { table.add(batch).execute().await.unwrap(); } + /// The gate's reproducer: `b = coalesce(a, 0)` refreshed before `a` + /// must not bake zeros from `a`'s placeholder null. It is refused, and + /// names the input, until `a` is filled -- after every append too. + #[tokio::test] + async fn test_dependent_refresh_refuses_an_unfilled_input() { + let table = table_with("dependent_refresh_order", vec![1, 2, 3]).await; + table + .add_columns() + .computed("a", "x + 1") + .computed("b", "coalesce(a, 0)") + .execute() + .await + .unwrap(); + + let err = table.refresh_column("b").await.unwrap_err(); + assert!( + matches!(&err, Error::InvalidInput { message } if message.contains("refresh 'a' first")), + "{err}" + ); + assert_eq!(read(&table, "b").await, vec![None, None, None]); + + assert_eq!(table.refresh_column("a").await.unwrap().rows_filled, 3); + assert_eq!(table.refresh_column("b").await.unwrap().rows_filled, 3); + assert_eq!(read(&table, "b").await, vec![Some(2), Some(3), Some(4)]); + + append(&table, vec![10]).await; + assert!(table.refresh_column("b").await.is_err()); + table.refresh_column("a").await.unwrap(); + assert_eq!(table.refresh_column("b").await.unwrap().rows_filled, 1); + assert_eq!( + table.count_rows(Some("b = 0".to_string())).await.unwrap(), + 0 + ); + } + + /// Names that need quoting, and a nested input, survive the trip through + /// declaration metadata and the dependency check: the recorded inputs + /// are matched by name, never re-parsed as SQL. + #[tokio::test] + async fn test_dependent_refresh_handles_awkward_column_names() { + use arrow_array::{Int32Array, StructArray}; + use arrow_schema::{DataType, Field, Fields}; + + let conn = connect("memory://").execute().await.unwrap(); + let age_fields = Fields::from(vec![Field::new("age", DataType::Int32, true)]); + let meta = StructArray::new( + age_fields.clone(), + vec![Arc::new(Int32Array::from(vec![10, 20])) as _], + None, + ); + let schema = Arc::new(arrow_schema::Schema::new(vec![ + Field::new("camelCase", DataType::Int32, true), + Field::new("with-hyphen", DataType::Int32, true), + Field::new("meta", DataType::Struct(age_fields), true), + ])); + let batch = arrow_array::RecordBatch::try_new( + schema, + vec![ + Arc::new(Int32Array::from(vec![1, 2])) as _, + Arc::new(Int32Array::from(vec![100, 200])) as _, + Arc::new(meta) as _, + ], + ) + .unwrap(); + let table = conn + .create_table("awkward_names", batch) + .execute() + .await + .unwrap(); + + table + .add_columns() + .computed("y", "`camelCase` * 2") + .computed("z", "coalesce(y, 0) + `with-hyphen` + meta.age") + .execute() + .await + .unwrap(); + let z = crate::table::computed_columns::computed_columns( + table.schema().await.unwrap().as_ref(), + ) + .into_iter() + .find(|c| c.name == "z") + .unwrap(); + assert_eq!(z.inputs, vec!["meta.age", "with-hyphen", "y"]); + + let err = table.refresh_column("z").await.unwrap_err(); + assert!(err.to_string().contains("refresh 'y' first"), "{err}"); + assert_eq!(table.refresh_column("y").await.unwrap().rows_filled, 2); + assert_eq!(table.refresh_column("z").await.unwrap().rows_filled, 2); + assert_eq!(read(&table, "z").await, vec![Some(112), Some(224)]); + } + #[tokio::test] async fn test_refresh_fills_a_declared_column() { let table = table_with("refresh_fills", vec![1, 2, 3]).await; @@ -651,7 +806,8 @@ mod tests { let read_back = read(&table, "doubled").await; assert_eq!(read_back.len(), 20_000); - let mut expected: Vec> = values.iter().map(|v| Some(v * 2)).collect(); + let mut expected: Vec> = + values.iter().map(|v| Some(i64::from(v * 2))).collect(); expected.sort(); assert_eq!(read_back, expected); }