Compare commits

..

60 Commits

Author SHA1 Message Date
Xuanwo 3a3ddfda01 test: verify function state survives server restart 2026-08-14 16:55:38 +08:00
Xuanwo c4371eb500 test: add enterprise function reliability e2e 2026-08-14 16:19:35 +08:00
Xuanwo 7a08580400 fix(remote): project generated column job results 2026-08-14 14:27:43 +08:00
Xuanwo 2fccab172f test: add enterprise first-class function e2e 2026-08-14 12:10:33 +08:00
Xuanwo e478b80985 feat: author generated column change jobs 2026-08-12 20:34:39 +08:00
Xuanwo 72767b17fa feat: project definitions from binding snapshots 2026-08-12 20:13:04 +08:00
Xuanwo d7d25cd5ef feat: author generated column refresh jobs 2026-08-12 20:01:09 +08:00
Xuanwo 4843445a7e feat: project generated column definitions 2026-08-12 19:42:36 +08:00
Xuanwo 713510375b feat: resolve generated column functions by id 2026-08-12 19:34:16 +08:00
Xuanwo a9617bf830 feat: submit generated column refresh jobs 2026-08-12 19:25:18 +08:00
Xuanwo 208787ae7b feat: submit generated column change jobs 2026-08-12 19:18:55 +08:00
Xuanwo 9b46e7a448 feat: guard generated metadata on overwrite 2026-08-12 19:08:13 +08:00
Xuanwo 7a46da2e67 feat: guard generated metadata on add columns 2026-08-12 18:52:21 +08:00
Xuanwo 705f7e7760 feat: guard generated metadata on table creation 2026-08-12 18:41:58 +08:00
Xuanwo d38a566282 feat: reserve generated column metadata updates 2026-08-12 18:25:31 +08:00
Xuanwo 5a27c71ab8 feat: reject merge insert on generated columns 2026-08-12 18:12:30 +08:00
Xuanwo 76f6487d92 feat: invalidate generated columns on native delete 2026-08-12 17:57:26 +08:00
Xuanwo 6f6c3c33e0 build: pin Lance zero-row delete attachment fix 2026-08-12 17:31:34 +08:00
Xuanwo 1aa3665d67 feat: invalidate generated columns on native update 2026-08-12 17:03:31 +08:00
Xuanwo 0194f2317a build: pin Lance no-op update attachment fix 2026-08-12 16:44:38 +08:00
Xuanwo c2a647189d feat: invalidate generated columns on native append 2026-08-12 16:14:22 +08:00
Xuanwo df67ee4028 chore: pin Lance A4 substrate 2026-08-12 15:49:25 +08:00
Xuanwo d9e41228c8 feat: plan generated column invalidation 2026-08-12 15:11:25 +08:00
Xuanwo 68597070d2 feat: project remote query function errors 2026-08-12 14:49:01 +08:00
Xuanwo c825737780 feat: guard native generated column queries 2026-08-12 14:14:41 +08:00
Xuanwo 0a43795996 feat: add generated column query guard analysis 2026-08-12 13:45:13 +08:00
Xuanwo 0fa2fa05ad feat(python): expose generated column status 2026-08-12 12:52:27 +08:00
Xuanwo 93ba442ac2 feat: expose generated column status 2026-08-12 12:17:02 +08:00
Xuanwo 7a94ab7d6c feat: add generated column Python API 2026-08-12 11:50:54 +08:00
Xuanwo 6ed1a25439 feat: submit generated column creation jobs 2026-08-12 10:58:04 +08:00
Xuanwo ca1d04db25 feat(python): bind function calls to table snapshots 2026-08-12 10:31:46 +08:00
Xuanwo efe3300404 feat: validate bound function call fields 2026-08-12 10:05:42 +08:00
Xuanwo ecf87f6371 feat: add atomic generated column binding snapshots 2026-08-12 09:49:51 +08:00
Xuanwo 47213e31f8 feat(python): add first-class function call authoring 2026-08-12 09:19:49 +08:00
Xuanwo f65bf89c98 feat(python): add exact function revocation 2026-08-12 08:22:49 +08:00
Xuanwo d902144605 feat(rust): add exact function revocation 2026-08-12 08:06:30 +08:00
Xuanwo a49dc5c71d feat(python): add conditional function name removal 2026-08-12 07:56:07 +08:00
Xuanwo 98fed41efa feat(rust): add conditional function name removal 2026-08-12 07:33:19 +08:00
Xuanwo 1524ee0669 feat(python): add conditional function replacement 2026-08-12 07:02:38 +08:00
Xuanwo 29be3e5509 feat(python): expose function job error codes 2026-08-12 06:44:03 +08:00
Xuanwo 8cedd50495 feat: expose function lookup in Python 2026-08-12 06:25:33 +08:00
Xuanwo b71ada0fae feat: add function catalog lookup 2026-08-12 06:06:18 +08:00
Xuanwo 206efd98ff feat: register functions from Python 2026-08-12 05:34:28 +08:00
Xuanwo 65c0968c0f feat: submit function registration jobs 2026-08-12 05:03:17 +08:00
Xuanwo 2b10f2a7ce feat: bridge Python UDF definitions to Rust 2026-08-12 04:34:58 +08:00
Xuanwo f8bb90405f feat: declare Python function capabilities 2026-08-12 03:57:52 +08:00
Xuanwo 76aac96749 feat: validate Python UDF source packages 2026-08-12 03:47:17 +08:00
Xuanwo 0093bc8179 feat: add Python UDF declarations 2026-08-12 03:28:05 +08:00
Xuanwo ac35a687f1 feat: expose Python function job results 2026-08-12 03:12:36 +08:00
Xuanwo 203f6536a6 feat: expose typed remote job results 2026-08-12 02:09:33 +08:00
Xuanwo 9d3d0d0640 feat: decode remote job results 2026-08-12 01:39:51 +08:00
Xuanwo a9ed8dba27 feat: return results from jobs 2026-08-12 01:18:41 +08:00
Xuanwo 04acf1d3b5 feat: add first-class function job result 2026-08-12 00:43:18 +08:00
Xuanwo 3746118374 feat: add generated column change job spec 2026-08-12 00:26:01 +08:00
Xuanwo d0b5cbe510 feat: add generated column refresh job spec 2026-08-12 00:09:57 +08:00
Xuanwo 7b195adc3a feat: add generated column create job spec 2026-08-11 23:44:27 +08:00
Xuanwo 818d6d1f59 feat: add function registration job spec 2026-08-11 23:26:05 +08:00
Xuanwo 9d589bea44 feat: add function definition contract 2026-08-11 23:11:01 +08:00
Xuanwo 1798ece362 feat: add stable function error codes 2026-08-11 22:43:39 +08:00
Xuanwo 82b82711ba feat: add first-class function value model 2026-08-11 22:25:09 +08:00
161 changed files with 43014 additions and 9080 deletions
+1 -1
View File
@@ -1,5 +1,5 @@
[tool.bumpversion]
current_version = "0.38.0-beta.2"
current_version = "0.37.1-beta.1"
parse = """(?x)
(?P<major>0|[1-9]\\d*)\\.
(?P<minor>0|[1-9]\\d*)\\.
@@ -4,14 +4,14 @@ on:
workflow_call:
inputs:
tag:
description: "Tag name from Lance (e.g. `v7.2.0-beta.1`). If omitted, the newest release is resolved automatically — stable releases are preferred over pre-releases — and the run is skipped if it is not newer than the version currently pinned in Cargo.toml."
description: "Tag name from Lance. If omitted, the skill will use the latest Lance release that needs an update."
required: false
default: ""
type: string
workflow_dispatch:
inputs:
tag:
description: "Tag name from Lance (e.g. `v7.2.0-beta.1`). Leave empty to resolve the newest release automatically — stable releases are preferred over pre-releases — and skip the run if it is not newer than the version currently pinned in Cargo.toml."
description: "Tag name from Lance. Leave empty to use the latest Lance release that needs an update."
required: false
default: ""
type: string
-10
View File
@@ -69,16 +69,6 @@ jobs:
uses: actions/setup-python@v6
with:
python-version: "3.10"
- name: Add swap for Arm fat LTO
if: matrix.config.platform == 'aarch64'
shell: bash
run: |
swap_file="$RUNNER_TEMP/lancedb-swap"
sudo fallocate --length 16G "$swap_file"
sudo chmod 600 "$swap_file"
sudo mkswap "$swap_file"
sudo swapon "$swap_file"
free -h
- uses: ./.github/workflows/build_linux_wheel
with:
python-minor-version: 10
Generated
+68 -54
View File
@@ -959,7 +959,7 @@ dependencies = [
"aws-smithy-runtime-api",
"aws-smithy-types",
"h2 0.3.27",
"h2 0.4.16",
"h2 0.4.14",
"http 0.2.12",
"http 1.5.0",
"http-body 0.4.6",
@@ -3455,8 +3455,8 @@ checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c"
[[package]]
name = "fsst"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"rand 0.9.5",
@@ -3877,9 +3877,9 @@ dependencies = [
[[package]]
name = "h2"
version = "0.4.16"
version = "0.4.14"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a9f37a958b41b3b19ee2707c06439c0e9e547e847223eb791ecb0cb821c65e27"
checksum = "171fefbc92fe4a4de27e0698d6a5b392d6a0e333506bc49133760b3bcf948733"
dependencies = [
"atomic-waker",
"bytes",
@@ -4188,7 +4188,7 @@ dependencies = [
"bytes",
"futures-channel",
"futures-core",
"h2 0.4.16",
"h2 0.4.14",
"http 1.5.0",
"http-body 1.1.0",
"httparse",
@@ -4815,8 +4815,8 @@ checksum = "e037a2e1d8d5fdbd49b16a4ea09d5d6401c1f29eca5ff29d03d3824dba16256a"
[[package]]
name = "lance"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arc-swap",
"arrow",
@@ -4832,6 +4832,7 @@ dependencies = [
"async-recursion",
"async-trait",
"async_cell",
"aws-credential-types",
"aws-sdk-dynamodb",
"byteorder",
"bytes",
@@ -4847,6 +4848,7 @@ dependencies = [
"either",
"fst",
"futures",
"half",
"humantime",
"itertools 0.14.0",
"lance-arrow",
@@ -4888,8 +4890,8 @@ dependencies = [
[[package]]
name = "lance-arrow"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -4911,7 +4913,7 @@ dependencies = [
[[package]]
name = "lance-arrow-scalar"
version = "58.0.0"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -4925,7 +4927,7 @@ dependencies = [
[[package]]
name = "lance-arrow-stats"
version = "58.0.0"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-schema",
@@ -4934,8 +4936,8 @@ dependencies = [
[[package]]
name = "lance-bitpacking"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrayref",
"crunchy",
@@ -4945,8 +4947,8 @@ dependencies = [
[[package]]
name = "lance-core"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -4954,10 +4956,12 @@ dependencies = [
"arrow-schema",
"async-trait",
"blake3",
"byteorder",
"bytes",
"datafusion-common",
"datafusion-sql",
"futures",
"itertools 0.14.0",
"lance-arrow",
"lance-derive",
"libc",
@@ -4975,6 +4979,7 @@ dependencies = [
"snafu 0.9.0",
"tempfile",
"tokio",
"tokio-stream",
"tokio-util",
"tracing",
"twox-hash",
@@ -4983,8 +4988,8 @@ dependencies = [
[[package]]
name = "lance-datafusion"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow",
"arrow-array",
@@ -5003,6 +5008,7 @@ dependencies = [
"jsonb",
"lance-arrow",
"lance-core",
"lance-datagen",
"log",
"pin-project",
"prost",
@@ -5013,8 +5019,8 @@ dependencies = [
[[package]]
name = "lance-datagen"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow",
"arrow-array",
@@ -5031,8 +5037,8 @@ dependencies = [
[[package]]
name = "lance-derive"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"proc-macro2",
"quote",
@@ -5041,8 +5047,8 @@ dependencies = [
[[package]]
name = "lance-encoding"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-arith",
"arrow-array",
@@ -5067,6 +5073,7 @@ dependencies = [
"num-traits",
"prost",
"prost-build",
"rand 0.9.5",
"tokio",
"tracing",
"xxhash-rust",
@@ -5075,8 +5082,8 @@ dependencies = [
[[package]]
name = "lance-file"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-arith",
"arrow-array",
@@ -5107,8 +5114,8 @@ dependencies = [
[[package]]
name = "lance-index"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arc-swap",
"arrow",
@@ -5123,6 +5130,7 @@ dependencies = [
"async-trait",
"bitvec",
"bytes",
"chrono",
"crossbeam-queue",
"datafusion",
"datafusion-common",
@@ -5140,6 +5148,7 @@ dependencies = [
"lance-bitpacking",
"lance-core",
"lance-datafusion",
"lance-datagen",
"lance-encoding",
"lance-file",
"lance-index-core",
@@ -5168,12 +5177,13 @@ dependencies = [
"tempfile",
"tokio",
"tracing",
"uuid",
]
[[package]]
name = "lance-index-core"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-schema",
@@ -5195,8 +5205,8 @@ dependencies = [
[[package]]
name = "lance-io"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow",
"arrow-array",
@@ -5210,6 +5220,7 @@ dependencies = [
"futures",
"http 1.5.0",
"io-uring",
"lance-arrow",
"lance-core",
"lance-namespace",
"log",
@@ -5227,28 +5238,29 @@ dependencies = [
"tokio",
"tracing",
"url",
"uuid",
]
[[package]]
name = "lance-linalg"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-buffer",
"arrow-schema",
"cc",
"half",
"lance-arrow",
"lance-core",
"num-traits",
"rand 0.9.5",
"rayon",
]
[[package]]
name = "lance-namespace"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow",
"async-trait",
@@ -5260,8 +5272,8 @@ dependencies = [
[[package]]
name = "lance-namespace-impls"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow",
"arrow-ipc",
@@ -5300,9 +5312,9 @@ dependencies = [
[[package]]
name = "lance-namespace-reqwest-client"
version = "0.11.0"
version = "0.8.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0a030196da1c994b63a96a4f0bf5b0cfa459fe6dadc9e962320246ca328da22a"
checksum = "ba3f0a235e3ed5f8805205649ccc7d7d0f3df23ce1294242c9265ad488d7f19d"
dependencies = [
"reqwest 0.12.28",
"serde",
@@ -5314,13 +5326,14 @@ dependencies = [
[[package]]
name = "lance-select"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-buffer",
"arrow-schema",
"byteorder",
"bytes",
"itertools 0.14.0",
"lance-core",
"roaring",
@@ -5329,8 +5342,8 @@ dependencies = [
[[package]]
name = "lance-table"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow",
"arrow-array",
@@ -5370,8 +5383,8 @@ dependencies = [
[[package]]
name = "lance-testing"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-schema",
@@ -5384,8 +5397,8 @@ dependencies = [
[[package]]
name = "lance-tokenizer"
version = "11.0.0-beta.15"
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"frostem",
"icu_segmenter",
@@ -5398,7 +5411,7 @@ dependencies = [
[[package]]
name = "lancedb"
version = "0.38.0-beta.2"
version = "0.37.1-beta.1"
dependencies = [
"ahash",
"anyhow",
@@ -5419,6 +5432,7 @@ dependencies = [
"aws-sdk-kms",
"aws-sdk-s3",
"aws-smithy-runtime",
"base64 0.22.1",
"bytes",
"candle-core",
"candle-nn",
@@ -5486,7 +5500,7 @@ dependencies = [
[[package]]
name = "lancedb-nodejs"
version = "0.38.0-beta.2"
version = "0.37.1-beta.1"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -5511,7 +5525,7 @@ dependencies = [
[[package]]
name = "lancedb-python"
version = "0.38.0-beta.2"
version = "0.37.1-beta.1"
dependencies = [
"arrow",
"async-trait",
@@ -8426,7 +8440,7 @@ dependencies = [
"encoding_rs",
"futures-core",
"futures-util",
"h2 0.4.16",
"h2 0.4.14",
"http 1.5.0",
"http-body 1.1.0",
"http-body-util",
@@ -10082,7 +10096,7 @@ dependencies = [
"async-trait",
"base64 0.22.1",
"bytes",
"h2 0.4.16",
"h2 0.4.14",
"http 1.5.0",
"http-body 1.1.0",
"http-body-util",
+14 -17
View File
@@ -13,23 +13,20 @@ categories = ["database-implementations"]
rust-version = "1.91.0"
[workspace.dependencies]
# TEMPORARY: `ObjectStore::read_dir_page` is not in a lance release yet, so these point at
# lance-format/lance#8606 cherry-picked onto the v11.0.0-beta.15 tag. Put the tag back once that
# PR has merged and shipped in a release.
lance = { "version" = "=11.0.0-beta.15", default-features = false, "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-core = { "version" = "=11.0.0-beta.15", "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-datagen = { "version" = "=11.0.0-beta.15", "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-file = { "version" = "=11.0.0-beta.15", "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-io = { "version" = "=11.0.0-beta.15", default-features = false, "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-index = { "version" = "=11.0.0-beta.15", "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-linalg = { "version" = "=11.0.0-beta.15", "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-namespace = { "version" = "=11.0.0-beta.15", "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-namespace-impls = { "version" = "=11.0.0-beta.15", default-features = false, "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-table = { "version" = "=11.0.0-beta.15", "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-testing = { "version" = "=11.0.0-beta.15", "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-datafusion = { "version" = "=11.0.0-beta.15", "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-encoding = { "version" = "=11.0.0-beta.15", "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance-arrow = { "version" = "=11.0.0-beta.15", "rev" = "d7b1d570461c6d2adde8f3a84ae88db4823c726f", "git" = "https://github.com/lance-format/lance.git" }
lance = { version = "=11.0.0-beta.4", default-features = false, rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-core = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-datagen = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-file = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-io = { version = "=11.0.0-beta.4", default-features = false, rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-index = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-linalg = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-namespace = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-namespace-impls = { version = "=11.0.0-beta.4", default-features = false, rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-table = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-testing = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-datafusion = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-encoding = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-arrow = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
ahash = "0.8"
# Note that this one does not include pyarrow
arrow = { version = "58.0.0", optional = false }
-13
View File
@@ -101,19 +101,6 @@ ignore = [
# https://rustsec.org/advisories/RUSTSEC-2026-0195
{ id = "RUSTSEC-2026-0194", reason = "transitive via inferno/lance/opendal; XML from trusted cloud endpoints, not attacker-controlled" },
{ id = "RUSTSEC-2026-0195", reason = "transitive via inferno/lance/opendal; XML from trusted cloud endpoints, not attacker-controlled" },
# smartstring: unmaintained — the repository was archived by its author on
# 2026-05-03. Not a vulnerability. Reached only transitively through polars
# (polars-core/-io/-ops/-time/-utils); nothing in LanceDB depends on it directly.
# The advisory states no safe upgrade is available: upstream recommends
# compact_str/smol_str, so clearing this requires polars to migrate.
# https://rustsec.org/advisories/RUSTSEC-2026-0249
{ id = "RUSTSEC-2026-0249", reason = "smartstring unmaintained via polars; no fixed upstream release" },
# h2 0.3: empty DATA frames can be queued without limit. The patched
# h2 0.4 line is locked to 0.4.16, but no patched 0.3 release exists.
# The old copy is pulled in by aws-smithy's legacy hyper 0.14 client.
# https://rustsec.org/advisories/RUSTSEC-2026-0258
{ id = "RUSTSEC-2026-0258", reason = "h2 0.3 via legacy aws-smithy/hyper 0.14; no patched 0.3 release" },
]
# ---------------------------------------------------------------------------
+1 -33
View File
@@ -14,7 +14,7 @@ Add the following dependency to your `pom.xml`:
<dependency>
<groupId>com.lancedb</groupId>
<artifactId>lancedb-core</artifactId>
<version>0.38.0-beta.2</version>
<version>0.37.1-beta.1</version>
</dependency>
```
@@ -55,38 +55,6 @@ LanceNamespace namespaceClient = LanceDbNamespaceClientBuilder.newBuilder()
| `region(String)` | AWS region (default: "us-east-1") | No |
| `config(String, String)` | Additional configuration parameters | No |
### Opening a Table with Vended Credentials
When the catalog vends temporary object store credentials, open the table through the
namespace client. The Lance dataset builder fetches the table location and storage options
from the catalog and refreshes the credentials when they expire.
```java
import com.lancedb.LanceDbNamespaceClientBuilder;
import org.lance.Dataset;
import org.lance.namespace.LanceNamespace;
import java.util.Arrays;
LanceNamespace namespaceClient = LanceDbNamespaceClientBuilder.newBuilder()
.apiKey(System.getenv("LANCEDB_API_KEY"))
.database(System.getenv("LANCEDB_DATABASE"))
// Set the endpoint for a LanceDB Enterprise deployment.
// .endpoint("https://your-enterprise-endpoint")
.build();
try (Dataset dataset = Dataset.open()
.namespaceClient(namespaceClient)
.tableId(Arrays.asList("my_namespace", "my_table"))
.build()) {
System.out.println("Rows: " + dataset.countRows());
}
```
Do not call `describeTable()` and then open the returned location with `Dataset.open(uri)`.
Opening through `namespaceClient()` is what applies the vended storage options and enables
automatic credential refresh. No object store credentials need to be passed by the application.
## Metadata Operations
### Creating a Namespace Path
+1 -97
View File
@@ -386,29 +386,6 @@ Drop an existing table.
***
### dropTableAsync()
```ts
abstract dropTableAsync(name, namespacePath?): Promise<Job>
```
Start dropping a table and return its cleanup job.
The table may become unavailable before its data files are removed. Wait
on the returned job to know when cleanup has finished.
#### Parameters
* **name**: `string`
* **namespacePath?**: `string`[]
#### Returns
`Promise`&lt;[`Job`](Job.md)&gt;
***
### getJob()
```ts
@@ -529,71 +506,6 @@ Child namespace names and
***
### listTables()
#### listTables(options)
```ts
abstract listTables(options?): Promise<ListTablesResponse>
```
List a page of tables in this database.
Results may be paginated. To retrieve subsequent pages, pass the
`pageToken` returned by a previous call. A page may be shorter than
`limit` without being the last one, so walk until the response carries no
page token:
```ts
const names = [];
let pageToken = undefined;
do {
const page = await conn.listTables({ pageToken, limit: 100 });
names.push(...page.tables);
pageToken = page.pageToken;
} while (pageToken);
```
##### Parameters
* **options?**: `Partial`&lt;[`ListTablesOptions`](../interfaces/ListTablesOptions.md)&gt;
Pagination options
(`pageToken`, `limit`).
##### Returns
`Promise`&lt;[`ListTablesResponse`](../interfaces/ListTablesResponse.md)&gt;
Table names and an optional token
for fetching the next page.
#### listTables(namespacePath, options)
```ts
abstract listTables(namespacePath?, options?): Promise<ListTablesResponse>
```
List a page of tables in this database.
##### Parameters
* **namespacePath?**: `string`[]
The namespace path to list tables from
(defaults to root namespace)
* **options?**: `Partial`&lt;[`ListTablesOptions`](../interfaces/ListTablesOptions.md)&gt;
Pagination options
(`pageToken`, `limit`).
##### Returns
`Promise`&lt;[`ListTablesResponse`](../interfaces/ListTablesResponse.md)&gt;
Table names and an optional token
for fetching the next page.
***
### openTable()
```ts
@@ -655,7 +567,7 @@ a "not supported" error.
***
### ~~tableNames()~~
### tableNames()
#### tableNames(options)
@@ -677,10 +589,6 @@ Tables will be returned in lexicographical order.
`Promise`&lt;`string`[]&gt;
##### Deprecated
Use [Connection.listTables](Connection.md#listtables) instead.
#### tableNames(namespacePath, options)
```ts
@@ -703,7 +611,3 @@ Tables will be returned in lexicographical order.
##### Returns
`Promise`&lt;`string`[]&gt;
##### Deprecated
Use [Connection.listTables](Connection.md#listtables) instead.
+1 -182
View File
@@ -69,34 +69,14 @@ abstract addColumns(newColumnTransforms): Promise<AddColumnsResult>
Add new columns with defined values.
The `{ computed }` form stores the expression rather than evaluating it
now: the column is committed with no values, and rows get them from
[Table#refreshColumn](Table.md#refreshcolumn). Declaring one therefore costs the same on a
large table as on an empty one.
A refresh does not revisit rows it has already filled, so mutating an
input leaves the value computed at fill time; recomputing means dropping
the column and declaring it again. While a declaration reads a column,
that column cannot be renamed, retyped or dropped.
On LanceDB Cloud and Enterprise the expression is planned by the
server, and the refresh runs as a server job -- see
[Table#refreshColumnAsync](Table.md#refreshcolumnasync).
#### Parameters
* **newColumnTransforms**:
\| `Field`&lt;`any`&gt;
\| `Field`&lt;`any`&gt;[]
\| `Schema`&lt;`any`&gt;
\| [`AddColumnsSql`](../interfaces/AddColumnsSql.md)[]
\| `object`
* **newColumnTransforms**: `Field`&lt;`any`&gt; \| `Field`&lt;`any`&gt;[] \| `Schema`&lt;`any`&gt; \| [`AddColumnsSql`](../interfaces/AddColumnsSql.md)[]
Either:
- An array of objects with column names and SQL expressions to calculate values
- A single Arrow Field defining one column with its data type (column will be initialized with null values)
- An array of Arrow Fields defining columns with their data types (columns will be initialized with null values)
- An Arrow Schema defining columns with their data types (columns will be initialized with null values)
- `{ computed }`, declaring columns defined by a SQL expression whose type and inputs are derived from it
#### Returns
@@ -105,13 +85,6 @@ server, and the refresh runs as a server job -- see
A promise that resolves to an object
containing the new version number of the table after adding the columns.
#### Example
```ts
await table.addColumns({ computed: [{ name: "doubled", valueSql: "x * 2" }] });
const { rowsFilled } = await table.refreshColumn("doubled");
```
***
### alterColumns()
@@ -213,39 +186,6 @@ version of the table.
***
### checkpointLsm()
```ts
abstract checkpointLsm(): Promise<void>
```
Converge this table's LSM write path into its base table.
Seals once, then triggers compaction and polls until the L0 that existed
at the start is gone. The target set is fixed at the start, so
generations created *during* the checkpoint are ignored — that is what
lets it terminate under write load, and what makes it best-effort: it
converges the fresh tier as of some instant. Idempotent, abandonable at
any point, and safe to run on a cadence.
There is no liveness bound — the compactor pool is shared across tables,
so a checkpoint queued behind unrelated work looks exactly like one that
is merging. The caller owns the deadline.
#### Returns
`Promise`&lt;`void`&gt;
#### Example
```ts
const before = await table.getLsmStats();
await table.checkpointLsm();
const after = await table.getLsmStats();
```
***
### close()
```ts
@@ -283,24 +223,6 @@ It is a no-op when no writers are cached.
***
### compactLsm()
```ts
abstract compactLsm(): Promise<void>
```
Trigger a background L0 → base compaction pass per bucket.
Returns once the passes are *dispatched*, not once they finish — watch
[Table#getLsmStats](Table.md#getlsmstats) for progress, or use
[Table#checkpointLsm](Table.md#checkpointlsm) to wait for convergence.
#### Returns
`Promise`&lt;`void`&gt;
***
### countRows()
```ts
@@ -499,48 +421,6 @@ Drop an index from the table.
***
### flushLsm()
```ts
abstract flushLsm(): Promise<void>
```
Seal every bucket's active memtable into a new L0 generation.
Returns once the seal is committed. Sealing an empty memtable is a no-op,
so this is safe to call repeatedly.
#### Returns
`Promise`&lt;`void`&gt;
***
### getLsmStats()
```ts
abstract getLsmStats(includeGenerationRows?): Promise<undefined | LsmStats>
```
Read live per-bucket LSM state.
Answers "how far behind is my fresh tier", "which bucket is hot", and
"why is my fresh-tier vector search brute-force". Mutates no table state.
Resolves to `undefined` only when the LSM write path is not enabled.
#### Parameters
* **includeGenerationRows?**: `boolean`
Also count rows per L0 generation.
Off by default because each count opens an uncached Lance dataset.
#### Returns
`Promise`&lt;`undefined` \| [`LsmStats`](../interfaces/LsmStats.md)&gt;
***
### getLsmWriteSpec()
```ts
@@ -838,67 +718,6 @@ for await (const batch of table.query()) {
***
### refreshColumn()
```ts
abstract refreshColumn(column): Promise<RefreshColumnResult>
```
Fill the rows of a computed column that hold no value yet.
Rows appended since the last refresh are filled by the next one; rows
already filled are left as they are, so the call is idempotent and does
not observe a mutated input. Local tables only: a remote refresh runs
as a server job, through [Table#refreshColumnAsync](Table.md#refreshcolumnasync).
#### Parameters
* **column**: `string`
The name of the computed column to fill.
#### Returns
`Promise`&lt;[`RefreshColumnResult`](../interfaces/RefreshColumnResult.md)&gt;
A promise that resolves to the
number of rows filled and the new version number of the table.
***
### refreshColumnAsync()
```ts
abstract refreshColumnAsync(column): Promise<Job>
```
Like [Table#refreshColumn](Table.md#refreshcolumn), but returns a handle to the refresh
job instead of blocking until it completes.
The job may already be complete when returned; callers must not assume
the column is filled until [Job.wait](Job.md#wait) resolves. Invalid input --
an unknown column, or one that is not computed -- rejects here rather
than failing the job. On local tables the job runs in-process; on
LanceDB Cloud and Enterprise it is the server's backfill job.
#### Parameters
* **column**: `string`
The name of the computed column to fill.
#### Returns
`Promise`&lt;[`Job`](Job.md)&gt;
#### Example
```ts
const job = await table.refreshColumnAsync("doubled");
await job.wait();
console.log(await job.status()); // "finished"
```
***
### restore()
```ts
-7
View File
@@ -58,7 +58,6 @@
- [BranchDiff](interfaces/BranchDiff.md)
- [BranchIndexSummary](interfaces/BranchIndexSummary.md)
- [BranchRowCountSummary](interfaces/BranchRowCountSummary.md)
- [BucketStats](interfaces/BucketStats.md)
- [ClientConfig](interfaces/ClientConfig.md)
- [ColumnAlteration](interfaces/ColumnAlteration.md)
- [ColumnOrdering](interfaces/ColumnOrdering.md)
@@ -82,7 +81,6 @@
- [FtsToken](interfaces/FtsToken.md)
- [FullTextQuery](interfaces/FullTextQuery.md)
- [FullTextSearchOptions](interfaces/FullTextSearchOptions.md)
- [GenerationStats](interfaces/GenerationStats.md)
- [HnswPqOptions](interfaces/HnswPqOptions.md)
- [HnswSqOptions](interfaces/HnswSqOptions.md)
- [IndexConfig](interfaces/IndexConfig.md)
@@ -96,11 +94,7 @@
- [JobInfo](interfaces/JobInfo.md)
- [ListNamespacesOptions](interfaces/ListNamespacesOptions.md)
- [ListNamespacesResponse](interfaces/ListNamespacesResponse.md)
- [ListTablesOptions](interfaces/ListTablesOptions.md)
- [ListTablesResponse](interfaces/ListTablesResponse.md)
- [LsmStats](interfaces/LsmStats.md)
- [LsmWriteSpec](interfaces/LsmWriteSpec.md)
- [MemtableStats](interfaces/MemtableStats.md)
- [MergeBlocker](interfaces/MergeBlocker.md)
- [MergeBranchResult](interfaces/MergeBranchResult.md)
- [MergePreview](interfaces/MergePreview.md)
@@ -111,7 +105,6 @@
- [OptimizeOptions](interfaces/OptimizeOptions.md)
- [OptimizeStats](interfaces/OptimizeStats.md)
- [QueryExecutionOptions](interfaces/QueryExecutionOptions.md)
- [RefreshColumnResult](interfaces/RefreshColumnResult.md)
- [RemovalStats](interfaces/RemovalStats.md)
- [RenameTableOptions](interfaces/RenameTableOptions.md)
- [RestNamespaceConfig](interfaces/RestNamespaceConfig.md)
-116
View File
@@ -1,116 +0,0 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / BucketStats
# Interface: BucketStats
Live state of one bucket. A table is N buckets on one node; flattening to a
single number hides the one hot bucket that is usually why someone opened
this endpoint.
## Properties
### compacting
```ts
compacting: boolean;
```
Whether a pass owns this bucket's compaction latch right now. Says *a*
driver is running, not *whose*, and the latch is held from dispatch —
including while the pass queues for a pod-wide compactor permit. Read it
as "do not pile on", never as "mine is progressing".
***
### currentGeneration
```ts
currentGeneration: number;
```
The generation the active memtable will become.
***
### generations
```ts
generations: GenerationStats[];
```
Flushed L0 generations not yet merged into the base table.
***
### manifestVersion
```ts
manifestVersion: number;
```
Version of the shard manifest these numbers were read from.
***
### memtables?
```ts
optional memtables: MemtableStats[];
```
Oldest first, active last. Absent for a `"Sealed"` bucket, whose
in-memory state is torn down.
***
### replayAfterWalEntryPosition
```ts
replayAfterWalEntryPosition: number;
```
WAL position replay resumes from.
***
### shardId
```ts
shardId: string;
```
The shard this bucket writes.
***
### status
```ts
status: string;
```
`"Active"` or `"Sealed"` (drop-table 2PC in flight).
***
### walEntryPositionLastSeen
```ts
walEntryPositionLastSeen: number;
```
Highest WAL position the writer has seen. The difference against
`replayAfterWalEntryPosition` is the WAL lag.
***
### writerEpoch
```ts
writerEpoch: number;
```
Epoch of the writer that currently owns the shard.
-40
View File
@@ -1,40 +0,0 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / GenerationStats
# Interface: GenerationStats
One flushed L0 generation.
## Properties
### bytes
```ts
bytes: number;
```
On-disk size of the generation.
***
### generation
```ts
generation: number;
```
The generation number. Increases as memtables are sealed into L0.
***
### rows?
```ts
optional rows: number;
```
Present only when `includeGenerationRows` was requested. Off by default
because each count opens an uncached Lance dataset.
@@ -1,33 +0,0 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / ListTablesOptions
# Interface: ListTablesOptions
## Properties
### limit?
```ts
optional limit: number;
```
An upper bound on how many tables to return.
A page may hold fewer than this and still not be the last one, so continue
while the response carries a page token rather than while pages are full.
***
### pageToken?
```ts
optional pageToken: string;
```
Token from a previous response for pagination.
The token is opaque: it carries whatever the database needs to resume, and
callers should not construct or interpret one.
@@ -1,23 +0,0 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / ListTablesResponse
# Interface: ListTablesResponse
## Properties
### pageToken?
```ts
optional pageToken: string;
```
***
### tables
```ts
tables: string[];
```
-22
View File
@@ -1,22 +0,0 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / LsmStats
# Interface: LsmStats
Live per-bucket LSM state, as returned by `Table#getLsmStats`.
Nothing here is derived: sums and differences (total L0 bytes, WAL lag) are
the caller's to compute.
## Properties
### buckets
```ts
buckets: BucketStats[];
```
One entry per bucket backing this table.
-60
View File
@@ -1,60 +0,0 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / MemtableStats
# Interface: MemtableStats
One in-memory memtable.
## Properties
### batches
```ts
batches: number;
```
Record batches currently buffered.
***
### bytes
```ts
bytes: number;
```
Estimated in-memory size.
***
### generation
```ts
generation: number;
```
The generation this memtable will become once sealed.
***
### indexes
```ts
indexes: string[];
```
Names of the indexes this memtable carries. An absent name is the whole
answer to "why is my fresh-tier search on that column brute-force".
***
### rows
```ts
rows: number;
```
Rows currently buffered.
@@ -1,23 +0,0 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / RefreshColumnResult
# Interface: RefreshColumnResult
## Properties
### rowsFilled
```ts
rowsFilled: number;
```
***
### version
```ts
version: number;
```
+3 -8
View File
@@ -4,16 +4,11 @@
[@lancedb/lancedb](../globals.md) / TableNamesOptions
# Interface: ~~TableNamesOptions~~
## Deprecated
Use [ListTablesOptions](ListTablesOptions.md) with [Connection.listTables](../classes/Connection.md#listtables)
instead.
# Interface: TableNamesOptions
## Properties
### ~~limit?~~
### limit?
```ts
optional limit: number;
@@ -23,7 +18,7 @@ An optional limit to the number of results to return.
***
### ~~startAfter?~~
### startAfter?
```ts
optional startAfter: string;
-2
View File
@@ -52,8 +52,6 @@ listing a storage directory.
::: lancedb.table.Branches
::: lancedb.LsmWriteSpec
## Expressions
Type-safe expression builder for filters and projections. Use these instead
-42
View File
@@ -29,48 +29,6 @@ LanceNamespace namespaceClient = LanceDbNamespaceClientBuilder.newBuilder()
.build();
```
## MemWAL LSM write path
Most table operations reach LanceDB through the `LanceNamespace` above, which is
generated from the Lance Namespace specification. The MemWAL LSM routes are not part
of that specification, so they are issued through a separate client:
```java
import com.lancedb.LanceDbRestClient;
import com.lancedb.LanceDbTableLsm;
import com.lancedb.LsmWriteSpec;
LanceDbRestClient client = LanceDbNamespaceClientBuilder.newBuilder()
.apiKey("your_lancedb_cloud_api_key")
.database("your_database_name")
.buildRestClient();
LanceDbTableLsm lsm = new LanceDbTableLsm(client, "my_table");
// Route future merge_insert upserts through the MemWAL, hash-bucketed by `id`.
lsm.setLsmWriteSpec(LsmWriteSpec.bucket("id", 16));
// ... merge_insert traffic ...
// Converge the fresh tier into the base table.
lsm.checkpointLsm();
// Inspect live per-bucket state.
lsm.getLsmStats().ifPresent(stats -> stats.buckets().forEach(bucket ->
System.out.println(bucket.shardId() + ": " + bucket.generations().size() + " L0 generations")));
client.close();
```
`maintainedIndexes` is tri-state, and the null default is the opposite of what a Java
reader usually expects:
| Value | Meaning |
| --- | --- |
| unset (null) | Maintain **every** index the MemWAL can, resolved on install |
| `Collections.emptyList()` | Maintain **none** |
| `Arrays.asList("id_idx")` | Maintain exactly those |
## Development
Build:
+1 -15
View File
@@ -8,7 +8,7 @@
<parent>
<groupId>com.lancedb</groupId>
<artifactId>lancedb-parent</artifactId>
<version>0.38.0-beta.2</version>
<version>0.37.1-beta.1</version>
<relativePath>../pom.xml</relativePath>
</parent>
@@ -33,20 +33,6 @@
<artifactId>arrow-memory-netty</artifactId>
</dependency>
<!-- Transport for the LanceDB routes outside the Lance Namespace spec.
Versions match what lance-namespace-apache-client resolves to. -->
<dependency>
<groupId>org.apache.httpcomponents.client5</groupId>
<artifactId>httpclient5</artifactId>
<version>5.2.1</version>
</dependency>
<dependency>
<groupId>com.fasterxml.jackson.core</groupId>
<artifactId>jackson-databind</artifactId>
<version>2.17.1</version>
</dependency>
<dependency>
<groupId>org.junit.jupiter</groupId>
<artifactId>junit-jupiter</artifactId>
@@ -1,194 +0,0 @@
/*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.lancedb;
import com.fasterxml.jackson.databind.JsonNode;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
import java.util.Optional;
import java.util.OptionalLong;
/**
* Live state of one bucket. A table is N buckets on one node; flattening to a single number hides
* the one hot bucket that is usually why someone opened this endpoint.
*/
public class BucketStats {
private static final String CONTEXT = "bucket stats";
private final String shardId;
private final String status;
private final long writerEpoch;
private final long manifestVersion;
private final long currentGeneration;
private final long replayAfterWalEntryPosition;
private final long walEntryPositionLastSeen;
private final List<GenerationStats> generations;
private final boolean compacting;
private final List<MemtableStats> memtables;
BucketStats(
String shardId,
String status,
long writerEpoch,
long manifestVersion,
long currentGeneration,
long replayAfterWalEntryPosition,
long walEntryPositionLastSeen,
List<GenerationStats> generations,
boolean compacting,
List<MemtableStats> memtables) {
this.shardId = shardId;
this.status = status;
this.writerEpoch = writerEpoch;
this.manifestVersion = manifestVersion;
this.currentGeneration = currentGeneration;
this.replayAfterWalEntryPosition = replayAfterWalEntryPosition;
this.walEntryPositionLastSeen = walEntryPositionLastSeen;
this.generations = Collections.unmodifiableList(generations);
this.compacting = compacting;
this.memtables = memtables == null ? null : Collections.unmodifiableList(memtables);
}
/** The shard this bucket writes. */
public String shardId() {
return shardId;
}
/** {@code "Active"} or {@code "Sealed"} (drop-table 2PC in flight). */
public String status() {
return status;
}
/** Epoch of the writer that currently owns the shard. */
public long writerEpoch() {
return writerEpoch;
}
/** Version of the shard manifest these numbers were read from. */
public long manifestVersion() {
return manifestVersion;
}
/** The generation the active memtable will become. */
public long currentGeneration() {
return currentGeneration;
}
/** WAL position replay resumes from. */
public long replayAfterWalEntryPosition() {
return replayAfterWalEntryPosition;
}
/**
* Highest WAL position the writer has seen. The difference against {@link
* #replayAfterWalEntryPosition()} is the WAL lag.
*/
public long walEntryPositionLastSeen() {
return walEntryPositionLastSeen;
}
/** Flushed L0 generations not yet merged into the base table. */
public List<GenerationStats> generations() {
return generations;
}
/**
* Whether a pass owns this bucket's compaction latch right now. Says <em>a</em> driver is
* running, not <em>whose</em>, and the latch is held from dispatch — including while the pass
* queues for a pod-wide compactor permit. Read it as "do not pile on", never as "mine is
* progressing".
*/
public boolean compacting() {
return compacting;
}
/** Oldest first, active last. Empty for a {@code "Sealed"} bucket, whose state is torn down. */
public Optional<List<MemtableStats>> memtables() {
return Optional.ofNullable(memtables);
}
/** The newest flushed generation, or empty when L0 is empty. */
OptionalLong newestGeneration() {
OptionalLong newest = OptionalLong.empty();
for (GenerationStats generation : generations) {
if (!newest.isPresent() || generation.generation() > newest.getAsLong()) {
newest = OptionalLong.of(generation.generation());
}
}
return newest;
}
/**
* How many generations at or below {@code target} are still in L0.
*
* <p>A count, not a boolean: one pass drains a bounded prefix rather than the whole target set,
* so a boolean would read as "no progress" for every pass but the last. Compaction drains
* oldest-first, so this decreases monotonically.
*/
long outstandingGenerations(long target) {
long count = 0;
for (GenerationStats generation : generations) {
if (generation.generation() <= target) {
count++;
}
}
return count;
}
static BucketStats fromJson(JsonNode node) {
JsonFields.requiredObject(node, CONTEXT);
List<GenerationStats> generations = new ArrayList<GenerationStats>();
for (JsonNode generation : JsonFields.requiredArray(node, "generations", CONTEXT)) {
generations.add(GenerationStats.fromJson(generation));
}
JsonNode memtablesNode = JsonFields.optionalArray(node, "memtables", CONTEXT);
List<MemtableStats> memtables = null;
if (memtablesNode != null) {
memtables = new ArrayList<MemtableStats>();
for (JsonNode memtable : memtablesNode) {
memtables.add(MemtableStats.fromJson(memtable));
}
}
return new BucketStats(
JsonFields.requiredText(node, "shard_id", CONTEXT),
JsonFields.requiredText(node, "status", CONTEXT),
JsonFields.requiredLong(node, "writer_epoch", CONTEXT),
JsonFields.requiredLong(node, "manifest_version", CONTEXT),
JsonFields.requiredLong(node, "current_generation", CONTEXT),
JsonFields.requiredLong(node, "replay_after_wal_entry_position", CONTEXT),
JsonFields.requiredLong(node, "wal_entry_position_last_seen", CONTEXT),
generations,
JsonFields.requiredBoolean(node, "compacting", CONTEXT),
memtables);
}
@Override
public String toString() {
return "BucketStats{shardId="
+ shardId
+ ", status="
+ status
+ ", currentGeneration="
+ currentGeneration
+ ", generations="
+ generations
+ ", compacting="
+ compacting
+ "}";
}
}
@@ -1,64 +0,0 @@
/*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.lancedb;
import com.fasterxml.jackson.databind.JsonNode;
import java.util.OptionalLong;
/** One flushed L0 generation. */
public class GenerationStats {
private static final String CONTEXT = "generation stats";
private final long generation;
private final long bytes;
private final Long rows;
GenerationStats(long generation, long bytes, Long rows) {
this.generation = generation;
this.bytes = bytes;
this.rows = rows;
}
/** The generation number. Increases as memtables are sealed into L0. */
public long generation() {
return generation;
}
/** On-disk size of the generation. */
public long bytes() {
return bytes;
}
/**
* Rows in this generation, present only when {@code includeGenerationRows} was requested. Off by
* default because each count opens an uncached Lance dataset.
*/
public OptionalLong rows() {
return rows == null ? OptionalLong.empty() : OptionalLong.of(rows);
}
static GenerationStats fromJson(JsonNode node) {
JsonFields.requiredObject(node, CONTEXT);
return new GenerationStats(
JsonFields.requiredLong(node, "generation", CONTEXT),
JsonFields.requiredLong(node, "bytes", CONTEXT),
JsonFields.optionalLong(node, "rows", CONTEXT));
}
@Override
public String toString() {
return "GenerationStats{generation=" + generation + ", bytes=" + bytes + ", rows=" + rows + "}";
}
}
@@ -1,109 +0,0 @@
/*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.lancedb;
import com.fasterxml.jackson.databind.JsonNode;
/**
* Strict readers for decoding LanceDB JSON responses.
*
* <p>Every reader fails closed: a missing, null, or wrong-typed field throws rather than
* defaulting. That mirrors the serde decoding the Rust client applies to the same payloads in
* {@code rust/lancedb/src/table/lsm_stats.rs}, where a required field has no default and a
* malformed response is an error rather than a zero.
*
* <p>The alternative — Jackson's {@code path()}, which yields a missing node that reads as an empty
* array or a zero — is unsafe here because {@link LanceDbTableLsm#checkpointLsm()} decides
* convergence from these numbers. A defaulted {@code generations} array is indistinguishable from a
* drained one, so a malformed response would report a checkpoint that never happened.
*/
final class JsonFields {
private JsonFields() {}
/** The node itself, once confirmed to be a JSON object. */
static JsonNode requiredObject(JsonNode node, String context) {
if (node == null || !node.isObject()) {
throw new IllegalStateException(context + " is not a JSON object: " + node);
}
return node;
}
static String requiredText(JsonNode owner, String field, String context) {
JsonNode value = required(owner, field, context);
if (!value.isTextual()) {
throw new IllegalStateException(fieldIs(context, field, "a string", value));
}
return value.asText();
}
static long requiredLong(JsonNode owner, String field, String context) {
JsonNode value = required(owner, field, context);
if (!value.isIntegralNumber()) {
throw new IllegalStateException(fieldIs(context, field, "an integer", value));
}
return value.asLong();
}
static boolean requiredBoolean(JsonNode owner, String field, String context) {
JsonNode value = required(owner, field, context);
if (!value.isBoolean()) {
throw new IllegalStateException(fieldIs(context, field, "a boolean", value));
}
return value.asBoolean();
}
static JsonNode requiredArray(JsonNode owner, String field, String context) {
JsonNode value = required(owner, field, context);
if (!value.isArray()) {
throw new IllegalStateException(fieldIs(context, field, "an array", value));
}
return value;
}
/** Null when the field is absent or JSON null, mirroring a serde {@code Option}. */
static Long optionalLong(JsonNode owner, String field, String context) {
JsonNode value = owner.get(field);
if (value == null || value.isNull()) {
return null;
}
if (!value.isIntegralNumber()) {
throw new IllegalStateException(fieldIs(context, field, "an integer", value));
}
return value.asLong();
}
/** Null when the field is absent or JSON null, mirroring a serde {@code Option}. */
static JsonNode optionalArray(JsonNode owner, String field, String context) {
JsonNode value = owner.get(field);
if (value == null || value.isNull()) {
return null;
}
if (!value.isArray()) {
throw new IllegalStateException(fieldIs(context, field, "an array", value));
}
return value;
}
private static JsonNode required(JsonNode owner, String field, String context) {
JsonNode value = owner.get(field);
if (value == null || value.isNull()) {
throw new IllegalStateException(context + " is missing required field '" + field + "'");
}
return value;
}
private static String fieldIs(String context, String field, String expected, JsonNode value) {
return context + " field '" + field + "' is not " + expected + ": " + value;
}
}
@@ -136,48 +136,29 @@ public class LanceDbNamespaceClientBuilder {
* @throws IllegalStateException if required parameters are missing
*/
public LanceNamespace build() {
validate();
// Build configuration map
Map<String, String> config = new HashMap<>(additionalConfig);
config.put("header.x-lancedb-database", database);
config.put("header.x-api-key", apiKey);
config.put("uri", resolveUri());
return LanceNamespace.connect("rest", config, null);
}
/**
* Build a {@link LanceDbRestClient} for the same endpoint.
*
* <p>Needed only for LanceDB routes that the Lance Namespace specification does not cover — the
* MemWAL LSM write path, reached through {@link LanceDbTableLsm}. Every other table operation
* belongs on the {@link LanceNamespace} from {@link #build()}.
*
* <p>The returned client owns an HTTP connection pool; close it when you are done with it.
*
* @return A configured LanceDbRestClient
* @throws IllegalStateException if required parameters are missing
*/
public LanceDbRestClient buildRestClient() {
validate();
return new LanceDbRestClient(resolveUri(), apiKey, database);
}
private void validate() {
// Validate required fields
if (apiKey == null) {
throw new IllegalStateException("API key is required");
}
if (database == null) {
throw new IllegalStateException("Database is required");
}
}
/** The custom endpoint when set, else the LanceDB Cloud URL for this database and region. */
private String resolveUri() {
// Build configuration map
Map<String, String> config = new HashMap<>(additionalConfig);
config.put("header.x-lancedb-database", database);
config.put("header.x-api-key", apiKey);
// Determine base URL
String uri;
if (endpoint.isPresent()) {
return endpoint.get();
uri = endpoint.get();
} else {
String effectiveRegion = region.orElse(DEFAULT_REGION);
uri = String.format(CLOUD_URL_PATTERN, database, effectiveRegion);
}
return String.format(CLOUD_URL_PATTERN, database, region.orElse(DEFAULT_REGION));
config.put("uri", uri);
return LanceNamespace.connect("rest", config, null);
}
}
@@ -1,119 +0,0 @@
/*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.lancedb;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.apache.hc.client5.http.classic.methods.HttpPost;
import org.apache.hc.client5.http.impl.classic.CloseableHttpClient;
import org.apache.hc.client5.http.impl.classic.HttpClients;
import org.apache.hc.core5.http.ContentType;
import org.apache.hc.core5.http.io.entity.EntityUtils;
import org.apache.hc.core5.http.io.entity.StringEntity;
import java.io.Closeable;
import java.io.IOException;
import java.io.UncheckedIOException;
/**
* Minimal HTTP client for LanceDB Cloud and Enterprise routes that the Lance Namespace
* specification does not cover.
*
* <p>Most table operations reach LanceDB through {@link org.lance.namespace.LanceNamespace}, which
* is generated from the namespace spec. A handful of routes — the MemWAL LSM write path in
* particular — are served by the same endpoint but are not part of that spec, so they are issued
* directly here. See {@link LanceDbTableLsm}.
*
* <p>Obtain one from {@link LanceDbNamespaceClientBuilder#buildRestClient()}.
*/
public class LanceDbRestClient implements Closeable {
private static final ObjectMapper MAPPER = new ObjectMapper();
private final String baseUri;
private final String apiKey;
private final String database;
private final CloseableHttpClient http;
LanceDbRestClient(String baseUri, String apiKey, String database) {
this.baseUri = baseUri.endsWith("/") ? baseUri.substring(0, baseUri.length() - 1) : baseUri;
this.apiKey = apiKey;
this.database = database;
// Automatic retries off, deliberately. The default strategy retries 429 and 503 —
// exactly the two statuses LanceDbTableLsm.checkpointLsm() acts on — which would
// silently double its explicit retry budget and would also retry compact_lsm in
// place, where the loop is designed to fall through to a fresh stats poll instead.
// The checkpoint loop owns the 421/429/503 transitions; the transport must not.
this.http = HttpClients.custom().disableAutomaticRetries().build();
}
/**
* POST {@code path}, sending {@code body} as JSON when it is non-null.
*
* @param path Absolute request path, beginning with {@code /}.
* @param body Object to serialize as the request body, or null to send no body.
* @return The parsed response body, or null when the response carried no content.
* @throws HttpException if the server returned a non-2xx status.
*/
public JsonNode post(String path, Object body) {
HttpPost request = new HttpPost(baseUri + path);
request.setHeader("x-api-key", apiKey);
request.setHeader("x-lancedb-database", database);
try {
if (body != null) {
request.setEntity(
new StringEntity(MAPPER.writeValueAsString(body), ContentType.APPLICATION_JSON));
}
return http.execute(
request,
response -> {
String text =
response.getEntity() == null ? "" : EntityUtils.toString(response.getEntity());
int status = response.getCode();
if (status < 200 || status >= 300) {
throw new HttpException(status, "LanceDB request to " + path + " failed: " + text);
}
return text.isEmpty() ? null : MAPPER.readTree(text);
});
} catch (IOException e) {
throw new UncheckedIOException("LanceDB request to " + path + " failed", e);
}
}
@Override
public void close() throws IOException {
http.close();
}
/**
* A non-2xx response.
*
* <p>The status is exposed because callers act on it: {@link LanceDbTableLsm#checkpointLsm()}
* treats 429 and 503 as retryable and 421 as a lost node claim.
*/
public static class HttpException extends RuntimeException {
private static final long serialVersionUID = 1L;
private final int statusCode;
public HttpException(int statusCode, String message) {
super(message);
this.statusCode = statusCode;
}
/** The HTTP status the failed response carried. */
public int statusCode() {
return statusCode;
}
}
}
@@ -1,394 +0,0 @@
/*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.lancedb;
import com.fasterxml.jackson.databind.JsonNode;
import java.util.HashMap;
import java.util.LinkedHashMap;
import java.util.Map;
import java.util.Optional;
import java.util.OptionalLong;
/**
* The MemWAL LSM write path for one LanceDB Cloud or Enterprise table.
*
* <p>Installing an {@link LsmWriteSpec} routes {@code mergeInsert} upserts through Lance's MemWAL —
* an LSM-style append — instead of the standard merge path. Rows land in an in-memory memtable,
* seal into L0 generations, and are merged into the base table by compaction.
*
* <p>These routes are not part of the Lance Namespace specification, so they are issued directly
* rather than through {@link org.lance.namespace.LanceNamespace}.
*
* <pre>{@code
* LanceDbRestClient client = LanceDbNamespaceClientBuilder.newBuilder()
* .apiKey("your_lancedb_cloud_api_key")
* .database("your_database_name")
* .buildRestClient();
*
* LanceDbTableLsm lsm = new LanceDbTableLsm(client, "my_table");
* lsm.setLsmWriteSpec(LsmWriteSpec.bucket("id", 16));
* // ... merge_insert traffic ...
* lsm.checkpointLsm();
* }</pre>
*/
public class LanceDbTableLsm {
/**
* Interval between {@code get_lsm_stats} polls during a checkpoint. One interval is roughly one
* compaction pass, the granularity at which the answer can change.
*/
private static final long POLL_INTERVAL_MS = 5_000L;
/**
* Cap on re-issues from {@code flushLsm} after a 421, so a crash-looping node cannot turn flush →
* compact → 421 → flush into a spin.
*
* <p>Deliberately not shared with {@link #MAX_RETRIES}: a claim that keeps evaporating is a
* broken node, while contention is routine and wants a real budget.
*/
private static final int MAX_REISSUES = 3;
/**
* Retryable faults tolerated on a <em>single</em> request, reset on every success — scattered
* contention across a long checkpoint must not accumulate toward a cap.
*/
private static final int MAX_RETRIES = 8;
private static final long RETRY_BACKOFF_BASE_MS = 100L;
private static final long RETRY_BACKOFF_MAX_MS = 5_000L;
private final LanceDbRestClient client;
private final String tableIdentifier;
/**
* Bind the LSM routes for one table.
*
* @param client Transport for the LanceDB endpoint.
* @param tableIdentifier The table's full identifier, {@code $}-delimited when it sits inside a
* namespace, such as {@code analytics$events}.
*/
public LanceDbTableLsm(LanceDbRestClient client, String tableIdentifier) {
if (client == null) {
throw new IllegalArgumentException("Client cannot be null");
}
if (tableIdentifier == null || tableIdentifier.trim().isEmpty()) {
throw new IllegalArgumentException("Table identifier cannot be null or empty");
}
this.client = client;
this.tableIdentifier = tableIdentifier;
}
/**
* Install an {@link LsmWriteSpec} on this table, selecting the MemWAL LSM write path for future
* {@code mergeInsert} calls.
*
* <p>All variants require the table to have an unenforced primary key; bucket sharding
* additionally requires it to be the single column being bucketed.
*/
public void setLsmWriteSpec(LsmWriteSpec spec) {
if (spec == null) {
throw new IllegalArgumentException("Spec cannot be null");
}
client.post(route("set_lsm_write_spec"), spec.toRequestBody());
}
/**
* Remove the {@link LsmWriteSpec} from this table, reverting to the standard {@code mergeInsert}
* write path.
*
* <p>Errors if no spec is currently set.
*/
public void unsetLsmWriteSpec() {
client.post(route("unset_lsm_write_spec"), null);
}
/**
* Read the {@link LsmWriteSpec} currently installed on this table.
*
* <p>Empty when the LSM write path is not enabled. The returned spec mirrors what was installed,
* except that {@link LsmWriteSpec#maintainedIndexes()} always reports the concrete list resolved
* when the spec was set — a null selection never round-trips.
*/
public Optional<LsmWriteSpec> getLsmWriteSpec() {
JsonNode response = client.post(route("get_lsm_write_spec"), null);
if (response == null || !response.hasNonNull("lsm_write_spec")) {
return Optional.empty();
}
return Optional.of(LsmWriteSpec.fromJson(response.get("lsm_write_spec")));
}
/**
* Seal every bucket's active memtable into a new L0 generation.
*
* <p>Returns once the seal is committed. Sealing an empty memtable is a no-op, so this is safe to
* call repeatedly.
*/
public void flushLsm() {
client.post(route("flush_lsm"), null);
}
/**
* Trigger a background L0 → base compaction pass per bucket.
*
* <p>Returns once the passes are <em>dispatched</em>, not once they finish — watch {@link
* #getLsmStats}, or use {@link #checkpointLsm} to wait for convergence.
*/
public void compactLsm() {
client.post(route("compact_lsm"), null);
}
/**
* Read live per-bucket LSM state.
*
* <p>Answers "how far behind is my fresh tier", "which bucket is hot", and "why is my fresh-tier
* vector search brute-force". Mutates no table state.
*
* <p>Empty only when the LSM write path is not enabled — that is, when the server sends an absent
* or null {@code lsm_stats}. A stats object that is present is decoded strictly, and a malformed
* one throws rather than decoding to something empty, because {@link #checkpointLsm} reads
* convergence out of these numbers and cannot tell a defaulted array from a drained one.
*
* @param includeGenerationRows Also count rows per L0 generation. Off by default because each
* count opens an uncached Lance dataset.
* @throws IllegalStateException if the response is absent or does not decode.
*/
public Optional<LsmStats> getLsmStats(boolean includeGenerationRows) {
Map<String, Object> body = new LinkedHashMap<String, Object>();
body.put("include_generation_rows", includeGenerationRows);
JsonNode response = client.post(route("get_lsm_stats"), body);
if (response == null) {
throw new IllegalStateException("get_lsm_stats returned an empty response body");
}
JsonNode stats = response.get("lsm_stats");
if (stats == null || stats.isNull()) {
return Optional.empty();
}
return Optional.of(LsmStats.fromJson(stats));
}
/** Equivalent to {@code getLsmStats(false)}. */
public Optional<LsmStats> getLsmStats() {
return getLsmStats(false);
}
/**
* Converge this table's LSM write path into its base table.
*
* <p>Seals once, fixes a target watermark from the resulting L0, then triggers compaction and
* polls until that L0 is gone. The target set is fixed at the start, so generations created
* <em>during</em> the checkpoint are ignored — that is what lets it terminate under write load,
* and what makes it best-effort: it converges the fresh tier as of some instant. Idempotent,
* abandonable at any point, safe on a cadence.
*
* <p>The loop runs here, not on the server: {@link #compactLsm} dispatches a pass and returns, so
* nothing holds a socket and a client can vanish mid-operation with nothing to reconcile.
* Completion is read from generation numbers in the shard manifest — durable state, unlike a
* count in a compact response, which a concurrent write invalidates.
*
* <p>No liveness bound — the caller owns the deadline. The compactor pool is shared across
* tables, so a checkpoint queued behind unrelated work looks exactly like one that is merging.
*/
public void checkpointLsm() {
for (int reissue = 0; reissue <= MAX_REISSUES; reissue++) {
// The seal turns everything written before this call into a generation, so the
// watermark has to be read after it. Idempotent: sealing an empty memtable is a
// no-op, so a re-issue does not churn empty generations.
if (issueVoid(this::flushLsm)) {
backoff(reissue);
continue;
}
Attempt<Optional<LsmStats>> stats = issue(() -> getLsmStats(false));
if (stats.lostClaim) {
backoff(reissue);
continue;
}
if (!stats.value.isPresent()) {
// Not WAL-backed; flushLsm would have errored first but for a race.
return;
}
Map<String, Long> targets = newestGenerations(stats.value.get());
if (targets.isEmpty()) {
return;
}
if (drainToTargets(targets)) {
return;
}
backoff(reissue);
}
throw new IllegalStateException(
"checkpointLsm: the owning node kept losing its claim; re-issued from flush the maximum "
+ "number of times");
}
/**
* Trigger and poll until no bucket holds a generation at or below its target.
*
* @return true when the drain finished, false when the table needs re-claiming from flush.
*/
private boolean drainToTargets(Map<String, Long> targets) {
while (true) {
Attempt<Optional<LsmStats>> stats = issue(() -> getLsmStats(false));
if (stats.lostClaim) {
return false;
}
if (!stats.value.isPresent()) {
return true;
}
// `compacting` is the bucket's compaction latch, held from dispatch until the pass
// ends — including while it waits on a pod-wide permit. So it answers one question
// only: do not pile on. Buckets with nothing outstanding are skipped, not counted
// as idle.
long outstanding = 0;
boolean allCompacting = true;
for (BucketStats bucket : stats.value.get().buckets()) {
Long target = targets.get(bucket.shardId());
if (target == null) {
continue;
}
long remaining = bucket.outstandingGenerations(target);
if (remaining > 0) {
outstanding += remaining;
allCompacting &= bucket.compacting();
}
}
if (outstanding == 0) {
return true;
}
if (!allCompacting) {
try {
compactLsm();
} catch (LanceDbRestClient.HttpException e) {
if (isLostClaim(e)) {
return false;
}
if (!isRetryable(e)) {
throw e;
}
// A 429 here means the server could latch no bucket at all, which the poll
// above already handles. Not retried in place: the latch it would contend for
// is the one doing the work, so fall through and re-read — POLL_INTERVAL_MS is
// the backoff.
}
}
sleep(POLL_INTERVAL_MS);
}
}
/** The newest generation held by each bucket, skipping buckets holding none. */
private static Map<String, Long> newestGenerations(LsmStats stats) {
Map<String, Long> targets = new HashMap<String, Long>();
for (BucketStats bucket : stats.buckets()) {
OptionalLong newest = bucket.newestGeneration();
if (newest.isPresent()) {
targets.put(bucket.shardId(), newest.getAsLong());
}
}
return targets;
}
/**
* 429 (latch held, pool saturated, or the pod replaying its WAL) and 503 (a draining node, or a
* proxy between here and it).
*/
private static boolean isRetryable(LanceDbRestClient.HttpException e) {
return e.statusCode() == 429 || e.statusCode() == 503;
}
/**
* 421: the owning node holds no claim. Only {@code flush} re-claims and replays, so this cannot
* be retried in place — the caller has to start over.
*/
private static boolean isLostClaim(LanceDbRestClient.HttpException e) {
return e.statusCode() == 421;
}
/**
* Issue one LSM request, retrying in place while the fault is retryable.
*
* <p>The two recoverable faults have separate budgets: contention clears on its own and retries
* here against {@link #MAX_RETRIES}, while a 421 needs {@code flush} to re-claim, which only the
* caller can drive.
*
* <p>An exhausted budget propagates the last error as itself rather than a synthesized one — "429
* after nine tries" beats "checkpoint failed".
*/
private static <T> Attempt<T> issue(Call<T> call) {
int retries = 0;
while (true) {
try {
return new Attempt<T>(call.run(), false);
} catch (LanceDbRestClient.HttpException e) {
if (isLostClaim(e)) {
return new Attempt<T>(null, true);
}
if (!isRetryable(e) || retries >= MAX_RETRIES) {
throw e;
}
backoff(retries);
retries++;
}
}
}
/** {@link #issue} for a call with no return value. Returns true when the claim was lost. */
private static boolean issueVoid(Runnable call) {
return issue(
() -> {
call.run();
return Boolean.TRUE;
})
.lostClaim;
}
/** Sleep before re-issuing a retryable request. Doubles up to {@link #RETRY_BACKOFF_MAX_MS}. */
private static void backoff(int attempt) {
long delay = RETRY_BACKOFF_BASE_MS << Math.min(attempt, 8);
sleep(Math.min(delay, RETRY_BACKOFF_MAX_MS));
}
private static void sleep(long millis) {
try {
Thread.sleep(millis);
} catch (InterruptedException e) {
Thread.currentThread().interrupt();
throw new IllegalStateException("Interrupted while waiting on the LSM checkpoint", e);
}
}
private String route(String operation) {
return "/v1/table/" + tableIdentifier + "/" + operation + "/";
}
/** What one LSM request produced: its value, or word that the owning node holds no claim. */
private static final class Attempt<T> {
private final T value;
private final boolean lostClaim;
private Attempt(T value, boolean lostClaim) {
this.value = value;
this.lostClaim = lostClaim;
}
}
@FunctionalInterface
private interface Call<T> {
T run();
}
}
@@ -1,56 +0,0 @@
/*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.lancedb;
import com.fasterxml.jackson.databind.JsonNode;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
/**
* Live per-bucket LSM state, as returned by {@link LanceDbTableLsm#getLsmStats()}.
*
* <p>Nothing here is derived: sums and differences (total L0 bytes, WAL lag) are the caller's to
* compute. There is no "LSM is off" shape — that case is an empty {@link java.util.Optional},
* because a stats object of zeros would read as measurements.
*/
public class LsmStats {
private static final String CONTEXT = "lsm stats";
private final List<BucketStats> buckets;
LsmStats(List<BucketStats> buckets) {
this.buckets = Collections.unmodifiableList(buckets);
}
/** One entry per bucket. */
public List<BucketStats> buckets() {
return buckets;
}
static LsmStats fromJson(JsonNode node) {
JsonFields.requiredObject(node, CONTEXT);
List<BucketStats> buckets = new ArrayList<BucketStats>();
for (JsonNode bucket : JsonFields.requiredArray(node, "buckets", CONTEXT)) {
buckets.add(BucketStats.fromJson(bucket));
}
return new LsmStats(buckets);
}
@Override
public String toString() {
return "LsmStats{buckets=" + buckets + "}";
}
}
@@ -1,260 +0,0 @@
/*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.lancedb;
import com.fasterxml.jackson.databind.JsonNode;
import java.util.ArrayList;
import java.util.Collections;
import java.util.HashMap;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
/**
* Specification selecting Lance's MemWAL LSM-style write path for {@code mergeInsert}.
*
* <p>Construct via {@link #bucket}, {@link #identity}, or {@link #unsharded}, then optionally chain
* {@link #withMaintainedIndexes} and {@link #withWriterConfigDefaults}. Install it with {@link
* LanceDbTableLsm#setLsmWriteSpec} and remove it with {@link LanceDbTableLsm#unsetLsmWriteSpec}.
*
* <p>This is deliberately not {@code org.lance.memwal.InitializeMemWalParams}. That type is Lance's
* own, and its maintained-index default is the opposite of this one: it defaults to maintaining
* <em>nothing</em>, while a fresh spec here maintains <em>every</em> index. It also cannot express
* the null that asks the server to resolve the set.
*/
public class LsmWriteSpec {
/** How writes are routed to MemWAL shards. */
public enum Sharding {
/** Hash-bucket writes by a scalar column. */
BUCKET("bucket"),
/** Shard by the raw value of a scalar column. */
IDENTITY("identity"),
/** Route every write to a single shard. */
UNSHARDED("unsharded");
private final String wireName;
Sharding(String wireName) {
this.wireName = wireName;
}
String wireName() {
return wireName;
}
static Sharding fromWireName(String name) {
for (Sharding s : values()) {
if (s.wireName.equals(name)) {
return s;
}
}
throw new IllegalArgumentException("Unknown sharding mode: " + name);
}
}
private final Sharding sharding;
private final String column;
private final Integer numBuckets;
private final List<String> maintainedIndexes;
private final Map<String, String> writerConfigDefaults;
private LsmWriteSpec(
Sharding sharding,
String column,
Integer numBuckets,
List<String> maintainedIndexes,
Map<String, String> writerConfigDefaults) {
this.sharding = sharding;
this.column = column;
this.numBuckets = numBuckets;
this.maintainedIndexes = maintainedIndexes;
this.writerConfigDefaults = writerConfigDefaults;
}
/**
* Hash-bucket sharding by a scalar column, maintaining every index on the table.
*
* <p>Iceberg-compatible Murmur3-x86-32 (seed 0) is used, so each row's {@code bucket(column,
* numBuckets)} value is stable across processes.
*
* @param column A non-nested column with a supported scalar type.
* @param numBuckets The number of buckets, in {@code [1, 1024]}.
*/
public static LsmWriteSpec bucket(String column, int numBuckets) {
if (column == null || column.trim().isEmpty()) {
throw new IllegalArgumentException("Column cannot be null or empty");
}
return new LsmWriteSpec(
Sharding.BUCKET, column, numBuckets, null, new HashMap<String, String>());
}
/**
* Identity sharding — shard by the raw value of {@code column} — maintaining every index on the
* table.
*
* <p>{@code column} must be a deterministic function of the unenforced primary key: every row
* with a given primary key must always produce the same {@code column} value, or upserts of that
* key can land in different shards and a stale version can win.
*/
public static LsmWriteSpec identity(String column) {
if (column == null || column.trim().isEmpty()) {
throw new IllegalArgumentException("Column cannot be null or empty");
}
return new LsmWriteSpec(Sharding.IDENTITY, column, null, null, new HashMap<String, String>());
}
/** No sharding — every write goes to a single MemWAL shard — maintaining every index. */
public static LsmWriteSpec unsharded() {
return new LsmWriteSpec(Sharding.UNSHARDED, null, null, null, new HashMap<String, String>());
}
/**
* Set the indexes the MemWAL keeps up to date as rows are appended.
*
* <p>Pass {@code null} — the default for a fresh spec — to maintain every index the MemWAL can,
* resolved when the spec is installed. That is a snapshot: indexes created later are not
* maintained until the spec is unset and set again. Pass an empty list to maintain none.
*
* <p>Note that {@code null} and the empty list mean opposite things here.
*/
public LsmWriteSpec withMaintainedIndexes(List<String> maintainedIndexes) {
return new LsmWriteSpec(
sharding,
column,
numBuckets,
maintainedIndexes == null ? null : new ArrayList<String>(maintainedIndexes),
writerConfigDefaults);
}
/**
* Set default {@code ShardWriter} configuration recorded in the MemWAL index.
*
* <p>A sparse override map — only the keys you set are recorded. Recognized keys include {@code
* durable_write}, {@code max_wal_buffer_size}, {@code max_memtable_size}, {@code
* max_memtable_rows}, {@code max_memtable_batches}, {@code manifest_scan_batch_size}, {@code
* max_unflushed_memtable_bytes}, and {@code enable_memtable}. Duration knobs carry an {@code _ms}
* suffix, such as {@code max_wal_flush_interval_ms}.
*/
public LsmWriteSpec withWriterConfigDefaults(Map<String, String> writerConfigDefaults) {
if (writerConfigDefaults == null) {
throw new IllegalArgumentException("writerConfigDefaults cannot be null");
}
return new LsmWriteSpec(
sharding,
column,
numBuckets,
maintainedIndexes,
new HashMap<String, String>(writerConfigDefaults));
}
/** How writes are routed to shards. */
public Sharding sharding() {
return sharding;
}
/** The sharding column for {@link Sharding#BUCKET} and {@link Sharding#IDENTITY}, else null. */
public String column() {
return column;
}
/** The bucket count for {@link Sharding#BUCKET}, else null. */
public Integer numBuckets() {
return numBuckets;
}
/**
* The indexes the MemWAL maintains, or null to have the server resolve every maintainable index
* on install. An empty list means none.
*/
public List<String> maintainedIndexes() {
return maintainedIndexes == null ? null : Collections.unmodifiableList(maintainedIndexes);
}
/** Default {@code ShardWriter} configuration recorded in the MemWAL index. */
public Map<String, String> writerConfigDefaults() {
return Collections.unmodifiableMap(writerConfigDefaults);
}
/** Render this spec as the {@code set_lsm_write_spec} request body. */
Map<String, Object> toRequestBody() {
Map<String, Object> shardingBody = new LinkedHashMap<String, Object>();
shardingBody.put("mode", sharding.wireName());
if (column != null) {
shardingBody.put("column", column);
}
if (numBuckets != null) {
shardingBody.put("num_buckets", numBuckets);
}
Map<String, Object> body = new LinkedHashMap<String, Object>();
body.put("sharding", shardingBody);
// Null is meaningful: it asks the server to resolve every maintainable index.
body.put("maintained_indexes", maintainedIndexes);
body.put("writer_config_defaults", writerConfigDefaults);
return body;
}
/**
* Rebuild a spec from a {@code get_lsm_write_spec} response body.
*
* <p>The server always reports a concrete maintained-index list, so a null selection never
* round-trips.
*/
static LsmWriteSpec fromJson(JsonNode node) {
JsonNode shardingNode = node.get("sharding");
if (shardingNode == null || shardingNode.get("mode") == null) {
throw new IllegalStateException("get_lsm_write_spec response has no sharding mode");
}
Sharding sharding = Sharding.fromWireName(shardingNode.get("mode").asText());
String column = shardingNode.hasNonNull("column") ? shardingNode.get("column").asText() : null;
Integer numBuckets =
shardingNode.hasNonNull("num_buckets") ? shardingNode.get("num_buckets").asInt() : null;
List<String> maintainedIndexes = new ArrayList<String>();
JsonNode indexesNode = node.get("maintained_indexes");
if (indexesNode != null && indexesNode.isArray()) {
for (JsonNode index : indexesNode) {
maintainedIndexes.add(index.asText());
}
}
Map<String, String> defaults = new HashMap<String, String>();
JsonNode defaultsNode = node.get("writer_config_defaults");
if (defaultsNode != null && defaultsNode.isObject()) {
defaultsNode
.fieldNames()
.forEachRemaining(name -> defaults.put(name, defaultsNode.get(name).asText()));
}
return new LsmWriteSpec(sharding, column, numBuckets, maintainedIndexes, defaults);
}
@Override
public String toString() {
return "LsmWriteSpec{sharding="
+ sharding
+ ", column="
+ column
+ ", numBuckets="
+ numBuckets
+ ", maintainedIndexes="
+ maintainedIndexes
+ ", writerConfigDefaults="
+ writerConfigDefaults
+ "}";
}
}
@@ -1,99 +0,0 @@
/*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.lancedb;
import com.fasterxml.jackson.databind.JsonNode;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
/** One in-memory memtable. */
public class MemtableStats {
private static final String CONTEXT = "memtable stats";
private final long generation;
private final long rows;
private final long bytes;
private final long batches;
private final List<String> indexes;
MemtableStats(long generation, long rows, long bytes, long batches, List<String> indexes) {
this.generation = generation;
this.rows = rows;
this.bytes = bytes;
this.batches = batches;
this.indexes = Collections.unmodifiableList(indexes);
}
/** The generation this memtable will become once sealed. */
public long generation() {
return generation;
}
/** Rows currently buffered. */
public long rows() {
return rows;
}
/** Estimated in-memory size. */
public long bytes() {
return bytes;
}
/** Record batches currently buffered. */
public long batches() {
return batches;
}
/**
* Names of the indexes this memtable carries. An absent name is the whole answer to "why is my
* fresh-tier search on that column brute-force".
*/
public List<String> indexes() {
return indexes;
}
static MemtableStats fromJson(JsonNode node) {
JsonFields.requiredObject(node, CONTEXT);
List<String> indexes = new ArrayList<String>();
for (JsonNode index : JsonFields.requiredArray(node, "indexes", CONTEXT)) {
if (!index.isTextual()) {
throw new IllegalStateException(CONTEXT + " has a non-string index name: " + index);
}
indexes.add(index.asText());
}
return new MemtableStats(
JsonFields.requiredLong(node, "generation", CONTEXT),
JsonFields.requiredLong(node, "rows", CONTEXT),
JsonFields.requiredLong(node, "bytes", CONTEXT),
JsonFields.requiredLong(node, "batches", CONTEXT),
indexes);
}
@Override
public String toString() {
return "MemtableStats{generation="
+ generation
+ ", rows="
+ rows
+ ", bytes="
+ bytes
+ ", batches="
+ batches
+ ", indexes="
+ indexes
+ "}";
}
}
@@ -1,570 +0,0 @@
/*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.lancedb;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.sun.net.httpserver.HttpServer;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.io.InputStream;
import java.io.UncheckedIOException;
import java.net.InetSocketAddress;
import java.nio.charset.StandardCharsets;
import java.util.ArrayDeque;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.Collections;
import java.util.Deque;
import java.util.HashMap;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.concurrent.ConcurrentHashMap;
import static org.junit.jupiter.api.Assertions.*;
/**
* Unit tests for the MemWAL LSM routes, run against a scripted local HTTP server.
*
* <p>The wire assertions mirror the Rust mocked-endpoint tests in {@code
* rust/lancedb/src/remote/table.rs}, which are the contract these routes have to match.
*/
public class LanceDbTableLsmTest {
private static final ObjectMapper MAPPER = new ObjectMapper();
private HttpServer server;
private LanceDbRestClient client;
private LanceDbTableLsm lsm;
private final List<String> requestPaths = Collections.synchronizedList(new ArrayList<String>());
private final List<String> requestBodies = Collections.synchronizedList(new ArrayList<String>());
private final Map<String, Deque<Reply>> replies = new ConcurrentHashMap<String, Deque<Reply>>();
@BeforeEach
public void setUp() throws IOException {
start();
}
/** Tear down and restart the scripted server, for a test that scripts several exchanges. */
private void setUpFresh() {
try {
client.close();
server.stop(0);
requestPaths.clear();
requestBodies.clear();
replies.clear();
start();
} catch (IOException e) {
throw new UncheckedIOException(e);
}
}
private void start() throws IOException {
server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0);
server.createContext(
"/",
exchange -> {
String path = exchange.getRequestURI().getPath();
requestPaths.add(path);
requestBodies.add(readAll(exchange.getRequestBody()));
Reply reply = nextReply(path);
byte[] out = reply.body.getBytes(StandardCharsets.UTF_8);
exchange.sendResponseHeaders(reply.status, out.length == 0 ? -1 : out.length);
if (out.length > 0) {
exchange.getResponseBody().write(out);
}
exchange.close();
});
server.start();
client =
LanceDbNamespaceClientBuilder.newBuilder()
.apiKey("test-key")
.database("test-db")
.endpoint("http://127.0.0.1:" + server.getAddress().getPort())
.buildRestClient();
lsm = new LanceDbTableLsm(client, "my_table");
}
@AfterEach
public void tearDown() throws IOException {
client.close();
server.stop(0);
}
// ===========================================================================
// set / unset / get spec
// ===========================================================================
@Test
public void testSetLsmWriteSpecUnsharded() throws Exception {
enqueue("set_lsm_write_spec", 200, "");
lsm.setLsmWriteSpec(LsmWriteSpec.unsharded());
assertEquals("/v1/table/my_table/set_lsm_write_spec/", requestPaths.get(0));
JsonNode body = MAPPER.readTree(requestBodies.get(0));
assertEquals("unsharded", body.get("sharding").get("mode").asText());
assertFalse(body.get("sharding").has("column"));
assertFalse(body.get("sharding").has("num_buckets"));
}
@Test
public void testSetLsmWriteSpecBucket() throws Exception {
enqueue("set_lsm_write_spec", 200, "");
lsm.setLsmWriteSpec(
LsmWriteSpec.bucket("id", 16).withMaintainedIndexes(Arrays.asList("id_idx")));
JsonNode body = MAPPER.readTree(requestBodies.get(0));
assertEquals("bucket", body.get("sharding").get("mode").asText());
assertEquals("id", body.get("sharding").get("column").asText());
assertEquals(16, body.get("sharding").get("num_buckets").asInt());
assertEquals(1, body.get("maintained_indexes").size());
assertEquals("id_idx", body.get("maintained_indexes").get(0).asText());
}
@Test
public void testSetLsmWriteSpecIdentity() throws Exception {
enqueue("set_lsm_write_spec", 200, "");
lsm.setLsmWriteSpec(LsmWriteSpec.identity("tenant"));
JsonNode body = MAPPER.readTree(requestBodies.get(0));
assertEquals("identity", body.get("sharding").get("mode").asText());
assertEquals("tenant", body.get("sharding").get("column").asText());
assertFalse(body.get("sharding").has("num_buckets"));
}
/**
* The tri-state that motivated a LanceDB-owned spec type: a null selection asks the server to
* resolve every maintainable index, while an empty list asks for none. They must not collapse.
*/
@Test
public void testMaintainedIndexesNullAndEmptyAreDistinctOnTheWire() throws Exception {
enqueue("set_lsm_write_spec", 200, "");
lsm.setLsmWriteSpec(LsmWriteSpec.unsharded());
JsonNode fresh = MAPPER.readTree(requestBodies.get(0));
assertTrue(fresh.has("maintained_indexes"), "the key must be present");
assertTrue(fresh.get("maintained_indexes").isNull(), "a fresh spec sends null, not []");
lsm.setLsmWriteSpec(
LsmWriteSpec.unsharded().withMaintainedIndexes(Collections.<String>emptyList()));
JsonNode none = MAPPER.readTree(requestBodies.get(1));
assertTrue(none.get("maintained_indexes").isArray());
assertEquals(0, none.get("maintained_indexes").size());
}
@Test
public void testSetLsmWriteSpecWriterConfigDefaults() throws Exception {
enqueue("set_lsm_write_spec", 200, "");
Map<String, String> defaults = new HashMap<String, String>();
defaults.put("max_memtable_rows", "50000");
lsm.setLsmWriteSpec(LsmWriteSpec.unsharded().withWriterConfigDefaults(defaults));
JsonNode body = MAPPER.readTree(requestBodies.get(0));
assertEquals("50000", body.get("writer_config_defaults").get("max_memtable_rows").asText());
}
@Test
public void testUnsetLsmWriteSpec() {
enqueue("unset_lsm_write_spec", 200, "");
lsm.unsetLsmWriteSpec();
assertEquals("/v1/table/my_table/unset_lsm_write_spec/", requestPaths.get(0));
assertEquals("", requestBodies.get(0));
}
@Test
public void testGetLsmWriteSpec() {
enqueue(
"get_lsm_write_spec",
200,
"{\"lsm_write_spec\":{\"sharding\":{\"mode\":\"bucket\",\"column\":\"id\","
+ "\"num_buckets\":16},\"maintained_indexes\":[\"id_idx\"],"
+ "\"writer_config_defaults\":{\"durable_write\":\"true\"}}}");
Optional<LsmWriteSpec> spec = lsm.getLsmWriteSpec();
assertTrue(spec.isPresent());
assertEquals(LsmWriteSpec.Sharding.BUCKET, spec.get().sharding());
assertEquals("id", spec.get().column());
assertEquals(Integer.valueOf(16), spec.get().numBuckets());
assertEquals(Arrays.asList("id_idx"), spec.get().maintainedIndexes());
assertEquals("true", spec.get().writerConfigDefaults().get("durable_write"));
}
@Test
public void testGetLsmWriteSpecAbsent() {
enqueue("get_lsm_write_spec", 200, "{\"lsm_write_spec\":null}");
assertFalse(lsm.getLsmWriteSpec().isPresent());
}
// ===========================================================================
// stats
// ===========================================================================
@Test
public void testGetLsmStats() throws Exception {
enqueue("get_lsm_stats", 200, stats(bucket("shard-0", false, 7L, 8L)));
Optional<LsmStats> got = lsm.getLsmStats(true);
assertEquals("/v1/table/my_table/get_lsm_stats/", requestPaths.get(0));
assertTrue(MAPPER.readTree(requestBodies.get(0)).get("include_generation_rows").asBoolean());
assertTrue(got.isPresent());
BucketStats decoded = got.get().buckets().get(0);
assertEquals("shard-0", decoded.shardId());
assertEquals("Active", decoded.status());
assertEquals(1, decoded.writerEpoch());
assertEquals(2, decoded.manifestVersion());
assertEquals(9, decoded.currentGeneration());
assertFalse(decoded.compacting());
assertEquals(Arrays.asList(7L, 8L), generationNumbers(decoded));
assertEquals(1024, decoded.generations().get(0).bytes());
assertFalse(decoded.generations().get(0).rows().isPresent(), "rows absent unless requested");
assertFalse(decoded.memtables().isPresent(), "absent memtables stay absent");
}
/** The optional fields decode when the server does send them. */
@Test
public void testGetLsmStatsDecodesOptionalFields() {
enqueue(
"get_lsm_stats",
200,
"{\"lsm_stats\":{\"buckets\":[{\"shard_id\":\"shard-0\",\"status\":\"Active\","
+ "\"writer_epoch\":1,\"manifest_version\":2,\"current_generation\":9,"
+ "\"replay_after_wal_entry_position\":3,\"wal_entry_position_last_seen\":11,"
+ "\"generations\":[{\"generation\":7,\"bytes\":1024,\"rows\":42}],"
+ "\"compacting\":true,\"memtables\":[{\"generation\":8,\"rows\":5,"
+ "\"bytes\":64,\"batches\":2,\"indexes\":[\"id_idx\"]}]}]}}");
BucketStats decoded = lsm.getLsmStats(true).get().buckets().get(0);
assertEquals(3, decoded.replayAfterWalEntryPosition());
assertEquals(11, decoded.walEntryPositionLastSeen());
assertTrue(decoded.compacting());
assertEquals(42, decoded.generations().get(0).rows().getAsLong());
assertTrue(decoded.memtables().isPresent());
MemtableStats memtable = decoded.memtables().get().get(0);
assertEquals(8, memtable.generation());
assertEquals(5, memtable.rows());
assertEquals(64, memtable.bytes());
assertEquals(2, memtable.batches());
assertEquals(Arrays.asList("id_idx"), memtable.indexes());
}
@Test
public void testGetLsmStatsAbsentWhenLsmDisabled() {
enqueue("get_lsm_stats", 200, "{\"lsm_stats\":null}");
assertFalse(lsm.getLsmStats().isPresent());
}
@Test
public void testGetLsmStatsDefaultsToExcludingGenerationRows() throws Exception {
enqueue("get_lsm_stats", 200, stats());
lsm.getLsmStats();
assertFalse(MAPPER.readTree(requestBodies.get(0)).get("include_generation_rows").asBoolean());
}
// ===========================================================================
// flush / compact
// ===========================================================================
@Test
public void testFlushAndCompactRoutes() {
enqueue("flush_lsm", 200, "");
enqueue("compact_lsm", 200, "");
lsm.flushLsm();
lsm.compactLsm();
assertEquals("/v1/table/my_table/flush_lsm/", requestPaths.get(0));
assertEquals("/v1/table/my_table/compact_lsm/", requestPaths.get(1));
}
@Test
public void testHttpErrorCarriesStatus() {
enqueue("flush_lsm", 404, "no such table");
LanceDbRestClient.HttpException e =
assertThrows(LanceDbRestClient.HttpException.class, () -> lsm.flushLsm());
assertEquals(404, e.statusCode());
}
// ===========================================================================
// checkpoint
// ===========================================================================
@Test
public void testCheckpointReturnsWhenLsmDisabled() {
enqueue("flush_lsm", 200, "");
enqueue("get_lsm_stats", 200, "{\"lsm_stats\":null}");
lsm.checkpointLsm();
assertEquals(0, countCalls("compact_lsm"), "nothing to compact when the LSM path is off");
}
@Test
public void testCheckpointReturnsWhenNoGenerationsOutstanding() {
enqueue("flush_lsm", 200, "");
// A bucket with no L0 generations yields no target, so the drain never starts.
enqueue("get_lsm_stats", 200, stats(bucket("shard-0", false)));
lsm.checkpointLsm();
assertEquals(0, countCalls("compact_lsm"));
}
@Test
public void testCheckpointConvergesOnceTargetGenerationsAreGone() {
enqueue("flush_lsm", 200, "");
// Watermark read: shard-0 holds generations 7 and 8, so target = 8.
enqueue("get_lsm_stats", 200, stats(bucket("shard-0", false, 7L, 8L)));
// First drain poll: both still outstanding, nothing compacting -> dispatch a pass.
enqueue("get_lsm_stats", 200, stats(bucket("shard-0", false, 7L, 8L)));
// Second drain poll: drained past the target -> done.
enqueue("get_lsm_stats", 200, stats(bucket("shard-0", false, 9L)));
enqueue("compact_lsm", 200, "");
lsm.checkpointLsm();
assertEquals(1, countCalls("compact_lsm"), "one pass dispatched");
assertEquals(3, countCalls("get_lsm_stats"), "watermark read plus two drain polls");
}
@Test
public void testCheckpointDoesNotPileOnWhileEveryTargetBucketIsCompacting() {
enqueue("flush_lsm", 200, "");
enqueue("get_lsm_stats", 200, stats(bucket("shard-0", true, 4L)));
// Still compacting on the first poll, so no pass is dispatched; then it drains.
enqueue("get_lsm_stats", 200, stats(bucket("shard-0", true, 4L)));
enqueue("get_lsm_stats", 200, stats(bucket("shard-0", false, 5L)));
lsm.checkpointLsm();
assertEquals(0, countCalls("compact_lsm"), "a latched bucket is left alone");
}
@Test
public void testCheckpointRetriesFromFlushAfterLostClaim() {
// 421 on the watermark read: the node lost its claim, so the whole thing restarts
// from flush rather than retrying the read in place.
enqueue("flush_lsm", 200, "");
enqueue("get_lsm_stats", 421, "no claim");
enqueue("get_lsm_stats", 200, stats(bucket("shard-0", false)));
lsm.checkpointLsm();
assertEquals(2, countCalls("flush_lsm"), "re-issued from flush");
}
@Test
public void testCheckpointRetriesRetryableStatusInPlace() {
enqueue("flush_lsm", 429, "latch held");
enqueue("flush_lsm", 200, "");
enqueue("get_lsm_stats", 200, stats(bucket("shard-0", false)));
lsm.checkpointLsm();
assertEquals(2, countCalls("flush_lsm"), "429 retried in place, not re-issued");
}
@Test
public void testCheckpointPropagatesTerminalStatus() {
enqueue("flush_lsm", 400, "bad request");
LanceDbRestClient.HttpException e =
assertThrows(LanceDbRestClient.HttpException.class, () -> lsm.checkpointLsm());
assertEquals(400, e.statusCode());
assertEquals(1, countCalls("flush_lsm"), "a terminal status is not retried");
}
@Test
public void testCheckpointGivesUpAfterRepeatedLostClaims() {
enqueue("flush_lsm", 421, "no claim");
IllegalStateException e = assertThrows(IllegalStateException.class, () -> lsm.checkpointLsm());
assertTrue(e.getMessage().contains("kept losing its claim"), e.getMessage());
assertEquals(4, countCalls("flush_lsm"), "the initial attempt plus MAX_REISSUES");
}
// ===========================================================================
// strict decoding
// ===========================================================================
/**
* A stats payload that does not decode must fail closed. Every one of these bodies used to be
* read as "no buckets", which is indistinguishable from a drained table, so {@code checkpointLsm}
* reported convergence for a checkpoint that never ran.
*/
@Test
public void testCheckpointRejectsMalformedStats() {
Map<String, String> malformed = new LinkedHashMap<String, String>();
malformed.put("no response body at all", "");
malformed.put("stats object with no buckets", "{\"lsm_stats\":{}}");
malformed.put("bucket missing its required fields", "{\"lsm_stats\":{\"buckets\":[{}]}}");
malformed.put(
"bucket missing generations",
"{\"lsm_stats\":{\"buckets\":[{\"shard_id\":\"shard-0\",\"status\":\"Active\","
+ "\"writer_epoch\":1,\"manifest_version\":2,\"current_generation\":9,"
+ "\"replay_after_wal_entry_position\":0,\"wal_entry_position_last_seen\":0,"
+ "\"compacting\":false}]}}");
malformed.put(
"generation with a non-numeric generation number",
"{\"lsm_stats\":{\"buckets\":[{\"shard_id\":\"shard-0\",\"status\":\"Active\","
+ "\"writer_epoch\":1,\"manifest_version\":2,\"current_generation\":9,"
+ "\"replay_after_wal_entry_position\":0,\"wal_entry_position_last_seen\":0,"
+ "\"generations\":[{\"generation\":\"7\",\"bytes\":1024}],"
+ "\"compacting\":false}]}}");
for (Map.Entry<String, String> each : malformed.entrySet()) {
setUpFresh();
enqueue("flush_lsm", 200, "");
enqueue("get_lsm_stats", 200, each.getValue());
assertThrows(
IllegalStateException.class,
() -> lsm.checkpointLsm(),
each.getKey() + " must not report convergence");
}
}
/** The one shape that legitimately means "this table has no LSM write path". */
@Test
public void testCheckpointTreatsNullStatsAsNotWalBacked() {
enqueue("flush_lsm", 200, "");
enqueue("get_lsm_stats", 200, "{\"lsm_stats\":null}");
lsm.checkpointLsm();
assertEquals(1, countCalls("get_lsm_stats"));
}
// ===========================================================================
// retry budget
// ===========================================================================
/**
* The transport must not retry on the checkpoint loop's behalf. Apache HttpClient's default
* strategy retries exactly 429 and 503 — the two statuses {@code isRetryable} owns — which
* doubled every budget here and also retried {@code compact_lsm} in place, where the loop is
* built to fall through to a fresh stats poll instead.
*/
@Test
public void testCheckpointRetryBudgetIsNotDoubledByTheTransport() {
enqueue("flush_lsm", 429, "latch held");
LanceDbRestClient.HttpException e =
assertThrows(LanceDbRestClient.HttpException.class, () -> lsm.checkpointLsm());
assertEquals(429, e.statusCode(), "the exhausted budget propagates the last error as itself");
assertEquals(9, countCalls("flush_lsm"), "the initial request plus MAX_RETRIES, and no more");
}
// ===========================================================================
// harness
// ===========================================================================
private static List<Long> generationNumbers(BucketStats bucket) {
List<Long> numbers = new ArrayList<Long>();
for (GenerationStats generation : bucket.generations()) {
numbers.add(generation.generation());
}
return numbers;
}
/** Build an {@code lsm_stats} response body from bucket fragments. */
private static String stats(String... buckets) {
return "{\"lsm_stats\":{\"buckets\":[" + String.join(",", buckets) + "]}}";
}
private static String bucket(String shardId, boolean compacting, Long... generations) {
StringBuilder gens = new StringBuilder();
for (Long generation : generations) {
if (gens.length() > 0) {
gens.append(",");
}
gens.append("{\"generation\":").append(generation).append(",\"bytes\":1024}");
}
return "{\"shard_id\":\""
+ shardId
+ "\",\"status\":\"Active\",\"writer_epoch\":1,\"manifest_version\":2,"
+ "\"current_generation\":9,\"replay_after_wal_entry_position\":0,"
+ "\"wal_entry_position_last_seen\":0,\"generations\":["
+ gens
+ "],\"compacting\":"
+ compacting
+ "}";
}
/** Queue a reply for an operation. The last queued reply repeats once the queue drains. */
private void enqueue(String operation, int status, String body) {
replies.computeIfAbsent(operation, key -> new ArrayDeque<Reply>()).add(new Reply(status, body));
}
private Reply nextReply(String path) {
String operation = operationOf(path);
Deque<Reply> queued = replies.get(operation);
if (queued == null || queued.isEmpty()) {
return new Reply(200, "");
}
return queued.size() > 1 ? queued.poll() : queued.peek();
}
private long countCalls(String operation) {
return requestPaths.stream().filter(path -> operationOf(path).equals(operation)).count();
}
/** {@code /v1/table/my_table/flush_lsm/} -> {@code flush_lsm}. */
private static String operationOf(String path) {
String[] segments = path.split("/");
return segments.length == 0 ? "" : segments[segments.length - 1];
}
private static String readAll(InputStream in) throws IOException {
ByteArrayOutputStream out = new ByteArrayOutputStream();
byte[] buffer = new byte[4096];
int read;
while ((read = in.read(buffer)) != -1) {
out.write(buffer, 0, read);
}
return new String(out.toByteArray(), StandardCharsets.UTF_8);
}
private static final class Reply {
private final int status;
private final String body;
private Reply(int status, String body) {
this.status = status;
this.body = body;
}
}
}
+2 -2
View File
@@ -6,7 +6,7 @@
<groupId>com.lancedb</groupId>
<artifactId>lancedb-parent</artifactId>
<version>0.38.0-beta.2</version>
<version>0.37.1-beta.1</version>
<packaging>pom</packaging>
<name>${project.artifactId}</name>
<description>LanceDB Java SDK Parent POM</description>
@@ -28,7 +28,7 @@
<properties>
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
<arrow.version>15.0.0</arrow.version>
<lance-core.version>11.0.0-beta.15</lance-core.version>
<lance-core.version>11.0.0-beta.3</lance-core.version>
<spotless.skip>false</spotless.skip>
<spotless.version>2.30.0</spotless.version>
<spotless.java.googlejavaformat.version>1.7</spotless.java.googlejavaformat.version>
+1 -1
View File
@@ -1,7 +1,7 @@
[package]
name = "lancedb-nodejs"
edition.workspace = true
version = "0.38.0-beta.2"
version = "0.37.1-beta.1"
publish = false
license.workspace = true
description.workspace = true
+1 -67
View File
@@ -4,13 +4,7 @@
import { readdirSync } from "fs";
import { Field, Float64, Schema } from "apache-arrow";
import * as tmp from "tmp";
import {
Connection,
ListTablesResponse,
Table,
connect,
connectNamespace,
} from "../lancedb";
import { Connection, Table, connect, connectNamespace } from "../lancedb";
import { LocalTable } from "../lancedb/table";
describe("when connecting", () => {
@@ -95,16 +89,6 @@ describe("given a connection", () => {
await db.createTable("test4", [{ id: 1 }, { id: 2 }]);
});
it("should return a completed job when dropping a local table", async () => {
await db.createTable("async-drop", [{ id: 1 }]);
const job = await db.dropTableAsync("async-drop");
expect(job.id).toBeNull();
await expect(job.status()).resolves.toBe("finished");
await job.wait();
await expect(db.tableNames()).resolves.toEqual([]);
});
it("should fail if creating table twice, unless overwrite is true", async () => {
let tbl = await db.createTable("test", [{ id: 1 }, { id: 2 }]);
await expect(tbl.countRows()).resolves.toBe(2);
@@ -135,56 +119,6 @@ describe("given a connection", () => {
expect(tables).toEqual(["b", "c"]);
});
it("should list tables with a page token", async () => {
const db = await connect(tmpDir.name);
await db.createTable("b", [{ id: 1 }]);
await db.createTable("a", [{ id: 1 }]);
await db.createTable("c", [{ id: 1 }]);
const all = await db.listTables();
expect(all.tables).toEqual(["a", "b", "c"]);
expect(all.pageToken).toBeUndefined();
const first = await db.listTables({ limit: 1 });
expect(first.tables).toEqual(["a"]);
expect(first.pageToken).toBeDefined();
const second = await db.listTables({
limit: 1,
pageToken: first.pageToken,
});
expect(second.tables).toEqual(["b"]);
});
it("should visit every table exactly once when paging", async () => {
const db = await connect(tmpDir.name);
const created = ["a", "b", "c", "d", "e"];
for (const name of created) {
await db.createTable(name, [{ id: 1 }]);
}
const seen: string[] = [];
let pageToken: string | undefined = undefined;
do {
const page: ListTablesResponse = await db.listTables({
limit: 2,
pageToken,
});
seen.push(...page.tables);
pageToken = page.pageToken;
} while (pageToken);
expect(seen.sort()).toEqual(created);
});
it("should reject listTables on a closed connection", async () => {
const db = await connect(tmpDir.name);
db.close();
await expect(db.listTables()).rejects.toThrow("Connection is closed");
});
it("should create tables in v2 mode", async () => {
const db = await connect(tmpDir.name);
const data = [...Array(10000).keys()].map((i) => ({ id: i }));
-117
View File
@@ -3340,120 +3340,3 @@ describe("LSM merge insert", () => {
await expect(table.query().useLsm(true).toArray()).rejects.toThrow();
});
});
describe("LSM convergence and stats", () => {
let tmpDir: tmp.DirResult;
beforeEach(() => {
tmpDir = tmp.dirSync({ unsafeCleanup: true });
});
afterEach(() => tmpDir.removeCallback());
async function lsmTable(conn: Connection): Promise<Table> {
const table = await conn.createEmptyTable(
"t",
new arrow.Schema([new arrow.Field("id", new arrow.Utf8(), false)]),
);
await table.setUnenforcedPrimaryKey("id");
await table.setLsmWriteSpec({ specType: "unsharded" });
return table;
}
// These four route through the server that owns the MemWAL, so a local table
// rejects them rather than answering. What is asserted here is that the
// bindings reach the core at all; the behavior against a real endpoint is
// covered by the mocked endpoint tests in rust/lancedb/src/remote/table.rs.
it("rejects flushLsm on a local table", async () => {
const conn = await connect(tmpDir.name);
const table = await lsmTable(conn);
await expect(table.flushLsm()).rejects.toThrow(/not supported/i);
});
it("rejects compactLsm on a local table", async () => {
const conn = await connect(tmpDir.name);
const table = await lsmTable(conn);
await expect(table.compactLsm()).rejects.toThrow(/not supported/i);
});
it("rejects getLsmStats on a local table", async () => {
const conn = await connect(tmpDir.name);
const table = await lsmTable(conn);
await expect(table.getLsmStats()).rejects.toThrow(/not supported/i);
await expect(table.getLsmStats(true)).rejects.toThrow(/not supported/i);
});
it("rejects checkpointLsm on a local table", async () => {
const conn = await connect(tmpDir.name);
const table = await lsmTable(conn);
// checkpointLsm seals first, so it surfaces flushLsm's rejection.
await expect(table.checkpointLsm()).rejects.toThrow(/not supported/i);
});
});
describe("computed columns", () => {
let tmpDir: tmp.DirResult;
beforeEach(() => {
tmpDir = tmp.dirSync({ unsafeCleanup: true });
});
afterEach(() => tmpDir.removeCallback());
it("declares a column and fills it on refresh", async () => {
const db = await connect(tmpDir.name);
const table = await db.createTable("computed", [{ x: 1 }, { x: 2 }]);
await table.addColumns({
computed: [{ name: "doubled", valueSql: "x * 2" }],
});
let rows = await table.query().toArray();
expect(rows.map((r) => r.doubled)).toEqual([null, null]);
const result = await table.refreshColumn("doubled");
expect(result.rowsFilled).toBe(2);
rows = await table.query().toArray();
expect(rows.map((r) => r.doubled).sort()).toEqual([2, 4]);
});
it("returns a job handle from refreshColumnAsync", async () => {
const db = await connect(tmpDir.name);
const table = await db.createTable("computed_job", [{ x: 1 }, { x: 2 }]);
await table.addColumns({
computed: [{ name: "doubled", valueSql: "x * 2" }],
});
const job = await table.refreshColumnAsync("doubled");
expect(job.id).toBeNull();
await job.wait();
expect(await job.status()).toBe("finished");
const rows = await table.query().toArray();
expect(rows.map((r) => r.doubled).sort()).toEqual([2, 4]);
// Bad input rejects at the call, not through the job.
await expect(table.refreshColumnAsync("x")).rejects.toThrow(
"not a computed column",
);
});
it("fills rows added since the last refresh", async () => {
const db = await connect(tmpDir.name);
const table = await db.createTable("computed_append", [{ x: 1 }]);
await table.addColumns({
computed: [{ name: "doubled", valueSql: "x * 2" }],
});
await table.refreshColumn("doubled");
await table.add([{ x: 5 }]);
const result = await table.refreshColumn("doubled");
expect(result.rowsFilled).toBe(1);
const rows = await table.query().toArray();
expect(rows.map((r) => r.doubled).sort()).toEqual([10, 2]);
});
});
-101
View File
@@ -25,14 +25,12 @@ import type {
JobDescription,
JobInfo,
ListNamespacesResponse,
ListTablesResponse,
} from "./native";
export type {
CreateNamespaceResponse,
DescribeNamespaceResponse,
DropNamespaceResponse,
ListNamespacesResponse,
ListTablesResponse,
};
import { sanitizeTable } from "./sanitize";
import { LocalTable, Table } from "./table";
@@ -130,10 +128,6 @@ export interface OpenTableOptions {
indexCacheSize?: number;
}
/**
* @deprecated Use {@link ListTablesOptions} with {@link Connection.listTables}
* instead.
*/
export interface TableNamesOptions {
/**
* If present, only return names that come lexicographically after the
@@ -147,23 +141,6 @@ export interface TableNamesOptions {
limit?: number;
}
export interface ListTablesOptions {
/**
* Token from a previous response for pagination.
*
* The token is opaque: it carries whatever the database needs to resume, and
* callers should not construct or interpret one.
*/
pageToken?: string;
/**
* An upper bound on how many tables to return.
*
* A page may hold fewer than this and still not be the last one, so continue
* while the response carries a page token rather than while pages are full.
*/
limit?: number;
}
export interface ListNamespacesOptions {
/** Token from a previous response for pagination. */
pageToken?: string;
@@ -248,7 +225,6 @@ export abstract class Connection {
* @param {Partial<TableNamesOptions>} options - options to control the
* paging / start point (backwards compatibility)
*
* @deprecated Use {@link Connection.listTables} instead.
*/
abstract tableNames(options?: Partial<TableNamesOptions>): Promise<string[]>;
/**
@@ -259,54 +235,12 @@ export abstract class Connection {
* @param {Partial<TableNamesOptions>} options - options to control the
* paging / start point
*
* @deprecated Use {@link Connection.listTables} instead.
*/
abstract tableNames(
namespacePath?: string[],
options?: Partial<TableNamesOptions>,
): Promise<string[]>;
/**
* List a page of tables in this database.
*
* Results may be paginated. To retrieve subsequent pages, pass the
* `pageToken` returned by a previous call. A page may be shorter than
* `limit` without being the last one, so walk until the response carries no
* page token:
*
* ```ts
* const names = [];
* let pageToken = undefined;
* do {
* const page = await conn.listTables({ pageToken, limit: 100 });
* names.push(...page.tables);
* pageToken = page.pageToken;
* } while (pageToken);
* ```
*
* @param {Partial<ListTablesOptions>} options - Pagination options
* (`pageToken`, `limit`).
* @returns {Promise<ListTablesResponse>} Table names and an optional token
* for fetching the next page.
*/
abstract listTables(
options?: Partial<ListTablesOptions>,
): Promise<ListTablesResponse>;
/**
* List a page of tables in this database.
*
* @param {string[]} namespacePath - The namespace path to list tables from
* (defaults to root namespace)
* @param {Partial<ListTablesOptions>} options - Pagination options
* (`pageToken`, `limit`).
* @returns {Promise<ListTablesResponse>} Table names and an optional token
* for fetching the next page.
*/
abstract listTables(
namespacePath?: string[],
options?: Partial<ListTablesOptions>,
): Promise<ListTablesResponse>;
/**
* Open a table in the database.
* @param {string} name - The name of the table
@@ -393,14 +327,6 @@ export abstract class Connection {
*/
abstract dropTable(name: string, namespacePath?: string[]): Promise<void>;
/**
* Start dropping a table and return its cleanup job.
*
* The table may become unavailable before its data files are removed. Wait
* on the returned job to know when cleanup has finished.
*/
abstract dropTableAsync(name: string, namespacePath?: string[]): Promise<Job>;
/**
* Drop all tables in the database.
* @param {string[]} namespacePath The namespace path to drop tables from (defaults to root namespace).
@@ -597,29 +523,6 @@ export class LocalConnection extends Connection {
);
}
async listTables(
namespacePathOrOptions?: string[] | Partial<ListTablesOptions>,
options?: Partial<ListTablesOptions>,
): Promise<ListTablesResponse> {
// Detect if first argument is namespacePath array or options object
let namespacePath: string[] | undefined;
let listTablesOptions: Partial<ListTablesOptions> | undefined;
if (Array.isArray(namespacePathOrOptions)) {
namespacePath = namespacePathOrOptions;
listTablesOptions = options;
} else {
namespacePath = undefined;
listTablesOptions = namespacePathOrOptions;
}
return this.inner.listTables(
namespacePath ?? [],
listTablesOptions?.pageToken,
listTablesOptions?.limit,
);
}
async openTable(
name: string,
namespacePath?: string[],
@@ -802,10 +705,6 @@ export class LocalConnection extends Connection {
return this.inner.dropTable(name, namespacePath ?? []);
}
async dropTableAsync(name: string, namespacePath?: string[]): Promise<Job> {
return this.inner.dropTableAsync(name, namespacePath ?? []);
}
async dropAllTables(namespacePath?: string[]): Promise<void> {
return this.inner.dropAllTables(namespacePath ?? []);
}
-7
View File
@@ -50,7 +50,6 @@ export {
MergeResult,
AddResult,
AddColumnsResult,
RefreshColumnResult,
AlterColumnsResult,
UpdateFieldMetadataResult,
DeleteResult,
@@ -75,13 +74,11 @@ export {
Connection,
CreateTableOptions,
TableNamesOptions,
ListTablesOptions,
OpenTableOptions,
ListNamespacesOptions,
CreateNamespaceOptions,
DropNamespaceOptions,
ListNamespacesResponse,
ListTablesResponse,
CreateNamespaceResponse,
DropNamespaceResponse,
DescribeNamespaceResponse,
@@ -149,10 +146,6 @@ export {
FtsToken,
TokenizeTableOptions,
LsmWriteSpec,
LsmStats,
BucketStats,
GenerationStats,
MemtableStats,
ColumnAlteration,
FieldMetadataUpdate,
} from "./table";
+2 -160
View File
@@ -31,10 +31,8 @@ import {
IndexConfig,
IndexStatistics,
Job,
LsmStats,
Branches as NativeBranches,
OptimizeStats,
RefreshColumnResult,
TableStatistics,
Tags,
UpdateFieldMetadataResult,
@@ -51,12 +49,6 @@ import {
import { sanitizeType } from "./sanitize";
import { IntoSql, toSQL } from "./util";
export { IndexConfig } from "./native";
export {
BucketStats,
GenerationStats,
LsmStats,
MemtableStats,
} from "./native";
/**
* Progress snapshot for a write operation, delivered to the `progress`
@@ -533,75 +525,18 @@ export abstract class Table {
abstract vectorSearch(vector: IntoVector | MultiVector): VectorQuery;
/**
* Add new columns with defined values.
*
* The `{ computed }` form stores the expression rather than evaluating it
* now: the column is committed with no values, and rows get them from
* {@link Table#refreshColumn}. Declaring one therefore costs the same on a
* large table as on an empty one.
*
* A refresh does not revisit rows it has already filled, so mutating an
* input leaves the value computed at fill time; recomputing means dropping
* the column and declaring it again. While a declaration reads a column,
* that column cannot be renamed, retyped or dropped.
*
* On LanceDB Cloud and Enterprise the expression is planned by the
* server, and the refresh runs as a server job -- see
* {@link Table#refreshColumnAsync}.
* @param {AddColumnsSql[] | Field | Field[] | Schema} newColumnTransforms Either:
* - An array of objects with column names and SQL expressions to calculate values
* - A single Arrow Field defining one column with its data type (column will be initialized with null values)
* - An array of Arrow Fields defining columns with their data types (columns will be initialized with null values)
* - An Arrow Schema defining columns with their data types (columns will be initialized with null values)
* - `{ computed }`, declaring columns defined by a SQL expression whose type and inputs are derived from it
* @returns {Promise<AddColumnsResult>} A promise that resolves to an object
* containing the new version number of the table after adding the columns.
* @example
* ```ts
* await table.addColumns({ computed: [{ name: "doubled", valueSql: "x * 2" }] });
* const { rowsFilled } = await table.refreshColumn("doubled");
* ```
*/
abstract addColumns(
newColumnTransforms:
| AddColumnsSql[]
| Field
| Field[]
| Schema
| { computed: AddColumnsSql[] },
newColumnTransforms: AddColumnsSql[] | Field | Field[] | Schema,
): Promise<AddColumnsResult>;
/**
* Fill the rows of a computed column that hold no value yet.
*
* Rows appended since the last refresh are filled by the next one; rows
* already filled are left as they are, so the call is idempotent and does
* not observe a mutated input. Local tables only: a remote refresh runs
* as a server job, through {@link Table#refreshColumnAsync}.
* @param {string} column The name of the computed column to fill.
* @returns {Promise<RefreshColumnResult>} A promise that resolves to the
* number of rows filled and the new version number of the table.
*/
abstract refreshColumn(column: string): Promise<RefreshColumnResult>;
/**
* Like {@link Table#refreshColumn}, but returns a handle to the refresh
* job instead of blocking until it completes.
*
* The job may already be complete when returned; callers must not assume
* the column is filled until {@link Job.wait} resolves. Invalid input --
* an unknown column, or one that is not computed -- rejects here rather
* than failing the job. On local tables the job runs in-process; on
* LanceDB Cloud and Enterprise it is the server's backfill job.
* @param {string} column The name of the computed column to fill.
* @example
* ```ts
* const job = await table.refreshColumnAsync("doubled");
* await job.wait();
* console.log(await job.status()); // "finished"
* ```
*/
abstract refreshColumnAsync(column: string): Promise<Job>;
/**
* Alter the name or nullability of columns.
* @param {ColumnAlteration[]} columnAlterations One or more alterations to
@@ -713,59 +648,6 @@ export abstract class Table {
* @returns {Promise<void>}
*/
abstract closeLsmWriters(): Promise<void>;
/**
* Seal every bucket's active memtable into a new L0 generation.
*
* Returns once the seal is committed. Sealing an empty memtable is a no-op,
* so this is safe to call repeatedly.
* @returns {Promise<void>}
*/
abstract flushLsm(): Promise<void>;
/**
* Trigger a background L0 → base compaction pass per bucket.
*
* Returns once the passes are *dispatched*, not once they finish — watch
* {@link Table#getLsmStats} for progress, or use
* {@link Table#checkpointLsm} to wait for convergence.
* @returns {Promise<void>}
*/
abstract compactLsm(): Promise<void>;
/**
* Converge this table's LSM write path into its base table.
*
* Seals once, then triggers compaction and polls until the L0 that existed
* at the start is gone. The target set is fixed at the start, so
* generations created *during* the checkpoint are ignored — that is what
* lets it terminate under write load, and what makes it best-effort: it
* converges the fresh tier as of some instant. Idempotent, abandonable at
* any point, and safe to run on a cadence.
*
* There is no liveness bound — the compactor pool is shared across tables,
* so a checkpoint queued behind unrelated work looks exactly like one that
* is merging. The caller owns the deadline.
* @returns {Promise<void>}
* @example
* ```ts
* const before = await table.getLsmStats();
* await table.checkpointLsm();
* const after = await table.getLsmStats();
* ```
*/
abstract checkpointLsm(): Promise<void>;
/**
* Read live per-bucket LSM state.
*
* Answers "how far behind is my fresh tier", "which bucket is hot", and
* "why is my fresh-tier vector search brute-force". Mutates no table state.
*
* Resolves to `undefined` only when the LSM write path is not enabled.
* @param {boolean} includeGenerationRows Also count rows per L0 generation.
* Off by default because each count opens an uncached Lance dataset.
* @returns {Promise<LsmStats | undefined>}
*/
abstract getLsmStats(
includeGenerationRows?: boolean,
): Promise<LsmStats | undefined>;
/** Retrieve the version of the table */
abstract version(): Promise<number>;
@@ -1206,22 +1088,8 @@ export class LocalTable extends Table {
// TODO: Support BatchUDF
async addColumns(
newColumnTransforms:
| AddColumnsSql[]
| Field
| Field[]
| Schema
| { computed: AddColumnsSql[] },
newColumnTransforms: AddColumnsSql[] | Field | Field[] | Schema,
): Promise<AddColumnsResult> {
// Columns defined by an expression are declared, not materialized here.
if (
typeof newColumnTransforms === "object" &&
!Array.isArray(newColumnTransforms) &&
"computed" in newColumnTransforms
) {
return await this.inner.addComputedColumns(newColumnTransforms.computed);
}
// Handle single Field -> convert to array of Fields
if (newColumnTransforms instanceof Field) {
newColumnTransforms = [newColumnTransforms];
@@ -1256,14 +1124,6 @@ export class LocalTable extends Table {
throw new Error("Invalid input type for addColumns");
}
async refreshColumn(column: string): Promise<RefreshColumnResult> {
return await this.inner.refreshColumn(column);
}
async refreshColumnAsync(column: string): Promise<Job> {
return await this.inner.refreshColumnAsync(column);
}
async alterColumns(
columnAlterations: ColumnAlteration[],
): Promise<AlterColumnsResult> {
@@ -1326,24 +1186,6 @@ export class LocalTable extends Table {
return await this.inner.closeLsmWriters();
}
async flushLsm(): Promise<void> {
return await this.inner.flushLsm();
}
async compactLsm(): Promise<void> {
return await this.inner.compactLsm();
}
async checkpointLsm(): Promise<void> {
return await this.inner.checkpointLsm();
}
async getLsmStats(
includeGenerationRows: boolean = false,
): Promise<LsmStats | undefined> {
return (await this.inner.getLsmStats(includeGenerationRows)) ?? undefined;
}
async version(): Promise<number> {
return await this.inner.version();
}
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-darwin-arm64",
"version": "0.38.0-beta.2",
"version": "0.37.1-beta.1",
"os": ["darwin"],
"cpu": ["arm64"],
"main": "lancedb.darwin-arm64.node",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-linux-arm64-gnu",
"version": "0.38.0-beta.2",
"version": "0.37.1-beta.1",
"os": ["linux"],
"cpu": ["arm64"],
"main": "lancedb.linux-arm64-gnu.node",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-linux-arm64-musl",
"version": "0.38.0-beta.2",
"version": "0.37.1-beta.1",
"os": ["linux"],
"cpu": ["arm64"],
"main": "lancedb.linux-arm64-musl.node",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-linux-x64-gnu",
"version": "0.38.0-beta.2",
"version": "0.37.1-beta.1",
"os": ["linux"],
"cpu": ["x64"],
"main": "lancedb.linux-x64-gnu.node",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-linux-x64-musl",
"version": "0.38.0-beta.2",
"version": "0.37.1-beta.1",
"os": ["linux"],
"cpu": ["x64"],
"main": "lancedb.linux-x64-musl.node",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-win32-arm64-msvc",
"version": "0.38.0-beta.2",
"version": "0.37.1-beta.1",
"os": [
"win32"
],
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-win32-x64-msvc",
"version": "0.38.0-beta.2",
"version": "0.37.1-beta.1",
"os": ["win32"],
"cpu": ["x64"],
"main": "lancedb.win32-x64-msvc.node",
+2 -2
View File
@@ -1,12 +1,12 @@
{
"name": "@lancedb/lancedb",
"version": "0.38.0-beta.2",
"version": "0.37.1-beta.1",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"name": "@lancedb/lancedb",
"version": "0.38.0-beta.2",
"version": "0.37.1-beta.1",
"cpu": [
"x64",
"arm64"
+1 -1
View File
@@ -11,7 +11,7 @@
"ann"
],
"private": false,
"version": "0.38.0-beta.2",
"version": "0.37.1-beta.1",
"main": "dist/index.js",
"exports": {
".": "./dist/index.js",
-47
View File
@@ -36,12 +36,6 @@ pub struct ListNamespacesResponse {
pub page_token: Option<String>,
}
#[napi(object)]
pub struct ListTablesResponse {
pub tables: Vec<String>,
pub page_token: Option<String>,
}
#[napi(object)]
pub struct CreateNamespaceResponse {
pub properties: Option<HashMap<String, String>>,
@@ -195,8 +189,6 @@ impl Connection {
/// List all tables in the dataset.
#[napi(catch_unwind)]
// Deprecated in favour of `list_tables`, but still exposed to JavaScript.
#[allow(deprecated)]
pub async fn table_names(
&self,
namespace_path: Option<Vec<String>>,
@@ -214,29 +206,6 @@ impl Connection {
op.execute().await.default_error()
}
/// List a page of tables in the database.
#[napi(catch_unwind)]
pub async fn list_tables(
&self,
namespace_path: Option<Vec<String>>,
page_token: Option<String>,
limit: Option<u32>,
) -> napi::Result<ListTablesResponse> {
let mut op = self.get_inner()?.list_tables();
op = op.namespace(namespace_path.unwrap_or_default());
if let Some(page_token) = page_token {
op = op.page_token(page_token);
}
if let Some(limit) = limit {
op = op.limit(limit);
}
let resp = op.execute().await.default_error()?;
Ok(ListTablesResponse {
tables: resp.tables,
page_token: resp.page_token,
})
}
/// Create table from a Apache Arrow IPC (file) buffer.
///
/// Parameters:
@@ -365,22 +334,6 @@ impl Connection {
.default_error()
}
/// Start dropping a table and return its cleanup job.
#[napi(catch_unwind)]
pub async fn drop_table_async(
&self,
name: String,
namespace_path: Option<Vec<String>>,
) -> napi::Result<crate::job::Job> {
let ns = namespace_path.unwrap_or_default();
let job = self
.get_inner()?
.drop_table_async(&name, &ns)
.await
.default_error()?;
Ok(crate::job::Job::new(job))
}
#[napi(catch_unwind)]
pub async fn drop_all_tables(&self, namespace_path: Option<Vec<String>>) -> napi::Result<()> {
let ns = namespace_path.unwrap_or_default();
+11 -1
View File
@@ -42,9 +42,19 @@ impl Job {
}
/// Wait until the operation reaches a terminal state.
///
/// Jobs that complete without a resource result resolve successfully.
/// Resource results are not exposed on this binding yet; unsupported
/// success results reject with a generic error.
#[napi(catch_unwind)]
pub async fn wait(&self) -> napi::Result<()> {
self.inner.wait().await.default_error()
match self.inner.wait().await.default_error()? {
lancedb::JobResult::None => Ok(()),
// JobResult is non_exhaustive; Function and future variants fail closed.
_ => Err(napi::Error::from_reason(
"unsupported job result".to_string(),
)),
}
}
/// Request cancellation. Cancelling a finished operation is a no-op.
-200
View File
@@ -347,40 +347,6 @@ impl Table {
Ok(res.into())
}
#[napi(catch_unwind)]
pub async fn add_computed_columns(
&self,
columns: Vec<AddColumnsSql>,
) -> napi::Result<AddColumnsResult> {
let table = self.inner_ref()?;
let mut builder = table.add_columns();
for column in columns {
builder = builder.computed(column.name, column.value_sql);
}
let res = builder.execute().await.default_error()?;
Ok(res.into())
}
#[napi(catch_unwind)]
pub async fn refresh_column(&self, column: String) -> napi::Result<RefreshColumnResult> {
let res = self
.inner_ref()?
.refresh_column(column)
.await
.default_error()?;
Ok(res.into())
}
#[napi(catch_unwind)]
pub async fn refresh_column_async(&self, column: String) -> napi::Result<crate::job::Job> {
let job = self
.inner_ref()?
.refresh_column_async(column)
.await
.default_error()?;
Ok(crate::job::Job::new(job))
}
#[napi(catch_unwind)]
pub async fn add_columns_with_schema(
&self,
@@ -497,34 +463,6 @@ impl Table {
self.inner_ref()?.close_lsm_writers().await.default_error()
}
#[napi(catch_unwind)]
pub async fn flush_lsm(&self) -> napi::Result<()> {
self.inner_ref()?.flush_lsm().await.default_error()
}
#[napi(catch_unwind)]
pub async fn compact_lsm(&self) -> napi::Result<()> {
self.inner_ref()?.compact_lsm().await.default_error()
}
#[napi(catch_unwind)]
pub async fn checkpoint_lsm(&self) -> napi::Result<()> {
self.inner_ref()?.checkpoint_lsm().await.default_error()
}
#[napi(catch_unwind)]
pub async fn get_lsm_stats(
&self,
include_generation_rows: bool,
) -> napi::Result<Option<LsmStats>> {
let stats = self
.inner_ref()?
.get_lsm_stats(include_generation_rows)
.await
.default_error()?;
Ok(stats.map(LsmStats::from))
}
#[napi(catch_unwind)]
pub async fn version(&self) -> napi::Result<i64> {
self.inner_ref()?
@@ -917,129 +855,6 @@ impl From<lancedb::table::LsmWriteSpec> for LsmWriteSpec {
}
}
/// One flushed L0 generation.
#[napi(object)]
#[derive(Clone, Debug)]
pub struct GenerationStats {
/// The generation number. Increases as memtables are sealed into L0.
pub generation: i64,
/// On-disk size of the generation.
pub bytes: i64,
/// Present only when `includeGenerationRows` was requested. Off by default
/// because each count opens an uncached Lance dataset.
pub rows: Option<i64>,
}
impl From<lancedb::table::GenerationStats> for GenerationStats {
fn from(g: lancedb::table::GenerationStats) -> Self {
Self {
generation: g.generation as i64,
bytes: g.bytes as i64,
rows: g.rows.map(|r| r as i64),
}
}
}
/// One in-memory memtable.
#[napi(object)]
#[derive(Clone, Debug)]
pub struct MemtableStats {
/// The generation this memtable will become once sealed.
pub generation: i64,
/// Rows currently buffered.
pub rows: i64,
/// Estimated in-memory size.
pub bytes: i64,
/// Record batches currently buffered.
pub batches: i64,
/// Names of the indexes this memtable carries. An absent name is the whole
/// answer to "why is my fresh-tier search on that column brute-force".
pub indexes: Vec<String>,
}
impl From<lancedb::table::MemtableStats> for MemtableStats {
fn from(m: lancedb::table::MemtableStats) -> Self {
Self {
generation: m.generation as i64,
rows: m.rows as i64,
bytes: m.bytes as i64,
batches: m.batches as i64,
indexes: m.indexes,
}
}
}
/// Live state of one bucket. A table is N buckets on one node; flattening to a
/// single number hides the one hot bucket that is usually why someone opened
/// this endpoint.
#[napi(object)]
#[derive(Clone, Debug)]
pub struct BucketStats {
/// The shard this bucket writes.
pub shard_id: String,
/// `"Active"` or `"Sealed"` (drop-table 2PC in flight).
pub status: String,
/// Epoch of the writer that currently owns the shard.
pub writer_epoch: i64,
/// Version of the shard manifest these numbers were read from.
pub manifest_version: i64,
/// The generation the active memtable will become.
pub current_generation: i64,
/// WAL position replay resumes from.
pub replay_after_wal_entry_position: i64,
/// Highest WAL position the writer has seen. The difference against
/// `replayAfterWalEntryPosition` is the WAL lag.
pub wal_entry_position_last_seen: i64,
/// Flushed L0 generations not yet merged into the base table.
pub generations: Vec<GenerationStats>,
/// Whether a pass owns this bucket's compaction latch right now. Says *a*
/// driver is running, not *whose*, and the latch is held from dispatch —
/// including while the pass queues for a pod-wide compactor permit. Read it
/// as "do not pile on", never as "mine is progressing".
pub compacting: bool,
/// Oldest first, active last. Absent for a `"Sealed"` bucket, whose
/// in-memory state is torn down.
pub memtables: Option<Vec<MemtableStats>>,
}
impl From<lancedb::table::BucketStats> for BucketStats {
fn from(b: lancedb::table::BucketStats) -> Self {
Self {
shard_id: b.shard_id,
status: b.status,
writer_epoch: b.writer_epoch as i64,
manifest_version: b.manifest_version as i64,
current_generation: b.current_generation as i64,
replay_after_wal_entry_position: b.replay_after_wal_entry_position as i64,
wal_entry_position_last_seen: b.wal_entry_position_last_seen as i64,
generations: b.generations.into_iter().map(Into::into).collect(),
compacting: b.compacting,
memtables: b
.memtables
.map(|ms| ms.into_iter().map(Into::into).collect()),
}
}
}
/// Live per-bucket LSM state, as returned by `Table#getLsmStats`.
///
/// Nothing here is derived: sums and differences (total L0 bytes, WAL lag) are
/// the caller's to compute.
#[napi(object)]
#[derive(Clone, Debug)]
pub struct LsmStats {
/// One entry per bucket backing this table.
pub buckets: Vec<BucketStats>,
}
impl From<lancedb::table::LsmStats> for LsmStats {
fn from(stats: lancedb::table::LsmStats) -> Self {
Self {
buckets: stats.buckets.into_iter().map(Into::into).collect(),
}
}
}
/// Statistics about a compaction operation.
#[napi(object)]
#[derive(Clone, Debug)]
@@ -1381,21 +1196,6 @@ pub struct AddColumnsResult {
pub version: i64,
}
#[napi(object)]
pub struct RefreshColumnResult {
pub rows_filled: i64,
pub version: i64,
}
impl From<lancedb::table::RefreshColumnResult> for RefreshColumnResult {
fn from(value: lancedb::table::RefreshColumnResult) -> Self {
Self {
rows_filled: value.rows_filled as i64,
version: value.version as i64,
}
}
}
impl From<lancedb::table::AddColumnsResult> for AddColumnsResult {
fn from(value: lancedb::table::AddColumnsResult) -> Self {
Self {
+1 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "lancedb-python"
version = "0.38.0-beta.2"
version = "0.37.1-beta.1"
publish = false
edition.workspace = true
description = "Python bindings for LanceDB"
+5 -2
View File
@@ -12,7 +12,7 @@ __version__ = importlib.metadata.version("lancedb")
from ._lancedb import connect as lancedb_connect
from ._lancedb import FtsToken
from ._lancedb import LsmWriteSpec
from ._lancedb import Function
from ._lancedb import tokenize as _tokenize
from .common import URI, sanitize_uri
from urllib.parse import urlparse
@@ -24,6 +24,7 @@ from .schema import blob, vector, BlobType
from .job import AsyncJob, Job
from .table import AsyncTable, Table
from .types import BaseTokenizerType
from ._udf import FunctionCapability, udf
from ._lancedb import Session
from .namespace import (
connect_namespace,
@@ -508,6 +509,8 @@ __all__ = [
"FtsToken",
"col",
"Expr",
"Function",
"FunctionCapability",
"func",
"lit",
"URI",
@@ -519,9 +522,9 @@ __all__ = [
"Job",
"LanceDBConnection",
"LanceNamespaceDBConnection",
"LsmWriteSpec",
"RemoteDBConnection",
"Session",
"Table",
"udf",
"__version__",
]
+108
View File
@@ -0,0 +1,108 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Private first-class Function namespace facades for database connections.
These helpers are internal submission and lookup surfaces. They are not durable
resources and are not part of the public top-level export surface.
"""
from __future__ import annotations
from collections.abc import Callable
from typing import TYPE_CHECKING
from . import _udf
from ._lancedb import Function
from .job import AsyncJob, Job
if TYPE_CHECKING:
from .db import AsyncConnection, DBConnection
class _SyncFunctions:
"""Synchronous `db.functions` facade."""
__slots__ = ("_connection",)
def __init__(self, connection: DBConnection) -> None:
self._connection = connection
def __repr__(self) -> str:
return "_SyncFunctions()"
def register(self, name: str, decorated_udf: Callable[..., object]) -> Job:
"""Register a decorated UDF and return a synchronous [Job][lancedb.job.Job]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = self._connection._submit_register_function(name, definition)
return Job(AsyncJob(native_job))
def replace(
self, name: str, current: Function, decorated_udf: Callable[..., object]
) -> Job:
"""Conditionally replace a Function; return sync [Job][lancedb.job.Job]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = self._connection._submit_replace_function(
name, current, definition
)
return Job(AsyncJob(native_job))
def get(self, name: str) -> Function:
"""Return the Function currently bound to a database-scoped name."""
return self._connection._lookup_function_by_name(name)
def get_by_id(self, function_id: str) -> Function:
"""Return the immutable Function for an exact Function ID."""
return self._connection._lookup_function_by_id(function_id)
def remove(self, name: str, current: Function) -> None:
"""Conditionally remove a Function catalog name binding."""
return self._connection._remove_function_name(name, current)
def revoke(self, function: Function) -> None:
"""Revoke an exact immutable Function by administrator set-bit."""
return self._connection._revoke_function(function)
class _AsyncFunctions:
"""Asynchronous `async_db.functions` facade."""
__slots__ = ("_connection",)
def __init__(self, connection: AsyncConnection) -> None:
self._connection = connection
def __repr__(self) -> str:
return "_AsyncFunctions()"
async def register(
self, name: str, decorated_udf: Callable[..., object]
) -> AsyncJob:
"""Register a decorated UDF and return an [AsyncJob][lancedb.job.AsyncJob]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = await self._connection._register_function(name, definition)
return AsyncJob(native_job)
async def replace(
self, name: str, current: Function, decorated_udf: Callable[..., object]
) -> AsyncJob:
"""Conditionally replace a Function; return [AsyncJob][lancedb.job.AsyncJob]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = await self._connection._replace_function(name, current, definition)
return AsyncJob(native_job)
async def get(self, name: str) -> Function:
"""Return the Function currently bound to a database-scoped name."""
return await self._connection._lookup_function_by_name(name)
async def get_by_id(self, function_id: str) -> Function:
"""Return the immutable Function for an exact Function ID."""
return await self._connection._lookup_function_by_id(function_id)
async def remove(self, name: str, current: Function) -> None:
"""Conditionally remove a Function catalog name binding."""
return await self._connection._remove_function_name(name, current)
async def revoke(self, function: Function) -> None:
"""Revoke an exact immutable Function by administrator set-bit."""
return await self._connection._revoke_function(function)
+59 -13
View File
@@ -153,6 +153,16 @@ class Connection(object):
async def job_history(
self, job_id: Optional[str] = None
) -> List[pa.RecordBatch]: ...
async def _register_function(
self, name: str, definition: "_FunctionDefinition"
) -> Job: ...
async def _replace_function(
self, name: str, current: Function, definition: "_FunctionDefinition"
) -> Job: ...
async def _lookup_function_by_name(self, name: str) -> Function: ...
async def _lookup_function_by_id(self, function_id: str) -> Function: ...
async def _remove_function_name(self, name: str, current: Function) -> None: ...
async def _revoke_function(self, function: Function) -> None: ...
async def create_table(
self,
name: str,
@@ -198,9 +208,6 @@ class Connection(object):
async def drop_table(
self, name: str, namespace_path: Optional[List[str]] = None
) -> None: ...
async def drop_table_async(
self, name: str, namespace_path: Optional[List[str]] = None
) -> Job: ...
async def drop_all_tables(
self, namespace_path: Optional[List[str]] = None
) -> None: ...
@@ -219,11 +226,45 @@ class BlobFile:
def read_range(self, offset: int, length: int) -> bytes: ...
def read_up_to(self, length: int) -> bytes: ...
class Function:
@property
def id(self) -> str: ...
@property
def parameters(self) -> tuple[tuple[str, pa.DataType], ...]: ...
@property
def output_type(self) -> pa.DataType: ...
@property
def output_nullable(self) -> bool: ...
def __call__(self, **kwargs: Any) -> "_FunctionCall": ...
class _FunctionCall:
"""Private unresolved Function call authoring value (FF-028)."""
...
class _FunctionDefinition:
"""Private owner of the Rust FunctionDefinition registration input."""
def _to_json(self) -> str: ...
def _new_function_definition(
*,
parameters: list[tuple[str, pa.DataType]],
output_type: pa.DataType,
output_nullable: bool,
module: str,
callable_name: str,
source: str,
python: str,
packages: list[str],
capabilities: list[tuple[str, str, Optional[str]]],
) -> _FunctionDefinition: ...
class Job:
@property
def id(self) -> Optional[str]: ...
async def status(self) -> str: ...
async def wait(self) -> None: ...
async def wait(self) -> Optional[Function]: ...
async def cancel(self) -> None: ...
class JobInfo:
@@ -245,6 +286,8 @@ class JobFailureInfo:
def message(self) -> Optional[str]: ...
@property
def retryable(self) -> Optional[bool]: ...
@property
def error_code(self) -> Optional[str]: ...
class JobDescription:
@property
@@ -259,6 +302,8 @@ class JobDescription:
def spec_json(self) -> Optional[str]: ...
@property
def failure(self) -> Optional[JobFailureInfo]: ...
@property
def result(self) -> Optional[Function]: ...
class Table:
def name(self) -> str: ...
@@ -321,6 +366,16 @@ class Table:
name: Optional[str],
train: Optional[bool],
) -> Job: ...
async def _add_generated_column(
self, column_name: str, call: _FunctionCall
) -> Job: ...
async def _generated_column_status(
self, column_name: str
) -> Literal["complete", "incomplete"]: ...
async def _refresh_generated_column(self, column_name: str) -> Job: ...
async def _alter_generated_column(
self, column_name: str, new_call: _FunctionCall
) -> Job: ...
async def list_versions(self) -> List[Dict[str, Any]]: ...
async def version(self) -> int: ...
async def checkout(self, version: Union[int, str]): ...
@@ -338,11 +393,6 @@ class Table:
) -> list[FtsToken]: ...
async def delete(self, filter: Union[str, PyExpr]) -> DeleteResult: ...
async def add_columns(self, columns: list[tuple[str, str]]) -> AddColumnsResult: ...
async def add_computed_columns(
self, columns: list[tuple[str, str]]
) -> AddColumnsResult: ...
async def refresh_column(self, column: str) -> RefreshColumnResult: ...
async def refresh_column_async(self, column: str) -> Job: ...
async def add_columns_with_schema(self, schema: pa.Schema) -> AddColumnsResult: ...
async def alter_columns(
self, columns: list[dict[str, Any]]
@@ -688,10 +738,6 @@ class LsmWriteSpec:
class AddColumnsResult:
version: int
class RefreshColumnResult:
rows_filled: int
version: int
class AlterColumnsResult:
version: int
+538
View File
@@ -0,0 +1,538 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Local authoring declaration surface for first-class UDFs.
This module snapshots declaration metadata onto a Python function and privately
validates packagable callables into a source snapshot. It does not mint durable
identity or register anything with a database.
"""
from __future__ import annotations
import ast
import inspect
import stat
import symtable
import sys
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from types import CodeType, FunctionType
from typing import NoReturn, ParamSpec, TypeVar
import pyarrow as pa
from . import _lancedb
__all__ = ["FunctionCapability", "udf"]
_P = ParamSpec("_P")
_R = TypeVar("_R")
_CONFIG_ATTR = "__lancedb_udf_config__"
_SYNTHETIC_SOURCE_FILENAME = "<lancedb-udf>"
_PACKAGING_ERROR = "udf is not packagable"
_ALLOWED_PARAM_KINDS = (
inspect.Parameter.POSITIONAL_OR_KEYWORD,
inspect.Parameter.KEYWORD_ONLY,
)
class FunctionCapability:
"""Local capability declaration for a first-class UDF.
Construct via :meth:`network` or :meth:`secret`. Direct construction is
rejected so callers cannot create an uninitialized capability.
"""
__slots__ = ("_kind", "_origin", "_reference", "_environment_variable")
def __new__(cls, *args: object, **kwargs: object) -> FunctionCapability:
raise TypeError(
"FunctionCapability cannot be constructed directly; "
"use FunctionCapability.network() or FunctionCapability.secret()"
)
@classmethod
def _create(
cls,
kind: str,
origin: str | None,
reference: str | None,
environment_variable: str | None,
) -> FunctionCapability:
obj = object.__new__(cls)
object.__setattr__(obj, "_kind", kind)
object.__setattr__(obj, "_origin", origin)
object.__setattr__(obj, "_reference", reference)
object.__setattr__(obj, "_environment_variable", environment_variable)
return obj
@classmethod
def network(cls, origin: str) -> FunctionCapability:
if not isinstance(origin, str):
raise TypeError("origin must be a string")
if origin == "":
raise ValueError("origin must be non-empty")
return cls._create("network", origin, None, None)
@classmethod
def secret(cls, reference: str, *, environment_variable: str) -> FunctionCapability:
if not isinstance(reference, str):
raise TypeError("reference must be a string")
if not isinstance(environment_variable, str):
raise TypeError("environment_variable must be a string")
if reference == "":
raise ValueError("reference must be non-empty")
if environment_variable == "":
raise ValueError("environment_variable must be non-empty")
return cls._create("secret", None, reference, environment_variable)
@property
def kind(self) -> str:
return self._kind
@property
def origin(self) -> str | None:
return self._origin
@property
def reference(self) -> str | None:
return self._reference
@property
def environment_variable(self) -> str | None:
return self._environment_variable
def __setattr__(self, name: str, value: object) -> None:
raise AttributeError(
f"{type(self).__name__!r} object attribute {name!r} is read-only"
)
def __delattr__(self, name: str) -> None:
raise AttributeError(
f"{type(self).__name__!r} object attribute {name!r} is read-only"
)
def __eq__(self, other: object) -> bool:
if not isinstance(other, FunctionCapability):
return NotImplemented
return (
self._kind == other._kind
and self._origin == other._origin
and self._reference == other._reference
and self._environment_variable == other._environment_variable
)
def __hash__(self) -> int:
return hash(
(
self._kind,
self._origin,
self._reference,
self._environment_variable,
)
)
def __repr__(self) -> str:
if self._kind == "network":
return f"FunctionCapability.network({self._origin!r})"
return (
"FunctionCapability.secret("
f"environment_variable={self._environment_variable!r})"
)
@dataclass(frozen=True, slots=True)
class _UdfConfig:
"""Private frozen snapshot of a ``@udf`` declaration."""
inputs: tuple[tuple[str, pa.DataType], ...]
output: pa.DataType
output_nullable: bool
python: str
packages: tuple[str, ...]
capabilities: tuple[FunctionCapability, ...]
@dataclass(frozen=True, slots=True)
class _PackagedUdf:
"""Private frozen snapshot of a validated packagable UDF."""
source: str
module: str
callable_name: str
config: _UdfConfig
def __repr__(self) -> str:
return (
f"_PackagedUdf(source=<redacted>, module={self.module!r}, "
f"callable_name={self.callable_name!r}, config={self.config!r})"
)
def _validate_inputs(
inputs: object,
) -> tuple[tuple[str, pa.DataType], ...]:
if not isinstance(inputs, Mapping):
raise TypeError("udf inputs must be a Mapping of name to pyarrow DataType")
snapshot: list[tuple[str, pa.DataType]] = []
for key, value in inputs.items():
if not isinstance(key, str):
raise TypeError("udf input names must be strings")
if key == "":
raise ValueError("udf input names must be non-empty")
if not isinstance(value, pa.DataType):
raise TypeError("udf input types must be pyarrow DataType values")
snapshot.append((key, value))
return tuple(snapshot)
def _validate_packages(packages: object) -> tuple[str, ...]:
if isinstance(packages, (str, bytes, bytearray)):
raise TypeError("udf packages must be a sequence of strings, not a string")
if not isinstance(packages, Sequence):
raise TypeError("udf packages must be a sequence of strings")
snapshot: list[str] = []
seen: set[str] = set()
for package in packages:
if not isinstance(package, str):
raise TypeError("udf packages must contain only strings")
if package == "":
raise ValueError("udf packages must be non-empty strings")
if package in seen:
raise ValueError(f"duplicate udf package: {package}")
seen.add(package)
snapshot.append(package)
return tuple(snapshot)
def _reject_non_exact_capability() -> NoReturn:
# Exact-type only: subclasses are authoring inputs we never accept. Keep the
# message fixed so hostile markers never enter exception text.
raise TypeError(
"udf capabilities must contain only FunctionCapability values"
) from None
def _require_exact_capability(capability: object) -> FunctionCapability:
if type(capability) is not FunctionCapability:
_reject_non_exact_capability()
return capability
def _validate_capabilities(capabilities: object) -> tuple[FunctionCapability, ...]:
if isinstance(capabilities, (str, bytes, bytearray)):
raise TypeError(
"udf capabilities must be a sequence of FunctionCapability, not a string"
)
if not isinstance(capabilities, Sequence):
raise TypeError("udf capabilities must be a sequence of FunctionCapability")
return tuple(_require_exact_capability(capability) for capability in capabilities)
def udf(
*,
inputs: Mapping[str, pa.DataType],
output: pa.DataType,
python: str,
packages: Sequence[str] = (),
output_nullable: bool = True,
capabilities: Sequence[FunctionCapability] = (),
) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]:
"""Declare a local UDF without packaging or registration.
Applying the returned decorator attaches a private frozen config snapshot
and returns the exact same function object.
"""
input_snapshot = _validate_inputs(inputs)
if not isinstance(output, pa.DataType):
raise TypeError("udf output must be a pyarrow DataType")
if not isinstance(python, str):
raise TypeError("udf python must be a string")
if python == "":
raise ValueError("udf python must be a non-empty string")
package_snapshot = _validate_packages(packages)
if not isinstance(output_nullable, bool):
raise TypeError("udf output_nullable must be a bool")
capability_snapshot = _validate_capabilities(capabilities)
config = _UdfConfig(
inputs=input_snapshot,
output=output,
output_nullable=output_nullable,
python=python,
packages=package_snapshot,
capabilities=capability_snapshot,
)
def decorator(fn: Callable[_P, _R]) -> Callable[_P, _R]:
if not inspect.isfunction(fn):
raise TypeError("udf can only decorate a Python function")
if hasattr(fn, _CONFIG_ATTR):
raise ValueError("function is already decorated with @udf")
setattr(fn, _CONFIG_ATTR, config)
return fn
return decorator
def _get_udf_config(fn: object) -> _UdfConfig:
"""Return the private declaration snapshot for a ``@udf``-decorated function."""
config = getattr(fn, _CONFIG_ATTR, None)
if config is None:
raise TypeError("function is not decorated with @udf")
if not isinstance(config, _UdfConfig):
raise TypeError("function is not decorated with @udf")
return config
def _packaging_reject() -> NoReturn:
raise ValueError(_PACKAGING_ERROR) from None
def _is_ordinary_function(fn: FunctionType) -> bool:
if fn.__name__ == "<lambda>":
return False
if fn.__qualname__ != fn.__name__:
return False
if inspect.iscoroutinefunction(fn) or inspect.isasyncgenfunction(fn):
return False
if inspect.isgeneratorfunction(fn):
return False
return True
def _resolve_source_path(fn: FunctionType, module: object) -> Path:
try:
fn_source: str | None = inspect.getsourcefile(fn)
except TypeError:
fn_source = None
source_lookup_failed = True
else:
source_lookup_failed = False
if source_lookup_failed:
_packaging_reject()
module_file = vars(module).get("__file__")
if not fn_source or not isinstance(module_file, str) or module_file == "":
_packaging_reject()
try:
resolved_paths: tuple[Path, Path] | None = (
Path(fn_source).resolve(),
Path(module_file).resolve(),
)
except (OSError, RuntimeError):
resolved_paths = None
if resolved_paths is None:
_packaging_reject()
fn_path, module_path = resolved_paths
if fn_path != module_path:
_packaging_reject()
if fn_path.suffix != ".py":
_packaging_reject()
try:
mode: int | None = fn_path.stat().st_mode
except OSError:
mode = None
if mode is None:
_packaging_reject()
if not stat.S_ISREG(mode):
_packaging_reject()
return fn_path
def _validate_source(
source: str, callable_name: str
) -> tuple[CodeType, symtable.SymbolTable]:
try:
module_code = compile(
source,
_SYNTHETIC_SOURCE_FILENAME,
"exec",
optimize=sys.flags.optimize,
)
ast.parse(source, filename=_SYNTHETIC_SOURCE_FILENAME, mode="exec")
table = symtable.symtable(source, _SYNTHETIC_SOURCE_FILENAME, "exec")
parsed: tuple[CodeType, symtable.SymbolTable] | None = (module_code, table)
except Exception:
parsed = None
if parsed is None:
_packaging_reject()
module_code, table = parsed
for child in table.get_children():
if child.get_name() == callable_name and child.get_type() == "function":
return module_code, table
_packaging_reject()
def _source_bound_names(table: symtable.SymbolTable) -> set[str]:
names: set[str] = set()
for symbol in table.get_symbols():
if symbol.is_imported() or symbol.is_assigned() or symbol.is_namespace():
names.add(symbol.get_name())
return names
def _code_fingerprint(code: CodeType) -> tuple[object, ...]:
"""Structural fingerprint ignoring only location/debug fields."""
constants = tuple(
_code_fingerprint(constant) if isinstance(constant, CodeType) else constant
for constant in code.co_consts
)
return (
code.co_name,
getattr(code, "co_qualname", code.co_name),
code.co_argcount,
code.co_posonlyargcount,
code.co_kwonlyargcount,
code.co_flags,
code.co_code,
code.co_names,
code.co_varnames,
code.co_freevars,
code.co_cellvars,
getattr(code, "co_exceptiontable", b""),
constants,
)
def _toplevel_code_candidates(
module_code: CodeType, callable_name: str
) -> list[CodeType]:
candidates: list[CodeType] = []
for constant in module_code.co_consts:
if not isinstance(constant, CodeType):
continue
if constant.co_name != callable_name:
continue
if getattr(constant, "co_qualname", callable_name) != callable_name:
continue
candidates.append(constant)
return candidates
def _validate_loaded_code_matches_source(
fn: FunctionType, module_code: CodeType
) -> None:
candidates = _toplevel_code_candidates(module_code, fn.__name__)
if not candidates:
_packaging_reject()
target = _code_fingerprint(fn.__code__)
if not any(_code_fingerprint(candidate) == target for candidate in candidates):
_packaging_reject()
def _validate_signature(fn: FunctionType, config: _UdfConfig) -> None:
try:
signature: inspect.Signature | None = inspect.signature(fn)
except (TypeError, ValueError):
signature = None
if signature is None:
_packaging_reject()
parameters = list(signature.parameters.values())
expected = [name for name, _ in config.inputs]
actual = [parameter.name for parameter in parameters]
if actual != expected:
_packaging_reject()
for parameter in parameters:
if parameter.kind not in _ALLOWED_PARAM_KINDS:
_packaging_reject()
def _validate_ambient_globals(fn: FunctionType, table: symtable.SymbolTable) -> None:
try:
closure_vars: inspect.ClosureVars | None = inspect.getclosurevars(fn)
except (TypeError, ValueError):
closure_vars = None
if closure_vars is None:
_packaging_reject()
if closure_vars.nonlocals:
_packaging_reject()
bound_names = _source_bound_names(table)
for name in closure_vars.globals:
if name not in bound_names:
_packaging_reject()
def _package_udf(fn: object) -> _PackagedUdf:
"""Validate and snapshot a packagable ``@udf``-decorated function."""
config = _get_udf_config(fn)
if not isinstance(fn, FunctionType) or not _is_ordinary_function(fn):
_packaging_reject()
module_name = fn.__module__
if (
not isinstance(module_name, str)
or module_name == ""
or module_name == "__main__"
):
_packaging_reject()
module = sys.modules.get(module_name)
if module is None:
_packaging_reject()
callable_name = fn.__name__
if vars(module).get(callable_name) is not fn:
_packaging_reject()
source_path = _resolve_source_path(fn, module)
try:
source: str | None = source_path.read_text(encoding="utf-8")
except (OSError, UnicodeError):
source = None
if source is None:
_packaging_reject()
module_code, table = _validate_source(source, callable_name)
_validate_signature(fn, config)
_validate_ambient_globals(fn, table)
_validate_loaded_code_matches_source(fn, module_code)
return _PackagedUdf(
source=source,
module=module_name,
callable_name=callable_name,
config=config,
)
def _normalize_capability_triple(
capability: FunctionCapability,
) -> tuple[str, str, str | None]:
"""Normalize a local capability declaration to the native triple shape."""
# Private config is untrusted; re-check exact type before any property access.
capability = _require_exact_capability(capability)
if capability.kind == "network":
origin = capability.origin
if origin is None:
raise ValueError("invalid network capability") from None
return ("network", origin, None)
if capability.kind == "secret":
reference = capability.reference
environment_variable = capability.environment_variable
if reference is None or environment_variable is None:
raise ValueError("invalid secret capability") from None
return ("secret", reference, environment_variable)
# Fail closed without echoing the unknown kind.
raise ValueError("unsupported capability kind") from None
def _build_function_definition(fn: object) -> _lancedb._FunctionDefinition:
"""Package a ``@udf`` and bridge it to the private native definition."""
packaged = _package_udf(fn)
config = packaged.config
capabilities = [
_normalize_capability_triple(capability) for capability in config.capabilities
]
return _lancedb._new_function_definition(
parameters=list(config.inputs),
output_type=config.output,
output_nullable=config.output_nullable,
module=packaged.module,
callable_name=packaged.callable_name,
source=packaged.source,
python=config.python,
packages=list(config.packages),
capabilities=capabilities,
)
+126 -37
View File
@@ -63,8 +63,12 @@ if TYPE_CHECKING:
import pyarrow as pa
from .pydantic import LanceModel
from ._functions import _AsyncFunctions, _SyncFunctions
from ._lancedb import Connection as LanceDbConnection
from ._lancedb import Function
from ._lancedb import Job as NativeJob
from ._lancedb import JobDescription, JobInfo
from ._lancedb import _FunctionDefinition
from .common import DATA, URI
from .embeddings import EmbeddingFunctionConfig
from ._lancedb import Session
@@ -524,12 +528,6 @@ class DBConnection(EnforceOverrides):
namespace_path = []
raise NotImplementedError
def drop_table_async(
self, name: str, namespace_path: Optional[List[str]] = None
) -> Job:
"""Start dropping a table and return its cleanup job."""
raise NotImplementedError
def rename_table(
self,
cur_name: str,
@@ -656,6 +654,71 @@ class DBConnection(EnforceOverrides):
"job_history is not supported for this connection type"
)
@property
def functions(self) -> "_SyncFunctions":
"""First-class Function operations for this connection."""
from ._functions import _SyncFunctions
return _SyncFunctions(self)
def _submit_register_function(
self, name: str, definition: "_FunctionDefinition"
) -> NativeJob:
"""Submit a Function registration job via the native connection.
Connection subclasses that support registration override this hook.
"""
raise NotImplementedError(
"function registration is not supported for this connection type"
)
def _submit_replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> NativeJob:
"""Submit a Function conditional replace job via the native connection.
Connection subclasses that support registration override this hook.
"""
raise NotImplementedError(
"function replace is not supported for this connection type"
)
def _lookup_function_by_name(self, name: str) -> "Function":
"""Look up a Function by database-scoped name via the native connection.
Connection subclasses that share the native Connection override this hook.
"""
raise NotImplementedError(
"function lookup is not supported for this connection type"
)
def _lookup_function_by_id(self, function_id: str) -> "Function":
"""Look up a Function by exact Function ID via the native connection.
Connection subclasses that share the native Connection override this hook.
"""
raise NotImplementedError(
"function lookup is not supported for this connection type"
)
def _remove_function_name(self, name: str, current: "Function") -> None:
"""Conditionally remove a Function catalog name via the native connection.
Connection subclasses that share the native Connection override this hook.
"""
raise NotImplementedError(
"function name removal is not supported for this connection type"
)
def _revoke_function(self, function: "Function") -> None:
"""Revoke an exact immutable Function via the native connection.
Connection subclasses that share the native Connection override this hook.
"""
raise NotImplementedError(
"function revocation is not supported for this connection type"
)
class LanceDBConnection(DBConnection):
"""
@@ -1192,20 +1255,6 @@ class LanceDBConnection(DBConnection):
)
)
@override
def drop_table_async(
self, name: str, namespace_path: Optional[List[str]] = None
) -> Job:
"""Start dropping a table and return its cleanup job.
The table may become unavailable before its data files are removed.
Call :meth:`Job.wait` to wait for cleanup to finish.
"""
if namespace_path is None:
namespace_path = []
job = LOOP.run(self._conn.drop_table_async(name, namespace_path=namespace_path))
return Job(job if isinstance(job, AsyncJob) else AsyncJob(job))
@override
def drop_all_tables(self, namespace_path: Optional[List[str]] = None):
if namespace_path is None:
@@ -1287,6 +1336,34 @@ class LanceDBConnection(DBConnection):
"""
return LOOP.run(self._conn.job_history(job_id))
@override
def _submit_register_function(
self, name: str, definition: "_FunctionDefinition"
) -> NativeJob:
return LOOP.run(self._conn._register_function(name, definition))
@override
def _submit_replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> NativeJob:
return LOOP.run(self._conn._replace_function(name, current, definition))
@override
def _lookup_function_by_name(self, name: str) -> "Function":
return LOOP.run(self._conn._lookup_function_by_name(name))
@override
def _lookup_function_by_id(self, function_id: str) -> "Function":
return LOOP.run(self._conn._lookup_function_by_id(function_id))
@override
def _remove_function_name(self, name: str, current: "Function") -> None:
return LOOP.run(self._conn._remove_function_name(name, current))
@override
def _revoke_function(self, function: "Function") -> None:
return LOOP.run(self._conn._revoke_function(function))
@override
def namespace_client(self) -> LanceNamespace:
"""Get the equivalent namespace client for this connection.
@@ -1983,23 +2060,6 @@ class AsyncConnection(object):
if f"Table '{name}' was not found" not in str(e):
raise e
async def drop_table_async(
self,
name: str,
*,
namespace_path: Optional[List[str]] = None,
) -> AsyncJob:
"""Start dropping a table and return its cleanup job.
The table may become unavailable before its data files are removed.
Await :meth:`AsyncJob.wait` to wait for cleanup to finish.
"""
if namespace_path is None:
namespace_path = []
return AsyncJob(
await self._inner.drop_table_async(name, namespace_path=namespace_path)
)
async def drop_all_tables(self, namespace_path: Optional[List[str]] = None):
"""Drop all tables from the database.
@@ -2050,6 +2110,35 @@ class AsyncConnection(object):
"""
return await self._inner.job_history(job_id)
@property
def functions(self) -> "_AsyncFunctions":
"""First-class Function operations for this connection."""
from ._functions import _AsyncFunctions
return _AsyncFunctions(self)
async def _register_function(
self, name: str, definition: "_FunctionDefinition"
) -> NativeJob:
return await self._inner._register_function(name, definition)
async def _replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> NativeJob:
return await self._inner._replace_function(name, current, definition)
async def _lookup_function_by_name(self, name: str) -> "Function":
return await self._inner._lookup_function_by_name(name)
async def _lookup_function_by_id(self, function_id: str) -> "Function":
return await self._inner._lookup_function_by_id(function_id)
async def _remove_function_name(self, name: str, current: "Function") -> None:
return await self._inner._remove_function_name(name, current)
async def _revoke_function(self, function: "Function") -> None:
return await self._inner._revoke_function(function)
async def namespace_client(self) -> LanceNamespace:
"""Get the equivalent namespace client for this connection.
+39 -2
View File
@@ -3,6 +3,8 @@
"""Custom exception handling"""
from typing import Optional
class MissingValueError(ValueError):
"""Exception raised when a required value is missing."""
@@ -26,12 +28,47 @@ class MissingColumnError(KeyError):
class JobFailedError(RuntimeError):
"""Exception raised when an asynchronous job reaches the failed state."""
"""Exception raised when an asynchronous job reaches the failed state.
pass
``error_code`` is the optional exact category string projected from the
native job failure when the backend supplied one. The RuntimeError
message remains the existing diagnostic text and must not be used to
recover or override the code.
"""
__slots__ = ("_error_code",)
def __init__(self, message: str, error_code: Optional[str] = None) -> None:
super().__init__(message)
self._error_code = error_code
@property
def error_code(self) -> Optional[str]:
"""Exact job failure error category string, when supplied."""
return self._error_code
class JobCancelledError(RuntimeError):
"""Exception raised when an asynchronous job was cancelled."""
pass
class FunctionError(RuntimeError):
"""Exception raised when a first-class Function operation fails.
``code`` is the stable semantic category from the native error. The
message is a sanitized client diagnostic and must not be used to recover
or override the code.
"""
__slots__ = ("_code",)
def __init__(self, message: str, code: str) -> None:
super().__init__(message)
self._code = code
@property
def code(self) -> str:
"""Stable Function error category string."""
return self._code
+16 -7
View File
@@ -10,6 +10,7 @@ from typing import Optional
from lancedb.background_loop import LOOP
from . import _lancedb
from ._lancedb import Function
class AsyncJob:
@@ -44,18 +45,22 @@ class AsyncJob:
return "finished"
return await self._inner.status()
async def wait(self, timeout: Optional[timedelta] = None):
async def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
"""Wait until the operation reaches a terminal state.
Returns the success result when present (currently a
:class:`~lancedb.Function`), or `None` when the job finished without
one.
Raises `JobFailedError` if the operation failed, `JobCancelledError`
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
"""
if self._inner is None:
return
return None
if timeout is None:
await self._inner.wait()
return await self._inner.wait()
else:
await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
return await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
async def cancel(self):
"""Request cancellation. Cancelling a finished operation is a no-op."""
@@ -88,15 +93,19 @@ class Job:
return "finished"
return LOOP.run(self._inner.status())
def wait(self, timeout: Optional[timedelta] = None):
def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
"""Block until the operation reaches a terminal state.
Returns the success result when present (currently a
:class:`~lancedb.Function`), or `None` when the job finished without
one.
Raises `JobFailedError` if the operation failed, `JobCancelledError`
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
"""
if self._inner is None:
return
LOOP.run(self._inner.wait(timeout))
return None
return LOOP.run(self._inner.wait(timeout))
def cancel(self):
"""Request cancellation. Cancelling a finished operation is a no-op."""
-21
View File
@@ -49,7 +49,6 @@ from lancedb._lancedb import (
)
from lancedb.background_loop import LOOP
from lancedb.db import AsyncConnection, DBConnection
from lancedb.job import AsyncJob, Job
from lance_namespace import (
LanceNamespace,
connect as namespace_connect,
@@ -625,18 +624,6 @@ class LanceNamespaceDBConnection(DBConnection):
namespace_path = []
LOOP.run(self._inner.drop_table(name, namespace_path=namespace_path))
@override
def drop_table_async(
self, name: str, namespace_path: Optional[List[str]] = None
) -> Job:
"""Start dropping a table and return its cleanup job."""
if namespace_path is None:
namespace_path = []
job = LOOP.run(
self._inner.drop_table_async(name, namespace_path=namespace_path)
)
return Job(job if isinstance(job, AsyncJob) else AsyncJob(job))
@override
def rename_table(
self,
@@ -1147,14 +1134,6 @@ class AsyncLanceNamespaceDBConnection:
namespace_path = []
await self._inner.drop_table(name, namespace_path=namespace_path)
async def drop_table_async(
self, name: str, namespace_path: Optional[List[str]] = None
) -> AsyncJob:
"""Start dropping a table and return its cleanup job."""
if namespace_path is None:
namespace_path = []
return await self._inner.drop_table_async(name, namespace_path=namespace_path)
async def rename_table(
self,
cur_name: str,
+3 -9
View File
@@ -2235,7 +2235,6 @@ class LanceHybridQueryBuilder(LanceQueryBuilder):
reranker=self._reranker,
limit=self._limit,
with_row_ids=True,
offset=self._offset,
)
return self._finish_hybrid_results(results)
@@ -2257,7 +2256,6 @@ class LanceHybridQueryBuilder(LanceQueryBuilder):
reranker,
limit: int,
with_row_ids: bool,
offset: Optional[int] = None,
) -> pa.Table:
if norm == "rank":
vector_results = LanceHybridQueryBuilder._rank(vector_results, "_distance")
@@ -2334,7 +2332,7 @@ class LanceHybridQueryBuilder(LanceQueryBuilder):
score_i = results.column_names.index("_score")
results = results.set_column(score_i, "_score", original_scores)
results = results.slice(offset=offset or 0, length=limit)
results = results.slice(length=limit)
if not with_row_ids:
results = results.drop(["_rowid"])
@@ -2681,12 +2679,8 @@ class LanceHybridQueryBuilder(LanceQueryBuilder):
# Apply common configurations
if self._limit:
# The final offset/limit window is sliced out of the combined,
# reranked results, so each sub-query must fetch enough rows to
# cover the skipped prefix as well as the window itself.
sub_query_limit = self._limit + (self._offset or 0)
self._vector_query.limit(sub_query_limit)
self._fts_query.limit(sub_query_limit)
self._vector_query.limit(self._limit)
self._fts_query.limit(self._limit)
if self._columns:
self._vector_query.select(self._columns)
self._fts_query.select(self._columns)
+32 -12
View File
@@ -23,10 +23,12 @@ import pyarrow as pa
from ..common import DATA
from ..db import DBConnection, LOOP
from ..job import AsyncJob, Job
from ..job import Job
if TYPE_CHECKING:
from .._lancedb import JobDescription, JobInfo
from .._lancedb import Function
from .._lancedb import Job as NativeJob
from .._lancedb import JobDescription, JobInfo, _FunctionDefinition
from ..embeddings import EmbeddingFunctionConfig
from lance_namespace import (
LanceNamespace,
@@ -663,16 +665,6 @@ class RemoteDBConnection(DBConnection):
namespace_path = []
LOOP.run(self._conn.drop_table(name, namespace_path=namespace_path))
@override
def drop_table_async(
self, name: str, namespace_path: Optional[List[str]] = None
) -> Job:
"""Start dropping a table and return its cleanup job."""
if namespace_path is None:
namespace_path = []
job = LOOP.run(self._conn.drop_table_async(name, namespace_path=namespace_path))
return Job(job if isinstance(job, AsyncJob) else AsyncJob(job))
@override
def rename_table(
self,
@@ -744,6 +736,34 @@ class RemoteDBConnection(DBConnection):
"""
return LOOP.run(self._conn.job_history(job_id))
@override
def _submit_register_function(
self, name: str, definition: "_FunctionDefinition"
) -> "NativeJob":
return LOOP.run(self._conn._register_function(name, definition))
@override
def _submit_replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> "NativeJob":
return LOOP.run(self._conn._replace_function(name, current, definition))
@override
def _lookup_function_by_name(self, name: str) -> "Function":
return LOOP.run(self._conn._lookup_function_by_name(name))
@override
def _lookup_function_by_id(self, function_id: str) -> "Function":
return LOOP.run(self._conn._lookup_function_by_id(function_id))
@override
def _remove_function_name(self, name: str, current: "Function") -> None:
return LOOP.run(self._conn._remove_function_name(name, current))
@override
def _revoke_function(self, function: "Function") -> None:
return LOOP.run(self._conn._revoke_function(function))
@override
def namespace_client(self) -> LanceNamespace:
"""Get the equivalent namespace client for this connection.
+47 -37
View File
@@ -7,6 +7,7 @@ import logging
from functools import cached_property
import os
from typing import (
TYPE_CHECKING,
Any,
Callable,
Dict,
@@ -67,6 +68,9 @@ from ..query import (
from ..table import AsyncTable, BlobMode, Branches, IndexStatistics, Query, Table, Tags
from ..types import BaseTokenizerType
if TYPE_CHECKING:
from lancedb._lancedb import _FunctionCall
class RemoteTable(Table):
def __init__(
@@ -570,6 +574,45 @@ class RemoteTable(Table):
)
)
def add_generated_column(self, column_name: str, call: "_FunctionCall") -> Job:
"""Add a generated column from an authored Function call.
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
of the Job does not publish the column; callers must wait and re-read
the table to observe the new definition and values.
"""
return Job(LOOP.run(self._table.add_generated_column(column_name, call)))
def generated_column_status(
self, column_name: str
) -> Literal["complete", "incomplete"]:
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
Projection-only: reads the named column's stored definition status.
Does not refresh values, submit a Job, or mutate table state.
"""
return LOOP.run(self._table.generated_column_status(column_name))
def refresh_generated_column(self, column_name: str) -> Job:
"""Refresh values for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
of the Job does not publish new values; callers must wait and re-read
the table to observe refreshed results.
"""
return Job(LOOP.run(self._table.refresh_generated_column(column_name)))
def alter_generated_column(
self, column_name: str, new_call: "_FunctionCall"
) -> Job:
"""Alter the Function call for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
of the Job does not publish the new definition; callers must wait and
re-read the table to observe the updated column.
"""
return Job(LOOP.run(self._table.alter_generated_column(column_name, new_call)))
def _is_legacy_create_index_call(
self,
first_arg: str,
@@ -958,19 +1001,8 @@ class RemoteTable(Table):
def count_rows(self, filter: Optional[str] = None) -> int:
return LOOP.run(self._table.count_rows(filter))
def add_columns(
self,
transforms: Dict[str, str] | None = None,
*,
computed: Dict[str, str] | None = None,
) -> AddColumnsResult:
return LOOP.run(self._table.add_columns(transforms, computed=computed))
def refresh_column(self, column: str):
return LOOP.run(self._table.refresh_column(column))
def refresh_column_async(self, column: str) -> Job:
return Job(LOOP.run(self._table.refresh_column_async(column)))
def add_columns(self, transforms: Dict[str, str]) -> AddColumnsResult:
return LOOP.run(self._table.add_columns(transforms))
def alter_columns(
self, *alterations: Iterable[Dict[str, str]]
@@ -990,39 +1022,17 @@ class RemoteTable(Table):
return LOOP.run(self._table.set_unenforced_primary_key(columns))
def set_lsm_write_spec(self, spec: "LsmWriteSpec") -> None:
"""Install an LsmWriteSpec."""
"""Not supported on LanceDB Cloud."""
return LOOP.run(self._table.set_lsm_write_spec(spec))
def unset_lsm_write_spec(self) -> None:
"""Remove the LsmWriteSpec."""
"""Not supported on LanceDB Cloud."""
return LOOP.run(self._table.unset_lsm_write_spec())
def get_lsm_write_spec(self) -> Optional["LsmWriteSpec"]:
"""Read the installed LsmWriteSpec, or ``None``."""
return LOOP.run(self._table.get_lsm_write_spec())
def checkpoint_lsm(self) -> None:
"""Synchronous version of
[`AsyncTable.checkpoint_lsm`][lancedb.AsyncTable.checkpoint_lsm]."""
return LOOP.run(self._table.checkpoint_lsm())
def flush_lsm(self) -> None:
"""Synchronous version of
[`AsyncTable.flush_lsm`][lancedb.AsyncTable.flush_lsm]."""
return LOOP.run(self._table.flush_lsm())
def compact_lsm(self) -> None:
"""Synchronous version of
[`AsyncTable.compact_lsm`][lancedb.AsyncTable.compact_lsm]."""
return LOOP.run(self._table.compact_lsm())
def get_lsm_stats(self, *, include_generation_rows: bool = False) -> Optional[dict]:
"""Synchronous version of
[`AsyncTable.get_lsm_stats`][lancedb.AsyncTable.get_lsm_stats]."""
return LOOP.run(
self._table.get_lsm_stats(include_generation_rows=include_generation_rows)
)
def close_lsm_writers(self) -> None:
"""No-op on LanceDB Cloud (no local shard writers)."""
return LOOP.run(self._table.close_lsm_writers())
+127 -202
View File
@@ -176,7 +176,6 @@ if TYPE_CHECKING:
CompactionStats,
Tag,
AddColumnsResult,
RefreshColumnResult,
AddResult,
AlterColumnsResult,
UpdateFieldMetadataResult,
@@ -186,6 +185,7 @@ if TYPE_CHECKING:
LsmWriteSpec,
MergeResult,
UpdateResult,
_FunctionCall,
)
from .index import IndexConfig
import pandas
@@ -1008,6 +1008,43 @@ class Table(ABC):
"""
raise NotImplementedError
def add_generated_column(self, column_name: str, call: _FunctionCall) -> Job:
"""Add a generated column from an authored Function call.
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
of the Job does not publish the column; callers must wait and re-read
the table to observe the new definition and values.
"""
raise NotImplementedError
def generated_column_status(
self, column_name: str
) -> Literal["complete", "incomplete"]:
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
Projection-only: reads the named column's stored definition status.
Does not refresh values, submit a Job, or mutate table state.
"""
raise NotImplementedError
def refresh_generated_column(self, column_name: str) -> Job:
"""Refresh values for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
of the Job does not publish new values; callers must wait and re-read
the table to observe refreshed results.
"""
raise NotImplementedError
def alter_generated_column(self, column_name: str, new_call: _FunctionCall) -> Job:
"""Alter the Function call for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
of the Job does not publish the new definition; callers must wait and
re-read the table to observe the updated column.
"""
raise NotImplementedError
def drop_index(self, name: str) -> None:
"""
Drop an index from the table.
@@ -1917,14 +1954,7 @@ class Table(ABC):
@abstractmethod
def add_columns(
self,
transforms: Dict[str, str]
| pa.Field
| List[pa.Field]
| pa.Schema
| None = None,
*,
computed: Dict[str, str] | None = None,
self, transforms: Dict[str, str] | pa.Field | List[pa.Field] | pa.Schema
):
"""
Add new columns with defined values.
@@ -1938,95 +1968,11 @@ class Table(ABC):
Alternatively, a pyarrow Field or Schema can be provided to add
new columns with the specified data types. The new columns will
be initialized with null values.
computed: Dict[str, str], optional
A map of column name to a SQL expression defining the column. The
column's type and inputs are derived from the expression, so no
data type is supplied.
Unlike ``transforms``, the expression is stored rather than
evaluated now: the column is committed with no values, and rows get
them from [`refresh_column`][lancedb.table.Table.refresh_column].
Declaring one therefore costs the same on a large table as on an
empty one.
A refresh does not revisit rows it has already filled, so mutating
an input leaves the value computed at fill time; recomputing means
dropping the column and declaring it again. While a declaration
reads a column, that column cannot be renamed, retyped or dropped.
On LanceDB Cloud and Enterprise the expression is planned by the
server, and the refresh runs as a server job -- see
[`refresh_column_async`][lancedb.table.Table.refresh_column_async].
Cannot be combined with ``transforms``.
Returns
-------
AddColumnsResult
version: the new version number of the table after adding columns.
Examples
--------
>>> import lancedb
>>> db = lancedb.connect("./.lancedb")
>>> table = db.create_table("computed_demo", [{"x": 1}, {"x": 2}])
>>> table.add_columns(computed={"doubled": "x * 2"})
AddColumnsResult(version=2)
>>> table.refresh_column("doubled")
RefreshColumnResult(rows_filled=2, version=3)
>>> table.to_arrow().sort_by("x").to_pandas()
x doubled
0 1 2
1 2 4
"""
@abstractmethod
def refresh_column(self, column: str) -> "RefreshColumnResult":
"""
Fill the rows of a computed column that hold no value yet.
Declared with ``add_columns(computed=...)``, a column starts empty and
gets its values here. Rows appended since the last refresh are filled
by the next one; rows already filled are left as they are, so the call
is idempotent and does not observe a mutated input.
Local tables only: a remote refresh runs as a server job, through
[`refresh_column_async`][lancedb.table.Table.refresh_column_async].
Parameters
----------
column: str
The name of the computed column to fill.
Returns
-------
RefreshColumnResult
rows_filled: the number of rows given a value.
version: the new version number of the table.
"""
@abstractmethod
def refresh_column_async(self, column: str) -> Job:
"""
Like :meth:`refresh_column`, but returns a handle to the refresh job
instead of blocking until it completes.
The job may already be complete when returned; callers must not assume
the column is filled until :meth:`Job.wait` returns. Invalid input --
an unknown column, or one that is not computed -- raises here rather
than failing the job. On local tables the job runs in-process; on
LanceDB Cloud and Enterprise it is the server's backfill job.
Examples
--------
>>> import lancedb
>>> db = lancedb.connect("./.lancedb")
>>> table = db.create_table("computed_job_demo", [{"x": 1}, {"x": 2}])
>>> table.add_columns(computed={"doubled": "x * 2"})
AddColumnsResult(version=2)
>>> job = table.refresh_column_async("doubled")
>>> job.wait()
>>> job.status()
'finished'
"""
@abstractmethod
@@ -2941,6 +2887,43 @@ class LanceTable(Table):
)
)
def add_generated_column(self, column_name: str, call: _FunctionCall) -> Job:
"""Add a generated column from an authored Function call.
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
of the Job does not publish the column; callers must wait and re-read
the table to observe the new definition and values.
"""
return Job(LOOP.run(self._table.add_generated_column(column_name, call)))
def generated_column_status(
self, column_name: str
) -> Literal["complete", "incomplete"]:
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
Projection-only: reads the named column's stored definition status.
Does not refresh values, submit a Job, or mutate table state.
"""
return LOOP.run(self._table.generated_column_status(column_name))
def refresh_generated_column(self, column_name: str) -> Job:
"""Refresh values for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
of the Job does not publish new values; callers must wait and re-read
the table to observe refreshed results.
"""
return Job(LOOP.run(self._table.refresh_generated_column(column_name)))
def alter_generated_column(self, column_name: str, new_call: _FunctionCall) -> Job:
"""Alter the Function call for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
of the Job does not publish the new definition; callers must wait and
re-read the table to observe the updated column.
"""
return Job(LOOP.run(self._table.alter_generated_column(column_name, new_call)))
def _is_legacy_create_index_call(
self,
first_arg: str,
@@ -4031,28 +4014,9 @@ class LanceTable(Table):
return LOOP.run(self._table.index_stats(index_name))
def add_columns(
self,
transforms: Dict[str, str]
| pa.field
| List[pa.field]
| pa.Schema
| None = None,
*,
computed: Dict[str, str] | None = None,
self, transforms: Dict[str, str] | pa.field | List[pa.field] | pa.Schema
) -> AddColumnsResult:
return LOOP.run(self._table.add_columns(transforms, computed=computed))
def refresh_column(self, column: str) -> "RefreshColumnResult":
"""Fill a computed column's unfilled rows. See
[`AsyncTable.refresh_column`][lancedb.AsyncTable.refresh_column]."""
return LOOP.run(self._table.refresh_column(column))
def refresh_column_async(self, column: str) -> Job:
"""Fill a computed column's unfilled rows, returning a handle to the
refresh job. See
[`Table.refresh_column_async`][lancedb.table.Table.refresh_column_async].
"""
return Job(LOOP.run(self._table.refresh_column_async(column)))
return LOOP.run(self._table.add_columns(transforms))
def alter_columns(
self, *alterations: Iterable[Dict[str, str]]
@@ -4801,7 +4765,7 @@ class AsyncTable:
Examples
--------
>>> from lancedb import LsmWriteSpec
>>> from lancedb._lancedb import LsmWriteSpec
>>> # table.set_unenforced_primary_key("id")
>>> # table.set_lsm_write_spec(LsmWriteSpec.bucket("id", 16))
"""
@@ -4846,7 +4810,7 @@ class AsyncTable:
``asyncio.wait_for`` for a wall-clock bound; abandoning it partway
costs nothing.
"""
await self._inner.checkpoint_lsm()
return await self._inner.checkpoint_lsm()
async def flush_lsm(self) -> None:
"""Seal every bucket's active memtable into L0.
@@ -4855,7 +4819,7 @@ class AsyncTable:
`compact_lsm`. On a node that has not claimed this table, this claims
it and replays its WAL log first.
"""
await self._inner.flush_lsm()
return await self._inner.flush_lsm()
async def compact_lsm(self) -> None:
"""Trigger a background L0 to base compaction pass per bucket.
@@ -4864,7 +4828,7 @@ class AsyncTable:
``get_lsm_stats`` for progress, or use ``checkpoint_lsm`` to loop
until the current L0 has reached base.
"""
await self._inner.compact_lsm()
return await self._inner.compact_lsm()
async def get_lsm_stats(
self, *, include_generation_rows: bool = False
@@ -5177,6 +5141,50 @@ class AsyncTable:
)
return AsyncJob(job)
async def add_generated_column(
self, column_name: str, call: _FunctionCall
) -> AsyncJob:
"""Add a generated column from an authored Function call.
Returns an :class:`~lancedb.job.AsyncJob` for the create operation.
Acceptance of the Job does not publish the column; callers must wait
and re-read the table to observe the new definition and values.
"""
job = await self._inner._add_generated_column(column_name, call)
return AsyncJob(job)
async def generated_column_status(
self, column_name: str
) -> Literal["complete", "incomplete"]:
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
Projection-only: reads the named column's stored definition status.
Does not refresh values, submit a Job, or mutate table state.
"""
return await self._inner._generated_column_status(column_name)
async def refresh_generated_column(self, column_name: str) -> AsyncJob:
"""Refresh values for an existing generated column.
Returns an :class:`~lancedb.job.AsyncJob` for the refresh operation.
Acceptance of the Job does not publish new values; callers must wait
and re-read the table to observe refreshed results.
"""
job = await self._inner._refresh_generated_column(column_name)
return AsyncJob(job)
async def alter_generated_column(
self, column_name: str, new_call: _FunctionCall
) -> AsyncJob:
"""Alter the Function call for an existing generated column.
Returns an :class:`~lancedb.job.AsyncJob` for the change operation.
Acceptance of the Job does not publish the new definition; callers must
wait and re-read the table to observe the updated column.
"""
job = await self._inner._alter_generated_column(column_name, new_call)
return AsyncJob(job)
async def drop_index(self, name: str) -> None:
"""
Drop an index from the table.
@@ -5967,14 +5975,7 @@ class AsyncTable:
return await self._inner.update(updates_sql, where)
async def add_columns(
self,
transforms: dict[str, str]
| pa.field
| List[pa.field]
| pa.Schema
| None = None,
*,
computed: dict[str, str] | None = None,
self, transforms: dict[str, str] | pa.field | List[pa.field] | pa.Schema
) -> AddColumnsResult:
"""
Add new columns with defined values.
@@ -5987,22 +5988,6 @@ class AsyncTable:
each row in the table, and can reference existing columns.
Alternatively, you can pass a pyarrow field or schema to add
new columns with NULLs.
computed: Dict[str, str], optional
A map of column name to a SQL expression defining the column. The
column's type and inputs are derived from the expression.
Unlike ``transforms``, the expression is stored rather than
evaluated now: the column is committed with no values, and rows get
them from
[`refresh_column`][lancedb.table.AsyncTable.refresh_column].
A refresh does not revisit rows it has already filled, so mutating
an input leaves the value computed at fill time. While a
declaration reads a column, that column cannot be renamed, retyped
or dropped.
On LanceDB Cloud and Enterprise the expression is planned by
the server. Cannot be combined with ``transforms``.
Returns
-------
@@ -6016,71 +6001,11 @@ class AsyncTable:
{isinstance(f, pa.Field) for f in transforms}
):
transforms = pa.schema(transforms)
if computed:
if transforms:
raise ValueError(
"add_columns cannot take both transforms and computed columns"
)
return await self._inner.add_computed_columns(list(computed.items()))
if transforms is None:
raise ValueError("add_columns requires transforms or computed columns")
if isinstance(transforms, pa.Schema):
return await self._inner.add_columns_with_schema(transforms)
else:
return await self._inner.add_columns(list(transforms.items()))
async def refresh_column(self, column: str) -> RefreshColumnResult:
"""
Fill the rows of a computed column that hold no value yet.
Declared with ``add_columns(computed=...)``, a column starts empty and
gets its values here. Rows appended since the last refresh are filled
by the next one; rows already filled are left as they are, so the call
is idempotent and does not observe a mutated input.
Local tables only: a remote refresh runs as a server job, through
[`refresh_column_async`][lancedb.table.Table.refresh_column_async].
Parameters
----------
column: str
The name of the computed column to fill.
Returns
-------
RefreshColumnResult
The number of rows filled and the new version of the table.
"""
return await self._inner.refresh_column(column)
async def refresh_column_async(self, column: str) -> AsyncJob:
"""
Like :meth:`refresh_column`, but returns a handle to the refresh job
instead of blocking until it completes.
The job may already be complete when returned; callers must not assume
the column is filled until :meth:`AsyncJob.wait` resolves. Invalid
input -- an unknown column, or one that is not computed -- raises here
rather than failing the job. On local tables the job runs
in-process; on LanceDB Cloud and Enterprise it is the server's
backfill job.
Examples
--------
>>> import asyncio
>>> import lancedb
>>> async def refresh_in_background():
... db = await lancedb.connect_async("./.lancedb")
... table = await db.create_table("computed_job_async_demo", [{"x": 1}])
... await table.add_columns(computed={"doubled": "x * 2"})
... job = await table.refresh_column_async("doubled")
... await job.wait()
... return await job.status()
>>> asyncio.run(refresh_in_background())
'finished'
"""
return AsyncJob(await self._inner.refresh_column_async(column))
async def alter_columns(
self, *alterations: Iterable[dict[str, Any]]
) -> AlterColumnsResult:
+3 -16
View File
@@ -755,7 +755,8 @@ def test_delete_table(tmp_db: lancedb.DBConnection):
assert tmp_db.table_names() == []
def test_drop_table_async(tmp_db: lancedb.DBConnection):
@pytest.mark.asyncio
async def test_delete_table_async(tmp_db: lancedb.DBConnection):
data = pd.DataFrame(
{
"vector": [[3.1, 4.1], [5.9, 26.5]],
@@ -771,10 +772,7 @@ def test_drop_table_async(tmp_db: lancedb.DBConnection):
assert tmp_db.table_names() == ["test"]
job = tmp_db.drop_table_async("test")
assert job.id is None
assert job.status() == "finished"
job.wait()
tmp_db.drop_table("test")
assert tmp_db.table_names() == []
tmp_db.create_table("test", data=data)
@@ -783,17 +781,6 @@ def test_drop_table_async(tmp_db: lancedb.DBConnection):
tmp_db.drop_table("does_not_exist", ignore_missing=True)
@pytest.mark.asyncio
async def test_drop_table_async_connection(tmp_db_async: lancedb.AsyncConnection):
await tmp_db_async.create_table("test", data=pa.table({"id": [1, 2]}))
job = await tmp_db_async.drop_table_async("test")
assert job.id is None
assert await job.status() == "finished"
await job.wait()
assert await tmp_db_async.table_names() == []
def test_drop_database(tmp_db: lancedb.DBConnection):
data = pd.DataFrame(
{
-98
View File
@@ -632,101 +632,3 @@ class TestExprBytesIntegration:
.to_arrow()
)
assert result.num_rows == 2
# ── datetime / timezone integration for lit() (issue #3262) ──────────────────
class TestExprDatetimeTimezoneIntegration:
"""Integration coverage for lit(datetime) against table timestamp columns.
PyArrow stores naive timestamps as UTC wall-clock microseconds. Python's
datetime.timestamp() treats naive values as *local* time, which used to
shift lit(naive) by the host UTC offset and break equality filters on
non-UTC machines. These cases lock the expected semantics.
"""
def test_both_naive_match(self, tmp_path):
"""Table naive + lit naive with the same wall clock must match."""
db = lancedb.connect(str(tmp_path / "naive"))
ts = datetime(2024, 7, 1, 10, 0, 0)
table = db.create_table(
"t", [{"id": 1, "ts": ts}, {"id": 2, "ts": datetime(2024, 7, 2, 10, 0, 0)}]
)
result = table.search().where(col("ts") == lit(ts)).to_list()
assert len(result) == 1
assert result[0]["id"] == 1
def test_both_same_timezone_match(self, tmp_path):
"""Table UTC + lit UTC for the same instant must match."""
db = lancedb.connect(str(tmp_path / "utc"))
ts = datetime(2024, 7, 1, 10, 0, 0, tzinfo=timezone.utc)
table = db.create_table(
"t",
pa.table(
{
"id": [1, 2],
"ts": pa.array(
[ts, datetime(2024, 7, 2, 10, 0, 0, tzinfo=timezone.utc)],
type=pa.timestamp("us", tz="UTC"),
),
}
),
)
result = table.search().where(col("ts") == lit(ts)).to_list()
assert len(result) == 1
assert result[0]["id"] == 1
def test_different_timezones_same_instant(self, tmp_path):
"""UTC table row equals lit of the same instant in a different zone."""
db = lancedb.connect(str(tmp_path / "diff_tz"))
ts_utc = datetime(2024, 7, 1, 10, 0, 0, tzinfo=timezone.utc)
# Same instant as 06:00 in UTC-4
ts_est = datetime(2024, 7, 1, 6, 0, 0, tzinfo=timezone(timedelta(hours=-4)))
table = db.create_table(
"t",
pa.table(
{
"id": [1],
"ts": pa.array([ts_utc], type=pa.timestamp("us", tz="UTC")),
}
),
)
result = table.search().where(col("ts") == lit(ts_est)).to_list()
assert len(result) == 1
assert result[0]["id"] == 1
def test_table_tz_literal_naive(self, tmp_path):
"""UTC table + naive lit uses wall-clock equality (10:00 == 10:00 UTC)."""
db = lancedb.connect(str(tmp_path / "tz_naive"))
ts_utc = datetime(2024, 7, 1, 10, 0, 0, tzinfo=timezone.utc)
ts_naive = datetime(2024, 7, 1, 10, 0, 0)
table = db.create_table(
"t",
pa.table(
{
"id": [1],
"ts": pa.array([ts_utc], type=pa.timestamp("us", tz="UTC")),
}
),
)
result = table.search().where(col("ts") == lit(ts_naive)).to_list()
assert len(result) == 1
assert result[0]["id"] == 1
def test_table_naive_literal_aware(self, tmp_path):
"""Naive table + UTC lit with the same wall clock must match."""
db = lancedb.connect(str(tmp_path / "naive_aware"))
ts_naive = datetime(2024, 7, 1, 10, 0, 0)
ts_utc = datetime(2024, 7, 1, 10, 0, 0, tzinfo=timezone.utc)
table = db.create_table("t", [{"id": 1, "ts": ts_naive}])
result = table.search().where(col("ts") == lit(ts_utc)).to_list()
assert len(result) == 1
assert result[0]["id"] == 1
def test_naive_lit_sql_is_wall_clock_not_local_shifted(self):
"""Regression: naive lit must not apply the host local UTC offset."""
ts = datetime(2024, 7, 1, 10, 0, 0)
sql = lit(ts).to_sql()
# Must encode 10:00 wall clock, not 10:00+local_offset.
assert "2024-07-01 10:00:00" in sql
@@ -0,0 +1,372 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python exact Function handle call authoring (FF-028)."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import Iterator
from typing import Any, Callable
import pyarrow as pa
import pytest
import lancedb
from lancedb import _lancedb as _native
from lancedb.expr import Expr, col, func, lit
_CALL_PATH = "/v1/functions/lookup"
_CALL_CATALOG_NAME = "text.normalize.call-name"
_CALL_FUNCTION_ID = "fn.exact.call-handle"
_LITERAL_PAYLOAD_SENTINEL = "LITERAL_PAYLOAD_SENTINEL_call_xyz_42"
_INT_PAYLOAD_SENTINEL = 2_147_000_123
# Pinned Rust-canonical schema-only type IPC (base64).
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_LIST_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////+4AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAABAAAANz///8c"
"AAAADAAAAAAAAQxcAAAAAQAAABwAAAAEAAQABAAAABAAFAAQAA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECH"
"AAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAQAAABpdGVtAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////8AAAAAFAAAAAAAAAAMABQAEgAMAAgABAAMAAAAnAAAAKAAAAAQAAAAAAAEAAgACAAAAAQACAAAAAQAAAA"
"BAAAABAAAANz///8cAAAADAAAAAAAAQxcAAAAAQAAABwAAAAEAAQABAAAABAAFAAQAA4ADwAEAAAACAAQAAAA"
"GAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAQAAABpdGVtAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAAAwAAAAEFSUk9XMQ=="
)
_OVERDESIGN_ATTRS = (
"id",
"function_id",
"name",
"connection",
"table",
"snapshot",
"field_id",
"field_ids",
"job",
"job_id",
"artifact",
"digest",
"retry_key",
"idempotency_key",
"user_version",
"execute",
"status",
"wait",
"cancel",
"to_json",
"_to_json",
"serialize",
"geneva",
)
def _sample_function_wire(
*,
function_id: str = _CALL_FUNCTION_ID,
parameters: list[dict[str, str]] | None = None,
output_type_ipc: str = _UTF8_TYPE_IPC_B64,
) -> dict[str, Any]:
return {
"format_version": 1,
"id": function_id,
"signature": {
"parameters": parameters
or [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": output_type_ipc,
"nullable": True,
},
},
}
def _lookup_success_body(function: dict[str, Any] | None = None) -> bytes:
return json.dumps({"function": function or _sample_function_wire()}).encode("utf-8")
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _lookup_function(function: dict[str, Any] | None = None):
body = _lookup_success_body(function)
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _CALL_PATH
_read_body(request)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(body)
with _mock_remote_db(handler) as db:
return db.functions.get(_CALL_CATALOG_NAME)
def _authored_call_type():
cls = getattr(_native, "_FunctionCall", None)
if cls is None:
pytest.fail("lancedb._lancedb._FunctionCall is missing")
return cls
def _exception_text(exc: BaseException) -> str:
parts = [str(exc), repr(exc)]
current: BaseException | None = exc
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(f"{type(current).__name__}: {current}")
current = current.__cause__ or current.__context__
return "\n".join(parts)
def test_function_keyword_call_returns_private_frozen_authored_value():
function = _lookup_function()
assert callable(function)
authored = function(text=col("text"), limit=8)
authored_type = _authored_call_type()
assert type(authored) is authored_type
assert authored_type.__module__ == "lancedb._lancedb"
assert authored_type.__name__ == "_FunctionCall"
# Keyword order must not matter; bindings store/render in signature order.
authored_reversed = function(limit=8, text=col("text"))
assert type(authored_reversed) is authored_type
rendered = repr(authored_reversed)
assert rendered.index("text=") < rendered.index("limit=")
assert 'text=field("text")' in rendered
assert "limit=literal(Int32, null=false)" in rendered
def test_function_call_rejects_positional_missing_and_unknown_args():
function = _lookup_function()
with pytest.raises(TypeError, match="keyword"):
function(col("text"), 8)
with pytest.raises((TypeError, ValueError), match="limit"):
function(text=col("text"))
with pytest.raises((TypeError, ValueError), match="text"):
function(limit=8)
with pytest.raises((TypeError, ValueError), match="unknown|extra"):
function(text=col("text"), limit=8, extra=1)
def test_function_call_accepts_direct_case_sensitive_column_and_rejects_complex_exprs():
function = _lookup_function()
authored = function(text=col("firstName"), limit=1)
assert type(authored) is _authored_call_type()
rendered = repr(authored)
assert 'text=field("firstName")' in rendered
assert "limit=literal(Int32, null=false)" in rendered
complex_exprs = (
col("text") + lit("x"),
col("text").cast(pa.string()),
func("lower", col("text")),
col("text") == lit("x"),
col("text").lower(),
)
for expr in complex_exprs:
with pytest.raises((TypeError, ValueError)):
function(text=expr, limit=1)
# Raw native PyExpr is not the public col() wrapper.
with pytest.raises((TypeError, ValueError)):
function(text=col("text")._inner, limit=1)
# Non-expression / non-literal objects are rejected for field-shaped misuse
# when a column binding is required; plain strings are literals for utf8.
with pytest.raises((TypeError, ValueError)):
function(text=object(), limit=1)
def test_function_call_plain_literal_declared_type_null_and_nested():
function = _lookup_function()
authored = function(text="hello", limit=7)
assert type(authored) is _authored_call_type()
rendered = repr(authored)
assert "text=literal(Utf8, null=false)" in rendered
assert "limit=literal(Int32, null=false)" in rendered
# Plain Python int normalizes to declared Int32 and non-null.
authored_int32 = function(text="hello", limit=2_147_483_647)
assert type(authored_int32) is _authored_call_type()
rendered_int32 = repr(authored_int32)
assert "limit=literal(Int32, null=false)" in rendered_int32
assert "Int64" not in rendered_int32
assert "2147483647" not in rendered_int32
# Plain None keeps each declared parameter type with null=true.
authored_null = function(text=None, limit=None)
assert type(authored_null) is _authored_call_type()
rendered_null = repr(authored_null)
assert "text=literal(Utf8, null=true)" in rendered_null
assert "limit=literal(Int32, null=true)" in rendered_null
list_function = _lookup_function(
_sample_function_wire(
parameters=[
{"name": "values", "data_type_ipc": _LIST_INT32_TYPE_IPC_B64},
]
)
)
authored_list = list_function(values=[1, 2, 3])
assert type(authored_list) is _authored_call_type()
rendered_list = repr(authored_list)
assert "values=literal(List(Int32), null=false)" in rendered_list
assert "[1, 2, 3]" not in rendered_list
authored_list_null = list_function(values=None)
assert type(authored_list_null) is _authored_call_type()
rendered_list_null = repr(authored_list_null)
assert "values=literal(List(Int32), null=true)" in rendered_list_null
def test_function_call_direct_literal_expr_exact_type_only():
function = _lookup_function()
# lit(int) is Int64 in the expression builder; int32 parameter must reject it.
with pytest.raises((TypeError, ValueError), match="limit|int32|type") as raised:
function(text="hello", limit=lit(8))
reject_text = _exception_text(raised.value)
assert "Int64" in reject_text or "int64" in reject_text.lower()
assert "Int32" in reject_text or "int32" in reject_text.lower()
# Exact utf8 literal expression is accepted and stored as Utf8/non-null.
authored = function(text=lit("hello"), limit=8)
assert type(authored) is _authored_call_type()
rendered = repr(authored)
assert "text=literal(Utf8, null=false)" in rendered
assert "limit=literal(Int32, null=false)" in rendered
assert "hello" not in rendered
# Cast / arithmetic around a literal is not a direct Literal node.
with pytest.raises((TypeError, ValueError)):
function(text=lit("hello").cast(pa.string()), limit=8)
def test_function_call_conversion_error_and_repr_are_payload_free():
function = _lookup_function()
with pytest.raises((TypeError, ValueError)) as raised:
function(text="ok", limit=_LITERAL_PAYLOAD_SENTINEL)
text = _exception_text(raised.value)
assert _LITERAL_PAYLOAD_SENTINEL not in text
assert "limit" in text
assert "int32" in text.lower() or "Int32" in text
authored = function(text=_LITERAL_PAYLOAD_SENTINEL, limit=_INT_PAYLOAD_SENTINEL)
rendered = f"{authored!r}\n{authored!s}"
assert _LITERAL_PAYLOAD_SENTINEL not in rendered
assert str(_INT_PAYLOAD_SENTINEL) not in rendered
assert "text=literal(Utf8, null=false)" in rendered
assert "limit=literal(Int32, null=false)" in rendered
assert type(authored).__name__ == "_FunctionCall"
assert "_FunctionCall" in rendered
def test_function_call_private_type_nonconstructible_immutable_and_not_exported():
function = _lookup_function()
authored = function(text=col("text"), limit=1)
authored_type = _authored_call_type()
assert "_FunctionCall" not in getattr(lancedb, "__all__", [])
assert not hasattr(lancedb, "_FunctionCall")
assert getattr(_native, "_FunctionCall", None) is authored_type
with pytest.raises(TypeError):
authored_type()
for attr in _OVERDESIGN_ATTRS:
assert not hasattr(authored, attr)
for attr in ("function", "bindings", "arguments", "parameters", "text", "limit"):
with pytest.raises(AttributeError):
setattr(authored, attr, None)
# Existing Function handle stays frozen / connection-free / name-free.
assert not hasattr(function, "name")
assert not hasattr(function, "connection")
with pytest.raises(AttributeError):
function.id = "mutated"
def test_function_call_does_not_change_col_query_expression_behavior():
# Regression guard: authoring must not alter public col()/Expr query behavior.
expr = col("firstName") > lit(1)
assert isinstance(expr, Expr)
assert expr.to_sql() == "(`firstName` > 1)"
@@ -0,0 +1,181 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
from __future__ import annotations
import pyarrow as pa
from lancedb import udf
@udf(
inputs={"value": pa.int64()},
output=pa.int64(),
python="3.12",
packages=["pyarrow==24.0.0"],
output_nullable=True,
)
def double_nullable(value):
if value is None:
return None
return value * 2
def test_first_class_function_enterprise_lifecycle():
import json
import os
import uuid
from datetime import timedelta
import pytest
import lancedb
from lancedb.exceptions import FunctionError
from lancedb.expr import col
host = os.environ.get("LANCEDB_FCF_E2E_HOST")
if not host:
pytest.skip("LANCEDB_FCF_E2E_HOST is required for the live enterprise test")
database_uri = os.environ.get("LANCEDB_FCF_E2E_DB_URI", "db://fcf-e2e-local")
api_key = os.environ.get("LANCEDB_FCF_E2E_API_KEY", "fake")
run_suffix = uuid.uuid4().hex[:12]
table_name = f"fcf_e2e_{run_suffix}"
function_name = f"fcf_e2e.double_{run_suffix}"
job_timeout = timedelta(minutes=5)
query_timeout = timedelta(seconds=30)
def connect():
return lancedb.connect(
database_uri,
api_key=api_key,
host_override=host,
)
setup_db = connect()
setup_db.create_table(
table_name,
data=pa.Table.from_pylist(
[
{"row_id": 1, "value": 2},
{"row_id": 2, "value": 5},
{"row_id": 3, "value": None},
],
schema=pa.schema(
[
pa.field("row_id", pa.int64(), nullable=False),
pa.field("value", pa.int64(), nullable=True),
]
),
),
)
registration_job = setup_db.functions.register(function_name, double_nullable)
registration_job_id = registration_job.id
assert isinstance(registration_job_id, str) and registration_job_id
registered_function = registration_job.wait(timeout=job_timeout)
assert type(registered_function) is lancedb.Function
assert isinstance(registered_function.id, str) and registered_function.id
with pytest.raises(AttributeError):
registered_function.id = "mutated"
catalog_reader = connect()
function_by_name = catalog_reader.functions.get(function_name)
function_by_id = catalog_reader.functions.get_by_id(registered_function.id)
expected_signature = ((("value", pa.int64()),), pa.int64(), True)
expected_identity = (
registered_function.id,
*expected_signature,
)
for function in (registered_function, function_by_name, function_by_id):
assert type(function) is lancedb.Function
assert (
function.id,
function.parameters,
function.output_type,
function.output_nullable,
) == expected_identity
generated_column_table = catalog_reader.open_table(table_name)
generated_column_job = generated_column_table.add_generated_column(
"derived",
registered_function(value=col("value")),
)
generated_column_job_id = generated_column_job.id
assert isinstance(generated_column_job_id, str) and generated_column_job_id
assert generated_column_job.wait(timeout=job_timeout) is None
complete_reader = connect().open_table(table_name)
complete_status = complete_reader.generated_column_status("derived")
assert complete_status == "complete"
initial_rows = sorted(
complete_reader.search()
.select(["row_id", "value", "derived"])
.limit(3)
.to_list(timeout=query_timeout),
key=lambda row: row["row_id"],
)
assert initial_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 5, "derived": 10},
{"row_id": 3, "value": None, "derived": None},
]
update_result = complete_reader.update(
where="row_id = 2",
values={"value": 7},
)
assert update_result.rows_updated == 1
incomplete_reader = connect().open_table(table_name)
incomplete_status = incomplete_reader.generated_column_status("derived")
assert incomplete_status == "incomplete"
with pytest.raises(FunctionError) as raised:
(
incomplete_reader.search()
.select(["row_id", "derived"])
.limit(3)
.to_list(timeout=query_timeout)
)
assert raised.value.code == "generated_column_incomplete"
refresh_job = incomplete_reader.refresh_generated_column("derived")
refresh_job_id = refresh_job.id
assert isinstance(refresh_job_id, str) and refresh_job_id
assert refresh_job.wait(timeout=job_timeout) is None
refreshed_reader = connect().open_table(table_name)
refreshed_status = refreshed_reader.generated_column_status("derived")
assert refreshed_status == "complete"
final_rows = sorted(
refreshed_reader.search()
.select(["row_id", "value", "derived"])
.limit(3)
.to_list(timeout=query_timeout),
key=lambda row: row["row_id"],
)
assert final_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 7, "derived": 14},
{"row_id": 3, "value": None, "derived": None},
]
evidence = {
"run_suffix": run_suffix,
"database": database_uri.removeprefix("db://"),
"table": table_name,
"function": function_name,
"function_id": registered_function.id,
"job_ids": {
"register": registration_job_id,
"add_generated_column": generated_column_job_id,
"refresh_generated_column": refresh_job_id,
},
"status_transitions": [
complete_status,
incomplete_status,
refreshed_status,
],
"final_rows": final_rows,
}
print(json.dumps(evidence, sort_keys=True, separators=(",", ":")))
@@ -0,0 +1,595 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
from __future__ import annotations
import pyarrow as pa
from lancedb import udf
_RUNNING_DEADLINE_SECONDS = 30
@udf(
inputs={"value": pa.int64()},
output=pa.int64(),
python="3.12",
packages=["pyarrow==24.0.0"],
output_nullable=True,
)
def reliable_double(value):
if value is None:
return None
return value * 2
@udf(
inputs={"value": pa.int64()},
output=pa.int64(),
python="3.12",
packages=["pyarrow==24.0.0"],
output_nullable=True,
)
def terminate_worker_on_input(value):
if value is None:
return None
try:
if len(value) == 0:
return value
except TypeError:
pass
import os
os._exit(73)
@udf(
inputs={"value": pa.int64()},
output=pa.int64(),
python="3.12",
packages=["pyarrow==24.0.0"],
output_nullable=False,
)
def slow_triple(value):
import time
time.sleep(0.02)
return value * 3
def _require_live() -> str:
import os
import pytest
host = os.environ.get("LANCEDB_FCF_E2E_HOST")
if not host:
pytest.skip(
"LANCEDB_FCF_E2E_HOST is required for live enterprise reliability tests"
)
return host
def _job_timeout():
from datetime import timedelta
return timedelta(minutes=5)
def _query_timeout():
from datetime import timedelta
return timedelta(seconds=30)
def _connect():
import os
import lancedb
return lancedb.connect(
os.environ.get("LANCEDB_FCF_E2E_DB_URI", "db://fcf-e2e-local"),
api_key=os.environ.get("LANCEDB_FCF_E2E_API_KEY", "fake"),
host_override=_require_live(),
)
def _run_names(case: str) -> tuple[str, str]:
import uuid
suffix = uuid.uuid4().hex[:12]
return f"fcf_rel_{case}_{suffix}", f"fcf_rel.{case}_{suffix}"
def _read_rows(table, columns: list[str], row_count: int) -> list[dict]:
return sorted(
table.search()
.select(columns)
.limit(row_count)
.to_list(timeout=_query_timeout()),
key=lambda row: row["row_id"],
)
def _emit_evidence(case: str, evidence: dict) -> None:
import json
print(
json.dumps(
{"case": case, **evidence},
sort_keys=True,
separators=(",", ":"),
)
)
def test_enterprise_reliability_core_lifecycle():
import pytest
import lancedb
from lancedb.exceptions import FunctionError
from lancedb.expr import col
_require_live()
table_name, function_name = _run_names("lifecycle")
setup_db = _connect()
setup_db.create_table(
table_name,
data=pa.Table.from_pylist(
[
{"row_id": 1, "value": 2},
{"row_id": 2, "value": 5},
{"row_id": 3, "value": None},
],
schema=pa.schema(
[
pa.field("row_id", pa.int64(), nullable=False),
pa.field("value", pa.int64(), nullable=True),
]
),
),
)
registration_job = setup_db.functions.register(function_name, reliable_double)
registration_job_id = registration_job.id
assert isinstance(registration_job_id, str) and registration_job_id
registered = registration_job.wait(timeout=_job_timeout())
assert type(registered) is lancedb.Function
assert isinstance(registered.id, str) and registered.id
with pytest.raises(AttributeError):
registered.id = "mutated"
catalog_reader = _connect()
by_name = catalog_reader.functions.get(function_name)
by_id = catalog_reader.functions.get_by_id(registered.id)
expected_identity = (
registered.id,
(("value", pa.int64()),),
pa.int64(),
True,
)
for function in (registered, by_name, by_id):
assert type(function) is lancedb.Function
assert (
function.id,
function.parameters,
function.output_type,
function.output_nullable,
) == expected_identity
table = catalog_reader.open_table(table_name)
create_job = table.add_generated_column(
"derived",
registered(value=col("value")),
)
create_job_id = create_job.id
assert isinstance(create_job_id, str) and create_job_id
assert create_job.wait(timeout=_job_timeout()) is None
complete_reader = _connect().open_table(table_name)
complete_status = complete_reader.generated_column_status("derived")
assert complete_status == "complete"
initial_rows = _read_rows(
complete_reader,
["row_id", "value", "derived"],
3,
)
assert initial_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 5, "derived": 10},
{"row_id": 3, "value": None, "derived": None},
]
complete_reader.update(where="row_id = 2", values={"value": 7})
incomplete_reader = _connect().open_table(table_name)
changed_rows = _read_rows(incomplete_reader, ["row_id", "value"], 3)
assert changed_rows == [
{"row_id": 1, "value": 2},
{"row_id": 2, "value": 7},
{"row_id": 3, "value": None},
]
incomplete_status = incomplete_reader.generated_column_status("derived")
assert incomplete_status == "incomplete"
with pytest.raises(FunctionError) as raised:
(
incomplete_reader.search()
.select(["row_id", "derived"])
.limit(3)
.to_list(timeout=_query_timeout())
)
assert raised.value.code == "generated_column_incomplete"
refresh_job = incomplete_reader.refresh_generated_column("derived")
refresh_job_id = refresh_job.id
assert isinstance(refresh_job_id, str) and refresh_job_id
assert refresh_job.wait(timeout=_job_timeout()) is None
refreshed_reader = _connect().open_table(table_name)
refreshed_status = refreshed_reader.generated_column_status("derived")
assert refreshed_status == "complete"
final_rows = _read_rows(
refreshed_reader,
["row_id", "value", "derived"],
3,
)
assert final_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 7, "derived": 14},
{"row_id": 3, "value": None, "derived": None},
]
_emit_evidence(
"core_lifecycle",
{
"final_rows": final_rows,
"function_id": registered.id,
"job_ids": {
"create": create_job_id,
"refresh": refresh_job_id,
"register": registration_job_id,
},
"status": [
complete_status,
incomplete_status,
refreshed_status,
],
"table": table_name,
},
)
def test_enterprise_reliability_restart_retention():
import json
import os
import pytest
import lancedb
_require_live()
raw_evidence = os.environ.get("LANCEDB_FCF_E2E_RESTART_EVIDENCE")
if not raw_evidence:
pytest.skip(
"LANCEDB_FCF_E2E_RESTART_EVIDENCE is required for restart retention"
)
try:
evidence = json.loads(raw_evidence)
except json.JSONDecodeError as error:
pytest.fail(f"LANCEDB_FCF_E2E_RESTART_EVIDENCE must be valid JSON: {error.msg}")
assert isinstance(evidence, dict), (
"LANCEDB_FCF_E2E_RESTART_EVIDENCE must be a JSON object"
)
table_name = evidence.get("table")
function_id = evidence.get("function_id")
raw_job_ids = evidence.get("job_ids")
assert isinstance(table_name, str) and table_name, (
"restart evidence must contain a non-empty table"
)
assert isinstance(function_id, str) and function_id, (
"restart evidence must contain a non-empty function_id"
)
assert isinstance(raw_job_ids, dict), (
"restart evidence must contain a job_ids object"
)
job_ids = {}
for job_kind in ("register", "create", "refresh"):
job_id = raw_job_ids.get(job_kind)
assert isinstance(job_id, str) and job_id, (
f"restart evidence must contain a non-empty job_ids.{job_kind}"
)
job_ids[job_kind] = job_id
db = _connect()
function = db.functions.get_by_id(function_id)
expected_identity = (
function_id,
(("value", pa.int64()),),
pa.int64(),
True,
)
assert type(function) is lancedb.Function
assert (
function.id,
function.parameters,
function.output_type,
function.output_nullable,
) == expected_identity
jobs = {}
for job_kind in ("register", "create", "refresh"):
job = db.get_job(job_ids[job_kind])
assert job is not None
assert job.job_id == job_ids[job_kind]
assert job.state == "finished"
assert job.failure is None
jobs[job_kind] = job
registered_result = jobs["register"].result
assert type(registered_result) is lancedb.Function
assert (
registered_result.id,
registered_result.parameters,
registered_result.output_type,
registered_result.output_nullable,
) == expected_identity
assert jobs["create"].result is None
assert jobs["refresh"].result is None
table = db.open_table(table_name)
status = table.generated_column_status("derived")
assert status == "complete"
assert table.count_rows() == 3
final_rows = _read_rows(table, ["row_id", "value", "derived"], 3)
assert final_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 7, "derived": 14},
{"row_id": 3, "value": None, "derived": None},
]
_emit_evidence(
"restart_retention",
{
"final_rows": final_rows,
"function_id": function_id,
"generated_column_status": status,
"job_ids": job_ids,
"job_states": {
job_kind: jobs[job_kind].state
for job_kind in ("register", "create", "refresh")
},
"table": table_name,
},
)
def test_enterprise_reliability_failure_atomicity_and_worker_recovery():
import pytest
import lancedb
from lancedb.exceptions import JobFailedError
from lancedb.expr import col
_require_live()
table_name, failing_function_name = _run_names("worker_failure")
_, healthy_function_name = _run_names("worker_recovery")
row_count = 4
setup_db = _connect()
setup_db.create_table(
table_name,
data=pa.table(
{
"row_id": list(range(row_count)),
"value": [1, 2, 3, 4],
}
),
)
registration_job = setup_db.functions.register(
failing_function_name,
terminate_worker_on_input,
)
failing_function = registration_job.wait(timeout=_job_timeout())
assert type(failing_function) is lancedb.Function
table = setup_db.open_table(table_name)
failed_create_job = table.add_generated_column(
"must_not_publish",
failing_function(value=col("value")),
)
failed_job_id = failed_create_job.id
assert isinstance(failed_job_id, str) and failed_job_id
with pytest.raises(JobFailedError) as raised:
failed_create_job.wait(timeout=_job_timeout())
assert raised.value.error_code == "udf_execution_failure"
first_description = _connect().get_job(failed_job_id)
second_description = _connect().get_job(failed_job_id)
for description in (first_description, second_description):
assert description is not None
assert description.job_id == failed_job_id
assert description.state == "failed"
assert description.failure is not None
assert description.failure.error_code == "udf_execution_failure"
atomic_reader = _connect().open_table(table_name)
assert "must_not_publish" not in atomic_reader.schema.names
assert _read_rows(atomic_reader, ["row_id", "value"], row_count) == [
{"row_id": 0, "value": 1},
{"row_id": 1, "value": 2},
{"row_id": 2, "value": 3},
{"row_id": 3, "value": 4},
]
healthy_registration_job = setup_db.functions.register(
healthy_function_name,
reliable_double,
)
healthy_function = healthy_registration_job.wait(timeout=_job_timeout())
assert type(healthy_function) is lancedb.Function
recovery_job = atomic_reader.add_generated_column(
"recovered",
healthy_function(value=col("value")),
)
recovery_job_id = recovery_job.id
assert isinstance(recovery_job_id, str) and recovery_job_id
assert recovery_job.wait(timeout=_job_timeout()) is None
recovered_reader = _connect().open_table(table_name)
assert "must_not_publish" not in recovered_reader.schema.names
assert recovered_reader.generated_column_status("recovered") == "complete"
recovered_rows = _read_rows(
recovered_reader,
["row_id", "value", "recovered"],
row_count,
)
assert recovered_rows == [
{"row_id": 0, "value": 1, "recovered": 2},
{"row_id": 1, "value": 2, "recovered": 4},
{"row_id": 2, "value": 3, "recovered": 6},
{"row_id": 3, "value": 4, "recovered": 8},
]
_emit_evidence(
"failure_atomicity_and_worker_recovery",
{
"failure_code": first_description.failure.error_code,
"failed_job_id": failed_job_id,
"recovered_rows": recovered_rows,
"recovery_job_id": recovery_job_id,
"table": table_name,
},
)
def test_enterprise_reliability_concurrent_refresh_fencing():
import time
import pytest
import lancedb
from lancedb.exceptions import FunctionError, JobFailedError
from lancedb.expr import col
_require_live()
table_name, function_name = _run_names("refresh_fencing")
row_count = 1024
setup_db = _connect()
setup_db.create_table(
table_name,
data=pa.table(
{
"row_id": list(range(row_count)),
"value": list(range(row_count)),
}
),
)
registration_job = setup_db.functions.register(function_name, slow_triple)
function = registration_job.wait(timeout=_job_timeout())
assert type(function) is lancedb.Function
table = setup_db.open_table(table_name)
create_job = table.add_generated_column(
"derived",
function(value=col("value")),
)
assert create_job.wait(timeout=_job_timeout()) is None
initial_reader = _connect().open_table(table_name)
assert initial_reader.generated_column_status("derived") == "complete"
initial_reader.update(where="row_id = 0", values={"value": 10_000})
incomplete_reader = _connect().open_table(table_name)
assert incomplete_reader.generated_column_status("derived") == "incomplete"
refresh_job = incomplete_reader.refresh_generated_column("derived")
refresh_job_id = refresh_job.id
assert isinstance(refresh_job_id, str) and refresh_job_id
deadline = time.monotonic() + _RUNNING_DEADLINE_SECONDS
observed_states = []
running_observations = 0
while running_observations < 2:
state = refresh_job.status()
if not observed_states or observed_states[-1] != state:
observed_states.append(state)
if state == "running":
running_observations += 1
else:
running_observations = 0
assert state not in {"finished", "failed", "cancelled"}
assert time.monotonic() < deadline
if running_observations < 2:
time.sleep(0.05)
concurrent_writer = _connect().open_table(table_name)
concurrent_writer.update(where="row_id = 1", values={"value": 20_000})
with pytest.raises(JobFailedError) as raised:
refresh_job.wait(timeout=_job_timeout())
assert raised.value.error_code == "stale_or_conflicting_input"
stale_job = _connect().get_job(refresh_job_id)
assert stale_job is not None
assert stale_job.job_id == refresh_job_id
assert stale_job.state == "failed"
assert stale_job.failure is not None
assert stale_job.failure.error_code == raised.value.error_code
if observed_states[-1] != stale_job.state:
observed_states.append(stale_job.state)
stale_reader = _connect().open_table(table_name)
stale_rows = _read_rows(stale_reader, ["row_id", "value"], row_count)
assert len(stale_rows) == row_count
for row_id, row in enumerate(stale_rows):
expected_value = 10_000 if row_id == 0 else 20_000 if row_id == 1 else row_id
assert (row["row_id"], row["value"]) == (row_id, expected_value)
assert stale_reader.generated_column_status("derived") == "incomplete"
with pytest.raises(FunctionError) as incomplete:
(
stale_reader.search()
.select(["row_id", "derived"])
.limit(row_count)
.to_list(timeout=_query_timeout())
)
assert incomplete.value.code == "generated_column_incomplete"
resubmitted_job = stale_reader.refresh_generated_column("derived")
resubmitted_job_id = resubmitted_job.id
assert isinstance(resubmitted_job_id, str) and resubmitted_job_id
assert resubmitted_job.wait(timeout=_job_timeout()) is None
final_reader = _connect().open_table(table_name)
final_status = final_reader.generated_column_status("derived")
assert final_status == "complete"
final_rows = _read_rows(
final_reader,
["row_id", "value", "derived"],
row_count,
)
assert len(final_rows) == row_count
final_checksum = 0
for row_id, row in enumerate(final_rows):
expected_value = 10_000 if row_id == 0 else 20_000 if row_id == 1 else row_id
assert (row["row_id"], row["value"], row["derived"]) == (
row_id,
expected_value,
expected_value * 3,
)
final_checksum += row["derived"]
_emit_evidence(
"concurrent_refresh_fencing",
{
"failure_code": stale_job.failure.error_code,
"final_checksum": final_checksum,
"final_status": final_status,
"observed_states": observed_states,
"resubmitted_job_id": resubmitted_job_id,
"row_count": row_count,
"stale_job_id": refresh_job_id,
"table": table_name,
},
)
@@ -0,0 +1,268 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract: Python projection of JobFailure.error_code / JobFailedError.error_code.
Public Function failures expose eight stable string categories. Asynchronous
errors remain the unified JobFailedError and JobFailureInfo. Python must
project the optional exact error_code string already supplied structurally by
Rust: preserve a known code, preserve an unknown nonempty future code
byte-for-byte, and return None for legacy failure payloads without error_code.
Never infer or override a code from message, phase, retryable, HTTP status,
job type, or state.
"""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from datetime import timedelta
from typing import Any, Callable, Optional
import pytest
import lancedb
from lancedb.exceptions import JobFailedError
_DESCRIBE_PATH = "/v1/jobs/describe"
_KNOWN_CODE = "name_or_function_not_found"
_CONFLICTING_STABLE_IN_MESSAGE = "definition_validation_failure"
_UNKNOWN_CODE = "enterprise_future_category_xyz"
_WAIT_KNOWN_CODE = "unsupported_runtime_or_capability"
_WAIT_CONFLICTING_IN_MESSAGE = "revoked_function"
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _failed_describe_body(
*,
job_id: str,
error_code: Optional[str] = None,
include_error_code: bool = True,
phase: str = "execute",
message: str = "worker died",
retryable: bool = False,
job_type: str = "create_index",
) -> dict[str, Any]:
failure: dict[str, Any] = {
"phase": phase,
"message": message,
"retryable": retryable,
}
if include_error_code:
failure["error_code"] = error_code
return {
"job_id": job_id,
"job_type": job_type,
"job_state": "FAILED",
"creation_ms": 1000,
"spec": {},
"failure": failure,
}
def _describe_handler(bodies_by_job_id: dict[str, dict[str, Any]]):
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _DESCRIBE_PATH
payload = json.loads(_read_body(request).decode("utf-8") or "{}")
job_id = payload["job_id"]
body = bodies_by_job_id.get(job_id)
if body is None:
request.send_response(404)
request.end_headers()
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
return handler
def test_get_job_failure_error_code_known_not_inferred_from_message():
"""Structural error_code wins; conflicting message text must not override."""
body = _failed_describe_body(
job_id="job-known",
error_code=_KNOWN_CODE,
phase="validate",
message=f"looks like {_CONFLICTING_STABLE_IN_MESSAGE}",
retryable=False,
)
with _mock_remote_db(_describe_handler({"job-known": body})) as db:
description = db.get_job("job-known")
assert description is not None
failure = description.failure
assert failure is not None
assert failure.error_code == _KNOWN_CODE
assert failure.error_code != _CONFLICTING_STABLE_IN_MESSAGE
assert failure.phase == "validate"
assert failure.message == f"looks like {_CONFLICTING_STABLE_IN_MESSAGE}"
assert failure.retryable is False
def test_get_job_failure_error_code_unknown_preserved_byte_for_byte():
body = _failed_describe_body(
job_id="job-unknown",
error_code=_UNKNOWN_CODE,
phase="execute",
message=f"new category mentioning {_KNOWN_CODE}",
retryable=True,
)
with _mock_remote_db(_describe_handler({"job-unknown": body})) as db:
failure = db.get_job("job-unknown").failure
assert failure.error_code == _UNKNOWN_CODE
assert failure.error_code != _KNOWN_CODE
def test_get_job_failure_error_code_absent_is_none():
"""Legacy describe payloads without error_code must not invent a category."""
body = _failed_describe_body(
job_id="job-legacy",
include_error_code=False,
phase="execute",
message=f"{_KNOWN_CODE} in logs",
retryable=True,
)
with _mock_remote_db(_describe_handler({"job-legacy": body})) as db:
failure = db.get_job("job-legacy").failure
assert failure.error_code is None
assert failure.phase == "execute"
assert failure.retryable is True
def test_sync_job_wait_job_failed_error_code_known_not_inferred():
body = _failed_describe_body(
job_id="job-wait-known",
error_code=_WAIT_KNOWN_CODE,
phase="dispatch",
message=f"{_WAIT_CONFLICTING_IN_MESSAGE} in transport logs",
retryable=False,
)
with _mock_remote_db(_describe_handler({"job-wait-known": body})) as db:
with pytest.raises(JobFailedError) as exc_info:
db.job("job-wait-known").wait(timeout=timedelta(seconds=5))
err = exc_info.value
assert isinstance(err, JobFailedError)
assert err.error_code == _WAIT_KNOWN_CODE
assert err.error_code != _WAIT_CONFLICTING_IN_MESSAGE
def test_sync_job_wait_job_failed_error_code_absent_is_none():
body = _failed_describe_body(
job_id="job-wait-legacy",
include_error_code=False,
phase="execute",
message=f"{_WAIT_KNOWN_CODE} mentioned only in message",
retryable=True,
)
with _mock_remote_db(_describe_handler({"job-wait-legacy": body})) as db:
with pytest.raises(JobFailedError) as exc_info:
db.job("job-wait-legacy").wait(timeout=timedelta(seconds=5))
assert exc_info.value.error_code is None
@pytest.mark.asyncio
async def test_async_job_wait_job_failed_error_code_unknown_preserved():
body = _failed_describe_body(
job_id="job-wait-unknown",
error_code=_UNKNOWN_CODE,
phase="execute",
message=f"future code with {_WAIT_KNOWN_CODE} in text",
retryable=False,
)
async with _mock_remote_db_async(
_describe_handler({"job-wait-unknown": body})
) as db:
with pytest.raises(JobFailedError) as exc_info:
await db.job("job-wait-unknown").wait(timeout=timedelta(seconds=5))
err = exc_info.value
assert err.error_code == _UNKNOWN_CODE
assert err.error_code != _WAIT_KNOWN_CODE
def test_job_failed_error_legacy_message_construction_error_code_is_none():
err = JobFailedError("legacy construction with only a message")
assert err.error_code is None
def test_job_failed_error_error_code_is_read_only():
err = JobFailedError("message")
with pytest.raises(AttributeError):
err.error_code = _KNOWN_CODE
@@ -0,0 +1,634 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python first-class Function catalog lookup."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable
import pyarrow as pa
import pytest
import lancedb
from lancedb import _lancedb as _native
from lancedb.remote.errors import HttpError
def _function_error_cls(*, required: bool = True):
"""Resolve FunctionError from the live module (records RED when absent)."""
from lancedb import exceptions as exc_mod
cls = getattr(exc_mod, "FunctionError", None)
if cls is None:
if required:
pytest.fail("lancedb.exceptions.FunctionError is missing")
return type("MissingFunctionError", (), {})
return cls
_LOOKUP_PATH = "/v1/functions/lookup"
_LOOKUP_CATALOG_NAME = "text.normalize.lookup-name"
_LOOKUP_FUNCTION_ID = "fn.exact.lookup-handle"
_LOOKUP_SERVER_MESSAGE_MARKER = (
"SERVER_LOOKUP_DIAGNOSTIC_MARKER name=text.normalize.lookup-name "
"id=fn.exact.lookup-handle"
)
_SENSITIVE_BODY_MARKER = "SENSITIVE_LOOKUP_BODY_MARKER"
_UNKNOWN_CODE = "enterprise_future_lookup_category_xyz"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as job-result
# tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust FileWriter.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_LOOKUP_KEYWORDS = (
"idempotency_key",
"retry_key",
"user_version",
"replace",
"expected_current_function_id",
"list",
"alias",
"lineage",
"FunctionVersion",
)
def _sample_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _LOOKUP_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _lookup_success_body(
*,
function: dict[str, Any] | None = None,
extra_outer: dict[str, Any] | None = None,
) -> bytes:
body: dict[str, Any] = {"function": function or _sample_function_wire()}
if extra_outer:
body.update(extra_outer)
return json.dumps(body).encode("utf-8")
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _exception_chain_text(exc: BaseException) -> str:
parts: list[str] = []
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(str(current))
parts.append(repr(current))
current = current.__cause__
return "\n".join(parts)
def _assert_payload_free(exc: BaseException) -> None:
text = _exception_chain_text(exc)
assert _LOOKUP_SERVER_MESSAGE_MARKER not in text
assert _LOOKUP_CATALOG_NAME not in text
assert _LOOKUP_FUNCTION_ID not in text
assert _SENSITIVE_BODY_MARKER not in text
def _assert_exact_lookup_function(function: object) -> None:
assert type(function) is lancedb.Function
assert function.id == _LOOKUP_FUNCTION_ID
assert not hasattr(function, "name")
assert function.parameters == (
("text", pa.string()),
("limit", pa.int32()),
)
assert function.output_type == pa.string()
assert function.output_nullable is True
assert _LOOKUP_CATALOG_NAME not in repr(function)
assert _LOOKUP_CATALOG_NAME not in str(function)
def _assert_name_request(raw: bytes, body: dict[str, Any]) -> None:
assert raw
assert body == {"name": _LOOKUP_CATALOG_NAME}
assert "function_id" not in body
def _assert_id_request(raw: bytes, body: dict[str, Any]) -> None:
assert raw
assert body == {"function_id": _LOOKUP_FUNCTION_ID}
assert "name" not in body
def _assert_native_lookup_methods_present() -> None:
assert hasattr(_native.Connection, "_lookup_function_by_name")
assert hasattr(_native.Connection, "_lookup_function_by_id")
assert callable(getattr(_native.Connection, "_lookup_function_by_name"))
assert callable(getattr(_native.Connection, "_lookup_function_by_id"))
def test_native_connection_exposes_private_lookup_methods():
_assert_native_lookup_methods_present()
def test_sync_remote_get_by_name_exact_request_and_function_shape():
_assert_native_lookup_methods_present()
seen: dict[str, Any] = {}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _LOOKUP_PATH
raw = _read_body(request)
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
with _mock_remote_db(handler) as db:
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
assert not hasattr(db, "get_function")
function = db.functions.get(_LOOKUP_CATALOG_NAME)
_assert_name_request(seen["raw"], seen["body"])
_assert_exact_lookup_function(function)
def test_sync_remote_get_by_id_exact_request_and_function_shape():
_assert_native_lookup_methods_present()
seen: dict[str, Any] = {}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _LOOKUP_PATH
raw = _read_body(request)
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
with _mock_remote_db(handler) as db:
function = db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
_assert_id_request(seen["raw"], seen["body"])
_assert_exact_lookup_function(function)
@pytest.mark.asyncio
async def test_async_remote_get_by_name_and_id():
_assert_native_lookup_methods_present()
name_seen: dict[str, Any] = {}
id_seen: dict[str, Any] = {}
stage = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _LOOKUP_PATH
raw = _read_body(request)
body = json.loads(raw.decode("utf-8"))
stage["n"] += 1
if stage["n"] == 1:
name_seen["raw"] = raw
name_seen["body"] = body
else:
id_seen["raw"] = raw
id_seen["body"] = body
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
async with _mock_remote_db_async(handler) as db:
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
by_name = await db.functions.get(_LOOKUP_CATALOG_NAME)
by_id = await db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
_assert_name_request(name_seen["raw"], name_seen["body"])
_assert_id_request(id_seen["raw"], id_seen["body"])
_assert_exact_lookup_function(by_name)
_assert_exact_lookup_function(by_id)
def test_sync_remote_get_accepts_additive_outer_success_fields():
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
_lookup_success_body(
extra_outer={
"server_extra": {"ok": True},
"request_echo_name": _LOOKUP_CATALOG_NAME,
}
)
)
with _mock_remote_db(handler) as db:
function = db.functions.get(_LOOKUP_CATALOG_NAME)
_assert_exact_lookup_function(function)
def test_empty_name_and_id_reject_before_transport():
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
_read_body(request)
request.send_response(500)
request.end_headers()
request.wfile.write(b"should not be reached")
with _mock_remote_db(handler) as db:
with pytest.raises(ValueError):
db.functions.get("")
with pytest.raises(ValueError):
db.functions.get_by_id("")
assert received["n"] == 0
def test_local_sync_lookup_not_implemented_without_table_mutation(tmp_path):
_assert_native_lookup_methods_present()
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
with pytest.raises(NotImplementedError):
db.functions.get(_LOOKUP_CATALOG_NAME)
with pytest.raises(NotImplementedError):
db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
assert db.list_tables().tables == before
@pytest.mark.asyncio
async def test_local_async_lookup_not_implemented_without_table_mutation(tmp_path):
_assert_native_lookup_methods_present()
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
with pytest.raises(NotImplementedError):
await db.functions.get(_LOOKUP_CATALOG_NAME)
with pytest.raises(NotImplementedError):
await db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
assert (await db.list_tables()).tables == before
def test_explicit_known_code_is_function_error_with_exact_code():
body = {
"error_code": "name_or_function_not_found",
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
"looks_like": "definition_validation_failure",
}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
request.send_response(404)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(handler) as db:
with pytest.raises(function_error) as exc_info:
db.functions.get(_LOOKUP_CATALOG_NAME)
err = exc_info.value
assert isinstance(err, function_error)
assert err.code == "name_or_function_not_found"
assert err.code != "definition_validation_failure"
_assert_payload_free(err)
def test_explicit_unknown_code_preserved_despite_status_and_message():
body = {
"error_code": _UNKNOWN_CODE,
"message": (
f"{_LOOKUP_SERVER_MESSAGE_MARKER} revoked_function "
"name_or_function_not_found"
),
}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
raw = _read_body(request)
assert json.loads(raw.decode("utf-8")) == {"function_id": _LOOKUP_FUNCTION_ID}
request.send_response(409)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(handler) as db:
with pytest.raises(function_error) as exc_info:
db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
err = exc_info.value
assert err.code == _UNKNOWN_CODE
assert err.code != "revoked_function"
assert err.code != "name_or_function_not_found"
_assert_payload_free(err)
@pytest.mark.parametrize(
"label,status,response_body",
[
(
"missing_code_404",
404,
{
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
},
),
(
"empty_code",
400,
{
"error_code": "",
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
},
),
(
"wrong_type_code",
400,
{
"error_code": 123,
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
},
),
(
"null_code",
404,
{
"error_code": None,
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
},
),
(
"non_json",
404,
f"not-json {_LOOKUP_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
),
],
)
def test_invalid_or_missing_error_code_is_payload_free_http(
label: str, status: int, response_body: object
):
del label # parametrize label for failure diagnosis only
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
request.send_response(status)
request.send_header("Content-Type", "application/json")
request.end_headers()
if isinstance(response_body, str):
request.wfile.write(response_body.encode("utf-8"))
else:
request.wfile.write(json.dumps(response_body).encode("utf-8"))
with _mock_remote_db(handler) as db:
with pytest.raises(HttpError) as exc_info:
db.functions.get(_LOOKUP_CATALOG_NAME)
err = exc_info.value
assert isinstance(err, HttpError)
assert not isinstance(err, _function_error_cls(required=False))
_assert_payload_free(err)
@pytest.mark.parametrize(
"label,response_body",
[
(
"missing_function",
{
"server_extra": True,
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
},
),
(
"null_function",
{
"function": None,
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
},
),
(
"wrong_type_function",
{
"function": "not-an-object",
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
},
),
(
"invalid_function_shape",
{
"function": {
"format_version": 1,
"id": _LOOKUP_FUNCTION_ID,
# missing signature
_SENSITIVE_BODY_MARKER: True,
}
},
),
],
)
def test_malformed_success_is_payload_free_http(label: str, response_body: dict):
del label
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(response_body).encode("utf-8"))
with _mock_remote_db(handler) as db:
with pytest.raises(HttpError) as exc_info:
db.functions.get(_LOOKUP_CATALOG_NAME)
err = exc_info.value
assert isinstance(err, HttpError)
assert not isinstance(err, _function_error_cls(required=False))
_assert_payload_free(err)
def test_function_error_surface_omits_server_marker_name_and_id():
body = {
"error_code": "name_or_function_not_found",
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
"function_id": _LOOKUP_FUNCTION_ID,
"name": _LOOKUP_CATALOG_NAME,
_SENSITIVE_BODY_MARKER: True,
}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
_read_body(request)
request.send_response(404)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(handler) as db:
with pytest.raises(function_error) as exc_info:
db.functions.get(_LOOKUP_CATALOG_NAME)
err = exc_info.value
_assert_payload_free(err)
assert getattr(err, "code", None) == "name_or_function_not_found"
def test_no_direct_db_lookup_methods_and_no_deleted_keywords():
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
_read_body(request)
request.send_response(500)
request.end_headers()
request.wfile.write(b"should not be reached")
with _mock_remote_db(handler) as db:
assert not hasattr(db, "lookup_function")
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
assert not hasattr(db, "get_function")
assert not hasattr(db.functions, "get_by_name")
assert not hasattr(db.functions, "list")
for keyword in _DELETED_LOOKUP_KEYWORDS:
with pytest.raises(TypeError):
db.functions.get(_LOOKUP_CATALOG_NAME, **{keyword: True})
with pytest.raises(TypeError):
db.functions.get_by_id(_LOOKUP_FUNCTION_ID, **{keyword: True})
assert received["n"] == 0
def test_function_error_is_not_top_level_export():
assert not hasattr(lancedb, "FunctionError")
function_error = _function_error_cls()
assert issubclass(function_error, RuntimeError)
@@ -0,0 +1,398 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for Python first-class Function registration."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from datetime import timedelta
from typing import Any, Callable
from unittest import mock
import pyarrow as pa
import pytest
import lancedb
import lancedb._udf as _udf_mod
import lancedb.job
from lancedb import FunctionCapability, udf
from lancedb.remote.errors import HttpError
_SOURCE_MARKER = "registration-source-marker-unique-xyz"
_SECRET_REFERENCE = "secret://team/registration-redact-token-xyz"
_SECRET_ENV = "REGISTER_API_TOKEN"
_NETWORK_ORIGIN = "https://api.registration-example.com"
_FUNCTION_NAME = "text.normalize"
_FUNCTION_ID_RETRY = "fn.register-retry-1"
_JOB_ID_RETRY = "job-register-retry-1"
_JOB_ID_ASYNC = "job-register-async-1"
_REGISTER_PATH = "/v1/functions/register"
_DESCRIBE_PATH = "/v1/jobs/describe"
_DELETED_REGISTER_KEYWORDS = (
"idempotency_key",
"retry_key",
"user_version",
"deterministic",
"null_policy",
"replace",
"expected_current_function_id",
)
_SPEC_KEYS = {
"format_version",
"name",
"definition",
"expected_current_function_id",
}
@udf(
inputs={"text": pa.string(), "limit": pa.int32()},
output=pa.string(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=True,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
],
)
def packable_register_normalize(text, limit):
"""registration-source-marker-unique-xyz."""
return text[:limit]
def _definition_json(fn: object) -> dict[str, Any]:
payload = _udf_mod._build_function_definition(fn)._to_json()
if isinstance(payload, bytes):
return json.loads(payload.decode("utf-8"))
assert isinstance(payload, str)
return json.loads(payload)
def _expected_register_spec(name: str, fn: object) -> dict[str, Any]:
return {
"format_version": 1,
"name": name,
"definition": _definition_json(fn),
"expected_current_function_id": None,
}
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _exception_chain_text(exc: BaseException) -> str:
parts: list[str] = []
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(str(current))
parts.append(repr(current))
current = current.__cause__
return "\n".join(parts)
def _assert_markers_absent_from_exception(exc: BaseException) -> None:
text = _exception_chain_text(exc)
assert _SOURCE_MARKER not in text
assert _SECRET_REFERENCE not in text
def _assert_exact_register_spec(body: dict[str, Any], expected: dict[str, Any]) -> None:
assert set(body) == _SPEC_KEYS
assert body == expected
assert body["format_version"] == 1
assert body["expected_current_function_id"] is None
assert _SOURCE_MARKER in json.dumps(body["definition"])
assert any(
capability.get("reference") == _SECRET_REFERENCE
for capability in body["definition"]["capabilities"]
)
def test_sync_remote_register_retries_exact_wire_and_returns_job():
expected_spec = _expected_register_spec(_FUNCTION_NAME, packable_register_normalize)
attempts: list[dict[str, Any]] = []
describe_calls: list[dict[str, Any]] = []
function_result_wire = {
"kind": "function",
"format_version": 1,
"function": {
"format_version": 1,
"id": _FUNCTION_ID_RETRY,
"signature": expected_spec["definition"]["signature"],
},
}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _REGISTER_PATH:
request_id = request.headers.get("x-request-id")
attempts.append(
{
"request_id": request_id,
"raw": raw,
"body": json.loads(raw.decode("utf-8")),
}
)
if len(attempts) == 1:
request.send_response(500)
request.end_headers()
request.wfile.write(b"transient register failure")
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_RETRY}).encode("utf-8"))
return
assert request.path == _DESCRIBE_PATH
body = json.loads(raw.decode("utf-8"))
assert body["job_id"] == _JOB_ID_RETRY
describe_calls.append(body)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"job_id": _JOB_ID_RETRY,
"job_state": "DONE",
"job_type": "register_function",
"creation_ms": 1,
"spec": {},
"result": function_result_wire,
}
).encode("utf-8")
)
package_calls = {"n": 0}
original_package = _udf_mod._package_udf
def counting_package(fn: object):
package_calls["n"] += 1
return original_package(fn)
with _mock_remote_db(handler) as db:
assert not hasattr(db, "register_function")
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
job = db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_RETRY
waited = job.wait(timeout=timedelta(seconds=5))
assert package_calls["n"] == 1
assert len(attempts) == 2
first, second = attempts
assert isinstance(first["request_id"], str) and first["request_id"]
assert first["request_id"] == second["request_id"]
assert first["raw"] == second["raw"]
assert first["raw"]
_assert_exact_register_spec(first["body"], expected_spec)
_assert_exact_register_spec(second["body"], expected_spec)
assert len(describe_calls) == 1
assert describe_calls[0]["job_id"] == _JOB_ID_RETRY
assert type(waited) is lancedb.Function
assert waited.id == _FUNCTION_ID_RETRY
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
assert waited.output_type == pa.string()
assert waited.output_nullable is True
@pytest.mark.asyncio
async def test_async_remote_register_returns_async_job_with_exact_spec():
expected_spec = _expected_register_spec(_FUNCTION_NAME, packable_register_normalize)
seen: dict[str, Any] = {}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _REGISTER_PATH
raw = _read_body(request)
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
async with _mock_remote_db_async(handler) as db:
assert not hasattr(db, "register_function")
job = await db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert seen.get("raw")
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
_assert_exact_register_spec(seen["body"], expected_spec)
assert type(job) is lancedb.job.AsyncJob
assert job.id == _JOB_ID_ASYNC
def test_sync_remote_register_http_error_omits_source_and_secret_markers():
echoed = f"register failed with {_SOURCE_MARKER} and {_SECRET_REFERENCE}"
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
assert request.path == _REGISTER_PATH
_read_body(request)
request.send_response(400)
request.end_headers()
request.wfile.write(echoed.encode("utf-8"))
with _mock_remote_db(handler) as db:
with pytest.raises(HttpError) as exc_info:
db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert received["n"] == 1
err = exc_info.value
assert isinstance(err, HttpError)
assert err.status_code == 400
_assert_markers_absent_from_exception(err)
def test_empty_name_rejects_before_http():
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
_read_body(request)
request.send_response(500)
request.end_headers()
request.wfile.write(b"should not be reached")
with _mock_remote_db(handler) as db:
with pytest.raises(ValueError):
db.functions.register("", packable_register_normalize)
assert received["n"] == 0
def test_local_sync_register_not_implemented_without_table_mutation(tmp_path):
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "register_function")
with pytest.raises(NotImplementedError):
db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert db.list_tables().tables == before
@pytest.mark.asyncio
async def test_local_async_register_not_implemented_without_table_mutation(tmp_path):
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "register_function")
with pytest.raises(NotImplementedError):
await db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert (await db.list_tables()).tables == before
@pytest.mark.parametrize("keyword", _DELETED_REGISTER_KEYWORDS)
def test_register_rejects_deleted_overdesign_keywords_before_submission(keyword):
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
_read_body(request)
request.send_response(500)
request.end_headers()
request.wfile.write(b"should not be reached")
with _mock_remote_db(handler) as db:
with pytest.raises(TypeError):
db.functions.register(
_FUNCTION_NAME,
packable_register_normalize,
**{keyword: True},
)
assert received["n"] == 0
@@ -0,0 +1,719 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python conditional first-class Function name removal."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable
import pyarrow as pa
import pytest
import lancedb
from lancedb import _lancedb as _native
from lancedb.remote.errors import HttpError
_REMOVE_PATH = "/v1/functions/remove"
_LOOKUP_PATH = "/v1/functions/lookup"
_REMOVE_CATALOG_NAME = "text.normalize.remove-name"
_REMOVE_FUNCTION_ID = "fn.exact.remove-handle"
_REMOVE_SERVER_MESSAGE_MARKER = (
"SERVER_REMOVE_DIAGNOSTIC_MARKER name=text.normalize.remove-name "
"id=fn.exact.remove-handle"
)
_SENSITIVE_BODY_MARKER = "SENSITIVE_REMOVE_BODY_MARKER"
_CONFLICTING_MESSAGE_CODE = "revoked_function"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
# replace tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_REMOVE_KEYWORDS = (
"expected_current_function_id",
"function_id",
"idempotency_key",
"retry_key",
"user_version",
"version",
"force",
"if_exists",
"revoke",
"delete",
)
def _function_error_cls(*, required: bool = True):
"""Resolve FunctionError from the live module (records RED when absent)."""
from lancedb import exceptions as exc_mod
cls = getattr(exc_mod, "FunctionError", None)
if cls is None:
if required:
pytest.fail("lancedb.exceptions.FunctionError is missing")
return type("MissingFunctionError", (), {})
return cls
def _sample_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _REMOVE_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _lookup_success_body() -> bytes:
return json.dumps({"function": _sample_function_wire()}).encode("utf-8")
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
def _close_db(db: Any) -> None:
with contextlib.suppress(Exception):
inner = getattr(db, "_conn", None)
if inner is not None:
inner.close()
return
close = getattr(db, "close", None)
if callable(close):
close()
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
db = None
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
if db is not None:
_close_db(db)
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
db = None
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
if db is not None:
_close_db(db)
server.shutdown()
thread.join()
def _exception_chain_text(exc: BaseException) -> str:
parts: list[str] = []
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(str(current))
parts.append(repr(current))
current = current.__cause__
return "\n".join(parts)
def _assert_payload_free(exc: BaseException) -> None:
text = _exception_chain_text(exc)
assert _REMOVE_SERVER_MESSAGE_MARKER not in text
assert _REMOVE_CATALOG_NAME not in text
assert _REMOVE_FUNCTION_ID not in text
assert _SENSITIVE_BODY_MARKER not in text
def _assert_exact_remove_function(function: object) -> None:
assert type(function) is lancedb.Function
assert function.id == _REMOVE_FUNCTION_ID
assert not hasattr(function, "name")
assert function.parameters == (
("text", pa.string()),
("limit", pa.int32()),
)
assert function.output_type == pa.string()
assert function.output_nullable is True
assert _REMOVE_CATALOG_NAME not in repr(function)
assert _REMOVE_CATALOG_NAME not in str(function)
def _assert_exact_remove_request(
request: http.server.BaseHTTPRequestHandler,
raw: bytes,
body: dict[str, Any],
*,
expected_id: str,
) -> None:
assert request.command == "POST"
assert request.path == _REMOVE_PATH
assert "?" not in request.path
assert raw
assert body == {
"name": _REMOVE_CATALOG_NAME,
"expected_current_function_id": expected_id,
}
assert set(body) == {"name", "expected_current_function_id"}
assert "format_version" not in body
assert "function_id" not in body
assert "function" not in body
assert "signature" not in body
assert "job_id" not in body
assert "idempotency_key" not in body
assert "user_version" not in body
assert "force" not in body
assert "if_exists" not in body
request_id = request.headers.get("x-request-id")
assert isinstance(request_id, str) and request_id
def _assert_native_remove_method_present() -> None:
assert hasattr(_native.Connection, "_remove_function_name")
assert callable(getattr(_native.Connection, "_remove_function_name"))
def _lookup_success_handler(
counters: dict[str, int],
*,
after_lookup: (
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
) = None,
):
"""Serve exact name lookup; optionally continue for remove."""
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
counters["lookup"] = counters.get("lookup", 0) + 1
body = json.loads(raw.decode("utf-8"))
assert body == {"name": _REMOVE_CATALOG_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
if after_lookup is not None:
after_lookup(request, raw)
return
counters["remove"] = counters.get("remove", 0) + 1
request.send_response(500)
request.end_headers()
request.wfile.write(b"unexpected remove")
return handler
def _observe_current(db) -> lancedb.Function:
current = db.functions.get(_REMOVE_CATALOG_NAME)
_assert_exact_remove_function(current)
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
return current
async def _observe_current_async(db) -> lancedb.Function:
current = await db.functions.get(_REMOVE_CATALOG_NAME)
_assert_exact_remove_function(current)
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
return current
def test_native_connection_exposes_private_remove_function_name():
_assert_native_remove_method_present()
def test_sync_remote_remove_exact_body_path_request_id_returns_none():
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
remove_attempts: list[dict[str, Any]] = []
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REMOVE_PATH
counters["remove"] += 1
body = json.loads(payload.decode("utf-8"))
remove_attempts.append(
{
"request": request,
"raw": payload,
"body": body,
"request_id": request.headers.get("x-request-id"),
}
)
# Illegal body on 204 must be ignored; success is status-driven only.
request.send_response(204)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
_SENSITIVE_BODY_MARKER: True,
"message": _REMOVE_SERVER_MESSAGE_MARKER,
}
).encode("utf-8")
)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
result = db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert result is None
assert counters["lookup"] == 1
assert counters["remove"] == 1
assert len(remove_attempts) == 1
attempt = remove_attempts[0]
_assert_exact_remove_request(
attempt["request"],
attempt["raw"],
attempt["body"],
expected_id=current.id,
)
assert attempt["body"]["expected_current_function_id"] == current.id
@pytest.mark.asyncio
async def test_async_remote_remove_exact_body_returns_none():
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
seen: dict[str, Any] = {}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
assert request.path == _REMOVE_PATH
counters["remove"] += 1
seen["request"] = request
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(204)
request.end_headers()
async with _mock_remote_db_async(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
current = await _observe_current_async(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
result = await db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert result is None
assert counters["lookup"] == 1
assert counters["remove"] == 1
assert seen.get("raw")
_assert_exact_remove_request(
seen["request"],
seen["raw"],
seen["body"],
expected_id=current.id,
)
def test_after_remove_name_lookup_not_found_id_lookup_same_function():
"""Catalog-pointer SDK sequence via a stateful fixture; not server atomicity."""
_assert_native_remove_method_present()
counters: dict[str, int] = {
"lookup_name": 0,
"lookup_id": 0,
"remove": 0,
}
removed = {"yes": False}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
body = json.loads(raw.decode("utf-8"))
if "name" in body:
counters["lookup_name"] += 1
assert body == {"name": _REMOVE_CATALOG_NAME}
if removed["yes"]:
request.send_response(404)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"error_code": "name_or_function_not_found",
"message": _REMOVE_SERVER_MESSAGE_MARKER,
}
).encode("utf-8")
)
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
counters["lookup_id"] += 1
assert body == {"function_id": _REMOVE_FUNCTION_ID}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
assert request.path == _REMOVE_PATH
counters["remove"] += 1
body = json.loads(raw.decode("utf-8"))
_assert_exact_remove_request(
request, raw, body, expected_id=_REMOVE_FUNCTION_ID
)
removed["yes"] = True
request.send_response(204)
request.end_headers()
function_error = _function_error_cls()
with _mock_remote_db(handler) as db:
current = _observe_current(db)
assert counters["lookup_name"] == 1
assert counters["lookup_id"] == 0
assert counters["remove"] == 0
result = db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert result is None
assert counters["lookup_name"] == 1
assert counters["remove"] == 1
with pytest.raises(function_error) as exc_info:
db.functions.get(_REMOVE_CATALOG_NAME)
err = exc_info.value
assert err.code == "name_or_function_not_found"
_assert_payload_free(err)
by_id = db.functions.get_by_id(_REMOVE_FUNCTION_ID)
assert counters["lookup_name"] == 2
assert counters["lookup_id"] == 1
assert counters["remove"] == 1
_assert_exact_remove_function(by_id)
assert by_id.id == current.id
assert by_id.parameters == current.parameters
assert by_id.output_type == current.output_type
assert by_id.output_nullable is current.output_nullable
def test_explicit_name_conflict_is_function_error_payload_free():
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
body = {
"error_code": "name_conflict",
"message": (
f"{_REMOVE_SERVER_MESSAGE_MARKER} looks_like {_CONFLICTING_MESSAGE_CODE}"
),
"name": _REMOVE_CATALOG_NAME,
"function_id": _REMOVE_FUNCTION_ID,
_SENSITIVE_BODY_MARKER: True,
}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REMOVE_PATH
counters["remove"] += 1
parsed = json.loads(payload.decode("utf-8"))
_assert_exact_remove_request(
request, payload, parsed, expected_id=_REMOVE_FUNCTION_ID
)
request.send_response(409)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(function_error) as exc_info:
db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert counters["lookup"] == 1
assert counters["remove"] == 1
err = exc_info.value
assert isinstance(err, function_error)
assert err.code == "name_conflict"
assert err.code != _CONFLICTING_MESSAGE_CODE
_assert_payload_free(err)
@pytest.mark.parametrize(
"label,status,response_body",
[
(
"200_with_body",
200,
{
"ok": True,
"message": _REMOVE_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
"job_id": "must-not-infer-job",
},
),
(
"202_empty",
202,
f"{_REMOVE_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
),
("200_empty", 200, ""),
],
)
def test_http_200_202_cannot_return_success(
label: str, status: int, response_body: object
):
del label
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REMOVE_PATH
counters["remove"] += 1
parsed = json.loads(payload.decode("utf-8"))
_assert_exact_remove_request(
request, payload, parsed, expected_id=_REMOVE_FUNCTION_ID
)
request.send_response(status)
request.send_header("Content-Type", "application/json")
request.end_headers()
if isinstance(response_body, str):
request.wfile.write(response_body.encode("utf-8"))
else:
request.wfile.write(json.dumps(response_body).encode("utf-8"))
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(HttpError) as exc_info:
db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert counters["lookup"] == 1
assert counters["remove"] == 1
err = exc_info.value
assert isinstance(err, HttpError)
assert not isinstance(err, _function_error_cls(required=False))
_assert_payload_free(err)
def test_empty_name_rejects_before_remove_transport():
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(ValueError):
db.functions.remove("", current)
assert counters["lookup"] == 1
assert counters["remove"] == 0
@pytest.mark.parametrize(
"bad_current",
[
_REMOVE_FUNCTION_ID,
{"id": _REMOVE_FUNCTION_ID},
object(),
123,
],
)
def test_raw_id_or_arbitrary_current_rejected_without_remove(bad_current):
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
# Observe a real handle separately so the bad-current path is isolated.
_ = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(TypeError):
db.functions.remove(_REMOVE_CATALOG_NAME, bad_current)
assert counters["lookup"] == 1
assert counters["remove"] == 0
def test_local_sync_remove_not_implemented_without_table_mutation(tmp_path):
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
assert type(current) is lancedb.Function
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
with pytest.raises(NotImplementedError):
db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert db.list_tables().tables == before
_close_db(db)
@pytest.mark.asyncio
async def test_local_async_remove_not_implemented_without_table_mutation(tmp_path):
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
assert type(current) is lancedb.Function
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
with pytest.raises(NotImplementedError):
await db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert (await db.list_tables()).tables == before
db.close()
@pytest.mark.parametrize("keyword", _DELETED_REMOVE_KEYWORDS)
def test_remove_rejects_deleted_cas_retry_version_kwargs(keyword):
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(TypeError):
db.functions.remove(
_REMOVE_CATALOG_NAME,
current,
**{keyword: True},
)
assert counters["lookup"] == 1
assert counters["remove"] == 0
def test_no_direct_remove_methods_and_function_has_no_remove_facade_private():
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
assert callable(getattr(db.functions, "remove", None))
assert not hasattr(lancedb, "_SyncFunctions")
assert not hasattr(lancedb, "_AsyncFunctions")
assert type(db.functions).__name__.startswith("_")
assert "_SyncFunctions" not in getattr(lancedb, "__all__", [])
assert "_AsyncFunctions" not in getattr(lancedb, "__all__", [])
assert counters["lookup"] == 1
assert counters["remove"] == 0
@@ -0,0 +1,579 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python conditional first-class Function replacement."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from datetime import timedelta
from typing import Any, Callable
from unittest import mock
import pyarrow as pa
import pytest
import lancedb
import lancedb._udf as _udf_mod
import lancedb.job
from lancedb import FunctionCapability, udf
from lancedb.exceptions import JobFailedError
_SOURCE_MARKER = "replace-source-marker-unique-xyz"
_SECRET_REFERENCE = "secret://team/replace-redact-token-xyz"
_SECRET_ENV = "REPLACE_API_TOKEN"
_NETWORK_ORIGIN = "https://api.replace-example.com"
_FUNCTION_NAME = "text.normalize"
_CURRENT_FUNCTION_ID = "fn.replace-current-1"
_REPLACED_FUNCTION_ID = "fn.replace-result-1"
_JOB_ID_SYNC = "job-replace-sync-1"
_JOB_ID_ASYNC = "job-replace-async-1"
_JOB_ID_CONFLICT = "job-replace-conflict-1"
_REGISTER_PATH = "/v1/functions/register"
_LOOKUP_PATH = "/v1/functions/lookup"
_DESCRIBE_PATH = "/v1/jobs/describe"
_CONFLICTING_MESSAGE_CODE = "definition_validation_failure"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
# job-result tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_REPLACE_KEYWORDS = (
"idempotency_key",
"retry_key",
"user_version",
"version",
"deterministic",
"null_policy",
"replace",
"expected_current_function_id",
"alias",
"lineage",
)
_SPEC_KEYS = {
"format_version",
"name",
"definition",
"expected_current_function_id",
}
@udf(
inputs={"text": pa.string(), "limit": pa.int32()},
output=pa.string(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=True,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
],
)
def packable_replace_normalize(text, limit):
"""replace-source-marker-unique-xyz."""
return text[:limit]
def _current_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _CURRENT_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _definition_json(fn: object) -> dict[str, Any]:
payload = _udf_mod._build_function_definition(fn)._to_json()
if isinstance(payload, bytes):
return json.loads(payload.decode("utf-8"))
assert isinstance(payload, str)
return json.loads(payload)
def _expected_replace_spec(name: str, current_id: str, fn: object) -> dict[str, Any]:
return {
"format_version": 1,
"name": name,
"definition": _definition_json(fn),
"expected_current_function_id": current_id,
}
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _assert_exact_replace_spec(
body: dict[str, Any], expected: dict[str, Any], current_id: str
) -> None:
assert set(body) == _SPEC_KEYS
assert body == expected
assert body["format_version"] == 1
assert body["expected_current_function_id"] == current_id
assert body["expected_current_function_id"] is not None
assert _SOURCE_MARKER in json.dumps(body["definition"])
assert any(
capability.get("reference") == _SECRET_REFERENCE
for capability in body["definition"]["capabilities"]
)
def _lookup_success_handler(
counters: dict[str, int],
*,
after_lookup: (
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
) = None,
):
"""Serve exact lookup; optionally continue for register/describe."""
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
counters["lookup"] = counters.get("lookup", 0) + 1
body = json.loads(raw.decode("utf-8"))
assert body == {"name": _FUNCTION_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps({"function": _current_function_wire()}).encode("utf-8")
)
return
if after_lookup is not None:
after_lookup(request, raw)
return
counters["register"] = counters.get("register", 0) + 1
request.send_response(500)
request.end_headers()
request.wfile.write(b"unexpected register")
return handler
def _observe_current(db) -> lancedb.Function:
current = db.functions.get(_FUNCTION_NAME)
assert type(current) is lancedb.Function
assert current.id == _CURRENT_FUNCTION_ID
assert not hasattr(current, "name")
assert not hasattr(current, "replace")
return current
async def _observe_current_async(db) -> lancedb.Function:
current = await db.functions.get(_FUNCTION_NAME)
assert type(current) is lancedb.Function
assert current.id == _CURRENT_FUNCTION_ID
assert not hasattr(current, "name")
assert not hasattr(current, "replace")
return current
def test_sync_remote_replace_exact_body_one_package_job_and_function_result():
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
expected_spec = _expected_replace_spec(
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
)
register_attempts: list[dict[str, Any]] = []
function_result_wire = {
"kind": "function",
"format_version": 1,
"function": {
"format_version": 1,
"id": _REPLACED_FUNCTION_ID,
"signature": expected_spec["definition"]["signature"],
},
}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
if request.path == _REGISTER_PATH:
counters["register"] += 1
register_attempts.append(
{
"request_id": request.headers.get("x-request-id"),
"raw": payload,
"body": json.loads(payload.decode("utf-8")),
}
)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_SYNC}).encode("utf-8"))
return
assert request.path == _DESCRIBE_PATH
counters["describe"] += 1
body = json.loads(payload.decode("utf-8"))
assert body["job_id"] == _JOB_ID_SYNC
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"job_id": _JOB_ID_SYNC,
"job_state": "DONE",
"job_type": "register_function",
"creation_ms": 1,
"spec": {},
"result": function_result_wire,
}
).encode("utf-8")
)
package_calls = {"n": 0}
original_package = _udf_mod._package_udf
def counting_package(fn: object):
package_calls["n"] += 1
return original_package(fn)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "replace_function")
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
job = db.functions.replace(
_FUNCTION_NAME, current, packable_replace_normalize
)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_SYNC
waited = job.wait(timeout=timedelta(seconds=5))
assert package_calls["n"] == 1
assert counters["lookup"] == 1
assert counters["register"] == 1
assert counters["describe"] == 1
assert len(register_attempts) == 1
attempt = register_attempts[0]
assert isinstance(attempt["request_id"], str) and attempt["request_id"]
assert attempt["raw"]
_assert_exact_replace_spec(attempt["body"], expected_spec, current.id)
assert type(waited) is lancedb.Function
assert waited.id == _REPLACED_FUNCTION_ID
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
assert waited.output_type == pa.string()
assert waited.output_nullable is True
@pytest.mark.asyncio
async def test_async_remote_replace_exact_body_returns_async_job():
counters: dict[str, int] = {"lookup": 0, "register": 0}
expected_spec = _expected_replace_spec(
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
)
seen: dict[str, Any] = {}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
assert request.path == _REGISTER_PATH
counters["register"] += 1
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
async with _mock_remote_db_async(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "replace_function")
current = await _observe_current_async(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
job = await db.functions.replace(
_FUNCTION_NAME, current, packable_replace_normalize
)
assert counters["lookup"] == 1
assert counters["register"] == 1
assert seen.get("raw")
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
_assert_exact_replace_spec(seen["body"], expected_spec, current.id)
assert type(job) is lancedb.job.AsyncJob
assert job.id == _JOB_ID_ASYNC
def test_sync_remote_replace_failed_name_conflict_raises_job_failed_error_code():
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
if request.path == _REGISTER_PATH:
counters["register"] += 1
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps({"job_id": _JOB_ID_CONFLICT}).encode("utf-8")
)
return
assert request.path == _DESCRIBE_PATH
counters["describe"] += 1
body = json.loads(raw.decode("utf-8"))
assert body["job_id"] == _JOB_ID_CONFLICT
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"job_id": _JOB_ID_CONFLICT,
"job_type": "register_function",
"job_state": "FAILED",
"creation_ms": 1,
"spec": {},
"failure": {
"phase": "validate",
"message": (
f"looks like {_CONFLICTING_MESSAGE_CODE} during CAS"
),
"retryable": False,
"error_code": "name_conflict",
},
}
).encode("utf-8")
)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
job = db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
assert type(job) is lancedb.job.Job
with pytest.raises(JobFailedError) as exc_info:
job.wait(timeout=timedelta(seconds=5))
assert counters["lookup"] == 1
assert counters["register"] == 1
assert counters["describe"] == 1
err = exc_info.value
assert isinstance(err, JobFailedError)
assert err.error_code == "name_conflict"
assert err.error_code != _CONFLICTING_MESSAGE_CODE
def test_empty_name_rejects_before_register_transport():
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with pytest.raises(ValueError):
db.functions.replace("", current, packable_replace_normalize)
assert counters["lookup"] == 1
assert counters["register"] == 0
@pytest.mark.parametrize(
"bad_current",
[
_CURRENT_FUNCTION_ID,
{"id": _CURRENT_FUNCTION_ID},
object(),
123,
],
)
def test_raw_id_or_arbitrary_current_rejected_without_register(bad_current):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
# Observe a real handle separately so the bad-current path is isolated.
_ = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with pytest.raises(TypeError):
db.functions.replace(
_FUNCTION_NAME, bad_current, packable_replace_normalize
)
assert counters["lookup"] == 1
assert counters["register"] == 0
def test_local_sync_replace_not_implemented_without_table_mutation(tmp_path):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["register"] == 0
assert type(current) is lancedb.Function
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "replace_function")
with pytest.raises(NotImplementedError):
db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
assert db.list_tables().tables == before
@pytest.mark.asyncio
async def test_local_async_replace_not_implemented_without_table_mutation(tmp_path):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["register"] == 0
assert type(current) is lancedb.Function
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "replace_function")
with pytest.raises(NotImplementedError):
await db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
assert (await db.list_tables()).tables == before
@pytest.mark.parametrize("keyword", _DELETED_REPLACE_KEYWORDS)
def test_replace_rejects_deleted_cas_retry_version_kwargs(keyword):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with pytest.raises(TypeError):
db.functions.replace(
_FUNCTION_NAME,
current,
packable_replace_normalize,
**{keyword: True},
)
assert counters["lookup"] == 1
assert counters["register"] == 0
def test_no_direct_replace_function_methods_and_function_has_no_replace():
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert not hasattr(db, "replace_function")
assert not hasattr(db, "register_function")
assert not hasattr(current, "replace")
assert not hasattr(current, "replace_function")
assert callable(getattr(db.functions, "replace", None))
assert counters["lookup"] == 1
assert counters["register"] == 0
@@ -0,0 +1,728 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python exact first-class Function revocation."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable
import pyarrow as pa
import pytest
import lancedb
from lancedb import _lancedb as _native
from lancedb.remote.errors import HttpError
_REVOKE_PATH = "/v1/functions/revoke"
_LOOKUP_PATH = "/v1/functions/lookup"
_REVOKE_CATALOG_NAME = "text.normalize.revoke-name"
_REVOKE_FUNCTION_ID = "fn.exact.revoke-handle"
_REVOKE_SERVER_MESSAGE_MARKER = (
"SERVER_REVOKE_DIAGNOSTIC_MARKER id=fn.exact.revoke-handle "
"name=text.normalize.revoke-name"
)
_SENSITIVE_BODY_MARKER = "SENSITIVE_REVOKE_BODY_MARKER"
_CONFLICTING_MESSAGE_CODE = "revoked_function"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
# remove tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_REVOKE_KEYWORDS = (
"function_id",
"name",
"idempotency_key",
"retry_key",
"user_version",
"version",
"reason",
"expiry",
"force",
"if_exists",
"remove",
"delete",
)
def _function_error_cls(*, required: bool = True):
"""Resolve FunctionError from the live module (records RED when absent)."""
from lancedb import exceptions as exc_mod
cls = getattr(exc_mod, "FunctionError", None)
if cls is None:
if required:
pytest.fail("lancedb.exceptions.FunctionError is missing")
return type("MissingFunctionError", (), {})
return cls
def _sample_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _REVOKE_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _lookup_success_body() -> bytes:
return json.dumps({"function": _sample_function_wire()}).encode("utf-8")
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
def _close_db(db: Any) -> None:
with contextlib.suppress(Exception):
inner = getattr(db, "_conn", None)
if inner is not None:
inner.close()
return
close = getattr(db, "close", None)
if callable(close):
close()
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
db = None
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
if db is not None:
_close_db(db)
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
db = None
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
if db is not None:
_close_db(db)
server.shutdown()
thread.join()
def _exception_chain_text(exc: BaseException) -> str:
parts: list[str] = []
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(str(current))
parts.append(repr(current))
current = current.__cause__
return "\n".join(parts)
def _assert_payload_free(exc: BaseException) -> None:
text = _exception_chain_text(exc)
assert _REVOKE_SERVER_MESSAGE_MARKER not in text
assert _REVOKE_CATALOG_NAME not in text
assert _REVOKE_FUNCTION_ID not in text
assert _SENSITIVE_BODY_MARKER not in text
def _assert_exact_revoke_function(function: object) -> None:
assert type(function) is lancedb.Function
assert function.id == _REVOKE_FUNCTION_ID
assert not hasattr(function, "name")
assert function.parameters == (
("text", pa.string()),
("limit", pa.int32()),
)
assert function.output_type == pa.string()
assert function.output_nullable is True
assert _REVOKE_CATALOG_NAME not in repr(function)
assert _REVOKE_CATALOG_NAME not in str(function)
def _assert_exact_revoke_request(
request: http.server.BaseHTTPRequestHandler,
raw: bytes,
body: dict[str, Any],
*,
expected_id: str,
) -> None:
assert request.command == "POST"
assert request.path == _REVOKE_PATH
assert "?" not in request.path
assert "remove" not in request.path
assert raw
assert body == {"function_id": expected_id}
assert set(body) == {"function_id"}
assert "name" not in body
assert "expected_current_function_id" not in body
assert "format_version" not in body
assert "function" not in body
assert "signature" not in body
assert "job_id" not in body
assert "idempotency_key" not in body
assert "user_version" not in body
assert "reason" not in body
assert "expiry" not in body
assert "force" not in body
request_id = request.headers.get("x-request-id")
assert isinstance(request_id, str) and request_id
def _assert_native_revoke_method_present() -> None:
assert hasattr(_native.Connection, "_revoke_function")
assert callable(getattr(_native.Connection, "_revoke_function"))
def _lookup_success_handler(
counters: dict[str, int],
*,
after_lookup: (
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
) = None,
):
"""Serve exact name lookup; optionally continue for revoke."""
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
counters["lookup"] = counters.get("lookup", 0) + 1
body = json.loads(raw.decode("utf-8"))
assert body == {"name": _REVOKE_CATALOG_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
if after_lookup is not None:
after_lookup(request, raw)
return
counters["revoke"] = counters.get("revoke", 0) + 1
request.send_response(500)
request.end_headers()
request.wfile.write(b"unexpected revoke")
return handler
def _observe_current(db) -> lancedb.Function:
current = db.functions.get(_REVOKE_CATALOG_NAME)
_assert_exact_revoke_function(current)
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
return current
async def _observe_current_async(db) -> lancedb.Function:
current = await db.functions.get(_REVOKE_CATALOG_NAME)
_assert_exact_revoke_function(current)
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
return current
def test_native_connection_exposes_private_revoke_function():
_assert_native_revoke_method_present()
def test_sync_remote_revoke_exact_body_path_request_id_returns_none():
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
revoke_attempts: list[dict[str, Any]] = []
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
body = json.loads(payload.decode("utf-8"))
revoke_attempts.append(
{
"request": request,
"raw": payload,
"body": body,
"request_id": request.headers.get("x-request-id"),
}
)
# Illegal body on 204 must be ignored; success is status-driven only.
request.send_response(204)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
_SENSITIVE_BODY_MARKER: True,
"message": _REVOKE_SERVER_MESSAGE_MARKER,
}
).encode("utf-8")
)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "revoke_function")
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
result = db.functions.revoke(current)
assert result is None
assert counters["lookup"] == 1
assert counters["revoke"] == 1
assert len(revoke_attempts) == 1
attempt = revoke_attempts[0]
_assert_exact_revoke_request(
attempt["request"],
attempt["raw"],
attempt["body"],
expected_id=current.id,
)
assert attempt["body"]["function_id"] == current.id
_assert_exact_revoke_function(current)
@pytest.mark.asyncio
async def test_async_remote_revoke_exact_body_returns_none():
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
seen: dict[str, Any] = {}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
seen["request"] = request
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(204)
request.end_headers()
async with _mock_remote_db_async(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "revoke_function")
current = await _observe_current_async(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
result = await db.functions.revoke(current)
assert result is None
assert counters["lookup"] == 1
assert counters["revoke"] == 1
assert seen.get("raw")
_assert_exact_revoke_request(
seen["request"],
seen["raw"],
seen["body"],
expected_id=current.id,
)
def test_repeated_remote_revoke_204_both_return_none():
"""Two logical calls each receiving 204 both succeed (Python outcome only)."""
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
revoke_request_ids: list[str] = []
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
body = json.loads(payload.decode("utf-8"))
_assert_exact_revoke_request(
request, payload, body, expected_id=_REVOKE_FUNCTION_ID
)
request_id = request.headers.get("x-request-id")
assert isinstance(request_id, str) and request_id
revoke_request_ids.append(request_id)
request.send_response(204)
request.end_headers()
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
first = db.functions.revoke(current)
second = db.functions.revoke(current)
assert first is None
assert second is None
assert counters["lookup"] == 1
assert counters["revoke"] == 2
assert len(revoke_request_ids) == 2
_assert_exact_revoke_function(current)
def test_after_revoke_name_and_id_lookup_still_return_function():
"""Revoke does not unlink names; SDK-visible sequence only, not Sophon proof."""
_assert_native_revoke_method_present()
counters: dict[str, int] = {
"lookup_name": 0,
"lookup_id": 0,
"revoke": 0,
}
revoked = {"yes": False}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
body = json.loads(raw.decode("utf-8"))
if "name" in body:
counters["lookup_name"] += 1
assert body == {"name": _REVOKE_CATALOG_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
counters["lookup_id"] += 1
assert body == {"function_id": _REVOKE_FUNCTION_ID}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
body = json.loads(raw.decode("utf-8"))
_assert_exact_revoke_request(
request, raw, body, expected_id=_REVOKE_FUNCTION_ID
)
revoked["yes"] = True
request.send_response(204)
request.end_headers()
with _mock_remote_db(handler) as db:
current = _observe_current(db)
assert counters["lookup_name"] == 1
assert counters["lookup_id"] == 0
assert counters["revoke"] == 0
assert not revoked["yes"]
result = db.functions.revoke(current)
assert result is None
assert counters["lookup_name"] == 1
assert counters["revoke"] == 1
assert revoked["yes"]
by_name = db.functions.get(_REVOKE_CATALOG_NAME)
by_id = db.functions.get_by_id(_REVOKE_FUNCTION_ID)
assert counters["lookup_name"] == 2
assert counters["lookup_id"] == 1
assert counters["revoke"] == 1
_assert_exact_revoke_function(by_name)
_assert_exact_revoke_function(by_id)
assert by_name.id == current.id
assert by_id.id == current.id
assert by_name.parameters == current.parameters
assert by_id.parameters == current.parameters
assert by_name.output_type == current.output_type
assert by_id.output_type == current.output_type
assert by_name.output_nullable is current.output_nullable
assert by_id.output_nullable is current.output_nullable
def test_explicit_name_or_function_not_found_is_function_error_payload_free():
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
body = {
"error_code": "name_or_function_not_found",
"message": (
f"{_REVOKE_SERVER_MESSAGE_MARKER} looks_like {_CONFLICTING_MESSAGE_CODE} "
"name_conflict"
),
"name": _REVOKE_CATALOG_NAME,
"function_id": _REVOKE_FUNCTION_ID,
_SENSITIVE_BODY_MARKER: True,
}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
parsed = json.loads(payload.decode("utf-8"))
_assert_exact_revoke_request(
request, payload, parsed, expected_id=_REVOKE_FUNCTION_ID
)
request.send_response(404)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
with pytest.raises(function_error) as exc_info:
db.functions.revoke(current)
assert counters["lookup"] == 1
assert counters["revoke"] == 1
err = exc_info.value
assert isinstance(err, function_error)
assert err.code == "name_or_function_not_found"
assert err.code != _CONFLICTING_MESSAGE_CODE
assert err.code != "name_conflict"
_assert_payload_free(err)
@pytest.mark.parametrize(
"label,status,response_body",
[
(
"200_with_body",
200,
{
"ok": True,
"message": _REVOKE_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
"job_id": "must-not-infer-job",
},
),
(
"202_empty",
202,
f"{_REVOKE_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
),
("200_empty", 200, ""),
],
)
def test_http_200_202_cannot_return_success(
label: str, status: int, response_body: object
):
del label
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
parsed = json.loads(payload.decode("utf-8"))
_assert_exact_revoke_request(
request, payload, parsed, expected_id=_REVOKE_FUNCTION_ID
)
request.send_response(status)
request.send_header("Content-Type", "application/json")
request.end_headers()
if isinstance(response_body, str):
request.wfile.write(response_body.encode("utf-8"))
else:
request.wfile.write(json.dumps(response_body).encode("utf-8"))
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
with pytest.raises(HttpError) as exc_info:
db.functions.revoke(current)
assert counters["lookup"] == 1
assert counters["revoke"] == 1
err = exc_info.value
assert isinstance(err, HttpError)
assert not isinstance(err, _function_error_cls(required=False))
_assert_payload_free(err)
@pytest.mark.parametrize(
"bad_function",
[
_REVOKE_FUNCTION_ID,
{"id": _REVOKE_FUNCTION_ID},
object(),
123,
],
)
def test_raw_id_or_arbitrary_function_rejected_without_revoke(bad_function):
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
# Observe a real handle separately so the bad-function path is isolated.
_ = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
with pytest.raises(TypeError):
db.functions.revoke(bad_function)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
def test_local_sync_revoke_not_implemented_without_table_mutation(tmp_path):
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
assert type(current) is lancedb.Function
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "revoke_function")
with pytest.raises(NotImplementedError):
db.functions.revoke(current)
assert db.list_tables().tables == before
_close_db(db)
@pytest.mark.asyncio
async def test_local_async_revoke_not_implemented_without_table_mutation(tmp_path):
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
assert type(current) is lancedb.Function
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "revoke_function")
with pytest.raises(NotImplementedError):
await db.functions.revoke(current)
assert (await db.list_tables()).tables == before
db.close()
@pytest.mark.parametrize("keyword", _DELETED_REVOKE_KEYWORDS)
def test_revoke_rejects_overdesigned_kwargs(keyword):
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
with pytest.raises(TypeError):
db.functions.revoke(current, **{keyword: True})
assert counters["lookup"] == 1
assert counters["revoke"] == 0
def test_no_direct_revoke_methods_and_function_has_no_revoke_facade_private():
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert not hasattr(db, "revoke_function")
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
assert callable(getattr(db.functions, "revoke", None))
assert not hasattr(lancedb, "_SyncFunctions")
assert not hasattr(lancedb, "_AsyncFunctions")
assert type(db.functions).__name__.startswith("_")
assert "_SyncFunctions" not in getattr(lancedb, "__all__", [])
assert "_AsyncFunctions" not in getattr(lancedb, "__all__", [])
assert counters["lookup"] == 1
assert counters["revoke"] == 0
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,899 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python ``table.add_generated_column`` (FF-032).
Public user shape under test:
job = table.add_generated_column(
"normalized_text",
normalize(text=col("text")),
)
job.wait()
These tests exercise the live worktree PyO3 extension and public sync/async
wrappers. While the public methods and hidden native bridge are absent they
fail against that extension; once present they freeze the public contract
below. They must not fake success paths.
"""
from __future__ import annotations
import contextlib
import http.server
import inspect
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable
import pytest
import lancedb
import lancedb.job
from lancedb import _lancedb as _native
from lancedb.expr import col
from lancedb.remote.table import RemoteTable
from lancedb.table import AsyncTable, LanceTable, Table
_LOOKUP_PATH = "/v1/functions/lookup"
_JOB_DESCRIBE_PATH = "/v1/jobs/describe"
_TABLE_NAME = "articles"
_DESCRIBE_PATH = f"/v1/table/{_TABLE_NAME}/describe/"
_CREATE_PATH = f"/v1/table/{_TABLE_NAME}/generated_columns/create/"
_BRANCHES_CREATE_PATH = f"/v1/table/{_TABLE_NAME}/branches/create/"
_BRANCHES_LIST_PATH = f"/v1/table/{_TABLE_NAME}/branches/list/"
_CATALOG_NAME = "text.normalize"
_FUNCTION_ID = "fn.exact.normalize.gen-col"
_JOB_ID_SYNC = "job-create-gen-col-sync-1"
_JOB_ID_ASYNC = "job-create-gen-col-async-1"
_JOB_ID_BRANCH = "job-create-gen-col-branch-1"
_SOURCE_TABLE_VERSION = 42
_TEXT_FIELD_ID = 7
_BRANCH_NAME = "exp"
_BRANCH_SOURCE_VERSION = 9
_BRANCH_TEXT_FIELD_ID = 11
_DESCRIBE_BODY_MARKER = "SENSITIVE_DESCRIBE_BODY_MARKER_gen_col_xyz"
_CREATE_RESPONSE_MARKER = "SENSITIVE_CREATE_RESPONSE_MARKER_gen_col_xyz"
_LITERAL_PAYLOAD_SENTINEL = "LITERAL_PAYLOAD_SENTINEL_gen_col_xyz"
# Pinned Rust-canonical schema-only Utf8 type IPC (base64), shared with FF-028.
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_FORBIDDEN_PUBLIC_NAMES = (
"FunctionCall",
"BoundFunctionCall",
"AuthoredFunctionCall",
"CreateGeneratedColumnRequest",
"CreateGeneratedColumnJobSpec",
"GeneratedColumnBindingSnapshot",
"GeneratedColumnCreateRequest",
"geneva",
"GenevaFunction",
"VirtualColumnDefinition",
)
_FORBIDDEN_METHOD_KWARGS = (
"source_table_version",
"version",
"field_id",
"field_ids",
"output",
"output_type",
"output_nullable",
"nullable",
"spec",
"retry_key",
"idempotency_key",
"request",
"envelope",
"table_ref",
"branch",
)
def _sample_function_wire(
*,
function_id: str = _FUNCTION_ID,
parameters: list[dict[str, str]] | None = None,
) -> dict[str, Any]:
return {
"format_version": 1,
"id": function_id,
"signature": {
"parameters": parameters
or [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _text_schema_fields(
*, arrow_type: str = "string", nullable: bool = True
) -> dict[str, Any]:
return {
"fields": [
{
"name": "text",
"type": {"type": arrow_type},
"nullable": nullable,
}
]
}
def _describe_body(
*,
version: int = _SOURCE_TABLE_VERSION,
field_ids: list[int] | None = None,
arrow_type: str = "string",
include_marker: bool = True,
) -> dict[str, Any]:
body: dict[str, Any] = {
"version": version,
"schema": _text_schema_fields(arrow_type=arrow_type),
"field_ids": field_ids if field_ids is not None else [_TEXT_FIELD_ID],
}
if include_marker:
body["server_diagnostic"] = _DESCRIBE_BODY_MARKER
return body
def _create_gen_column_done_body(job_id: str) -> dict[str, Any]:
# DONE with omitted result: create_gen_column projects JobResult::None.
return {
"job_id": job_id,
"job_state": "DONE",
"job_type": "create_gen_column",
"creation_ms": 1,
"spec": {},
}
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _exception_text(exc: BaseException) -> str:
parts = [str(exc), repr(exc)]
current: BaseException | None = exc
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(f"{type(current).__name__}: {current}")
current = current.__cause__ or current.__context__
return "\n".join(parts)
def _json_response(
request: http.server.BaseHTTPRequestHandler, body: dict[str, Any]
) -> None:
payload = json.dumps(body).encode("utf-8")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(payload)
def _lookup_function(db: Any) -> lancedb.Function:
return db.functions.get(_CATALOG_NAME)
class _RequestLog:
"""Track lookup/describe/create after setup; setup traffic is excluded."""
def __init__(self) -> None:
self.lookup: list[dict[str, Any]] = []
self.describe: list[dict[str, Any]] = []
self.create: list[dict[str, Any]] = []
self.other_table: list[str] = []
self.recording = False
def start(self) -> None:
# Drop setup's explicit Function lookup and open_table describe so
# operation accounting cannot be polluted by fixture traffic.
self.lookup.clear()
self.describe.clear()
self.create.clear()
self.other_table.clear()
self.recording = True
def note(self, path: str, body: dict[str, Any] | None = None) -> None:
if not self.recording:
return
if path == _LOOKUP_PATH:
self.lookup.append(body or {})
elif path == _DESCRIBE_PATH:
self.describe.append(body or {})
elif path == _CREATE_PATH:
self.create.append(body or {})
elif path.startswith(f"/v1/table/{_TABLE_NAME}/"):
self.other_table.append(path)
def _assert_no_operation_traffic(log: _RequestLog) -> None:
assert log.lookup == []
assert log.describe == []
assert log.create == []
assert log.other_table == []
def _assert_exact_public_signature(method: Any) -> None:
"""Freeze ``(self, column_name, call)`` with no varargs/kwargs escape hatches."""
params = list(inspect.signature(method).parameters.values())
assert [p.name for p in params] == ["self", "column_name", "call"]
for param in params:
assert param.kind in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
)
assert param.default is inspect.Parameter.empty
assert param.kind is not inspect.Parameter.VAR_POSITIONAL
assert param.kind is not inspect.Parameter.VAR_KEYWORD
assert param.kind is not inspect.Parameter.KEYWORD_ONLY
def _open_table_and_function(
*,
describe_body: dict[str, Any] | None = None,
on_create: Callable[[dict[str, Any], http.server.BaseHTTPRequestHandler], None]
| None = None,
job_id: str = _JOB_ID_SYNC,
support_branch_create: bool = False,
function_wire: dict[str, Any] | None = None,
):
"""Open remote table + immutable Function; return (db, table, function, log, cm)."""
log = _RequestLog()
binding_describe = describe_body or _describe_body()
open_describe = {
"version": 1,
"schema": _text_schema_fields(),
}
state = {"opened": False}
wire = function_wire or _sample_function_wire()
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _LOOKUP_PATH:
log.note(request.path, body)
_json_response(request, {"function": wire})
return
if request.path == _JOB_DESCRIBE_PATH:
assert body["job_id"] == job_id
_json_response(request, _create_gen_column_done_body(job_id))
return
if support_branch_create and request.path == _BRANCHES_CREATE_PATH:
log.note(request.path, body)
_json_response(request, {})
return
if support_branch_create and request.path == _BRANCHES_LIST_PATH:
log.note(request.path, body)
_json_response(
request,
{
"branches": {
_BRANCH_NAME: {
"parentBranch": None,
"parentVersion": 1,
"createAt": 1,
"manifestSize": 1,
}
}
},
)
return
if request.path == _DESCRIBE_PATH:
# First describe seeds open_table; later ones are binding snapshots.
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding_describe)
return
if request.path == _CREATE_PATH:
log.note(request.path, body)
if on_create is not None:
on_create(body, request)
return
_json_response(
request,
{
"job_id": job_id,
"server_extra": {"marker": _CREATE_RESPONSE_MARKER},
},
)
return
if request.path.startswith(f"/v1/table/{_TABLE_NAME}/"):
log.note(request.path, body)
request.send_response(404)
request.end_headers()
request.wfile.write(b"unexpected path")
cm = _mock_remote_db(handler)
db = cm.__enter__()
function = _lookup_function(db)
table = db.open_table(_TABLE_NAME)
assert isinstance(table, RemoteTable)
# open_table consumed the seed describe; binding/create accounting starts now.
# Setup's one explicit lookup is cleared here and must not pollute counts.
log.start()
return db, table, function, log, cm
def _assert_exact_create_envelope(
body: dict[str, Any],
*,
source_table_version: int,
column_name: str,
field_id: int,
branch: str | None = None,
) -> None:
expected_keys = {"source_table_version", "spec"}
if branch is not None:
expected_keys.add("branch")
assert set(body) == expected_keys
assert body["source_table_version"] == source_table_version
assert "table_ref" not in body
if branch is None:
assert "branch" not in body
else:
assert body["branch"] == branch
spec = body["spec"]
assert set(spec) == {"format_version", "column_name", "function_call"}
assert spec["format_version"] == 1
assert spec["column_name"] == column_name
for forbidden in (
"table_ref",
"source_table_version",
"version",
"output",
"output_type",
"output_field_id",
"dependency_epoch",
"materialized_epoch",
"idempotency_key",
"retry_key",
"name",
"handle",
"artifact",
"geneva",
):
assert forbidden not in spec
call = spec["function_call"]
assert set(call) == {"function_id", "arguments"}
assert call["function_id"] == _FUNCTION_ID
assert len(call["arguments"]) == 1
binding = call["arguments"][0]
assert binding["parameter"] == "text"
value = binding["value"]
assert value["kind"] == "field"
assert value["field_id"] == field_id
assert value["data_type_ipc"] == _UTF8_TYPE_IPC_B64
assert "name" not in value
assert "column_name" not in value
assert "text" not in value
# Serialized call must not late-bind by column name anywhere relevant.
dumped = json.dumps(call)
assert '"column_name"' not in dumped
assert "normalized_text" not in dumped
def test_public_and_native_add_generated_column_seams_must_exist():
"""Public sync/async methods and the private native bridge must exist."""
assert hasattr(_native.Table, "_add_generated_column"), (
"native private bridge Table._add_generated_column is missing"
)
assert hasattr(AsyncTable, "add_generated_column"), (
"AsyncTable.add_generated_column is missing"
)
assert hasattr(Table, "add_generated_column"), (
"Table.add_generated_column is missing"
)
assert hasattr(LanceTable, "add_generated_column"), (
"LanceTable.add_generated_column is missing"
)
assert hasattr(RemoteTable, "add_generated_column"), (
"RemoteTable.add_generated_column is missing"
)
# Once present, freeze the exact public positional surface.
_assert_exact_public_signature(Table.add_generated_column)
_assert_exact_public_signature(LanceTable.add_generated_column)
_assert_exact_public_signature(RemoteTable.add_generated_column)
_assert_exact_public_signature(AsyncTable.add_generated_column)
def test_sync_remote_add_generated_column_returns_job_without_eager_wrapper_mutation():
db, table, normalize, log, cm = _open_table_and_function(job_id=_JOB_ID_SYNC)
try:
# Capture public wrapper state before the operation window.
schema_before = table.schema
version_before = table.version
log.start()
call = normalize(text=col("text"))
# Exact public argument order from the frozen user example.
job = table.add_generated_column(
"normalized_text",
call,
)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_SYNC
# Exact success path stops after submit: one binding describe, one create,
# and no catalog re-lookup. Do not wait yet.
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert len(log.create) == 1
assert log.other_table == []
# Public schema/version through the existing wrapper must still reflect
# the pre-submit table: generated column is not published by Job accept.
# Access both before wait so eager wrapper cache invalidation / refresh /
# version advancement is observable.
schema_after = table.schema
assert "normalized_text" not in schema_after.names
assert schema_after == schema_before
# Schema must be served from the existing wrapper cache — no extra
# describe beyond the one binding snapshot.
assert len(log.describe) == 1
assert len(log.lookup) == 0
assert len(log.create) == 1
version_after = table.version
assert version_after == version_before
# Public Remote ``version`` always describes once by design; that probe
# must not drag a schema-cache miss, create, or catalog lookup with it.
assert len(log.describe) == 2
assert len(log.lookup) == 0
assert len(log.create) == 1
assert log.other_table == []
waited = job.wait()
assert waited is None
assert len(log.lookup) == 0
assert len(log.create) == 1
finally:
cm.__exit__(None, None, None)
@pytest.mark.asyncio
async def test_async_remote_add_generated_column_returns_async_job_and_wait_none():
log = _RequestLog()
state = {"opened": False}
binding_describe = _describe_body()
open_describe = {"version": 1, "schema": _text_schema_fields()}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _LOOKUP_PATH:
log.note(request.path, body)
_json_response(request, {"function": _sample_function_wire()})
return
if request.path == _JOB_DESCRIBE_PATH:
assert body["job_id"] == _JOB_ID_ASYNC
_json_response(request, _create_gen_column_done_body(_JOB_ID_ASYNC))
return
if request.path == _DESCRIBE_PATH:
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding_describe)
return
if request.path == _CREATE_PATH:
log.note(request.path, body)
_json_response(request, {"job_id": _JOB_ID_ASYNC})
return
request.send_response(404)
request.end_headers()
async with _mock_remote_db_async(handler) as db:
normalize = await db.functions.get(_CATALOG_NAME)
table = await db.open_table(_TABLE_NAME)
log.start()
call = normalize(text=col("text"))
job = await table.add_generated_column("normalized_text", call)
assert type(job) is lancedb.job.AsyncJob
assert job.id == _JOB_ID_ASYNC
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert len(log.create) == 1
waited = await job.wait()
assert waited is None
assert len(log.lookup) == 0
assert len(log.create) == 1
def test_remote_add_generated_column_one_describe_one_create_exact_envelope():
db, table, normalize, log, cm = _open_table_and_function(job_id=_JOB_ID_SYNC)
try:
call = normalize(text=col("text"))
job = table.add_generated_column("normalized_text", call)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_SYNC
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert len(log.create) == 1
assert log.other_table == []
_assert_exact_create_envelope(
log.create[0],
source_table_version=_SOURCE_TABLE_VERSION,
column_name="normalized_text",
field_id=_TEXT_FIELD_ID,
)
finally:
cm.__exit__(None, None, None)
def test_remote_branch_add_generated_column_includes_exact_branch_identity():
branch_describe = _describe_body(
version=_BRANCH_SOURCE_VERSION,
field_ids=[_BRANCH_TEXT_FIELD_ID],
)
db, table, normalize, log, cm = _open_table_and_function(
describe_body=branch_describe,
job_id=_JOB_ID_BRANCH,
support_branch_create=True,
)
try:
branched = table.branches.create(_BRANCH_NAME)
assert isinstance(branched, RemoteTable)
assert branched.current_branch() == _BRANCH_NAME
log.start()
call = normalize(text=col("text"))
job = branched.add_generated_column("normalized_text", call)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_BRANCH
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert len(log.create) == 1
assert log.describe[0].get("branch") == _BRANCH_NAME
_assert_exact_create_envelope(
log.create[0],
source_table_version=_BRANCH_SOURCE_VERSION,
column_name="normalized_text",
field_id=_BRANCH_TEXT_FIELD_ID,
branch=_BRANCH_NAME,
)
finally:
cm.__exit__(None, None, None)
def test_empty_column_name_fails_locally_with_zero_table_requests():
db, table, normalize, log, cm = _open_table_and_function()
try:
# Authored call owns a real literal so payload-free failure is not vacuous.
call = normalize(text=_LITERAL_PAYLOAD_SENTINEL)
with pytest.raises((ValueError, TypeError)) as raised:
table.add_generated_column("", call)
text = _exception_text(raised.value)
lowered = text.lower()
assert "column" in lowered or "empty" in lowered or "non-empty" in lowered
assert _DESCRIBE_BODY_MARKER not in text
assert _CREATE_RESPONSE_MARKER not in text
assert _LITERAL_PAYLOAD_SENTINEL not in text
_assert_no_operation_traffic(log)
finally:
cm.__exit__(None, None, None)
@pytest.mark.parametrize(
("column_ref", "expected_token"),
[
("missing_text", "missing_text"),
("Text", "Text"), # exact-case mismatch against schema field "text"
],
)
def test_missing_or_case_mismatch_column_one_describe_zero_create(
column_ref: str, expected_token: str
):
db, table, normalize, log, cm = _open_table_and_function()
try:
call = normalize(text=col(column_ref))
with pytest.raises(ValueError) as raised:
table.add_generated_column("normalized_text", call)
text = _exception_text(raised.value)
assert expected_token in text
assert "text" in text # parameter name from the Function signature
assert "missing" in text.lower() or "field" in text.lower()
assert _DESCRIBE_BODY_MARKER not in text
assert _CREATE_RESPONSE_MARKER not in text
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert log.create == []
assert log.other_table == []
finally:
cm.__exit__(None, None, None)
def test_type_mismatch_one_describe_zero_create_identifies_parameter():
db, table, normalize, log, cm = _open_table_and_function(
describe_body=_describe_body(arrow_type="int32"),
)
try:
call = normalize(text=col("text"))
with pytest.raises(ValueError) as raised:
table.add_generated_column("normalized_text", call)
text = _exception_text(raised.value)
assert "text" in text
assert "type" in text.lower() or "mismatch" in text.lower()
assert _DESCRIBE_BODY_MARKER not in text
assert _CREATE_RESPONSE_MARKER not in text
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert log.create == []
assert log.other_table == []
finally:
cm.__exit__(None, None, None)
def test_literal_payload_stays_out_of_field_binding_failure():
"""Authored call owns a real literal; later field binding fails payload-free."""
wire = _sample_function_wire(
parameters=[
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "prefix", "data_type_ipc": _UTF8_TYPE_IPC_B64},
]
)
db, table, normalize, log, cm = _open_table_and_function(function_wire=wire)
try:
call = normalize(text=col("missing_text"), prefix=_LITERAL_PAYLOAD_SENTINEL)
with pytest.raises(ValueError) as raised:
table.add_generated_column("normalized_text", call)
text = _exception_text(raised.value)
assert "missing_text" in text
assert _LITERAL_PAYLOAD_SENTINEL not in text
assert _DESCRIBE_BODY_MARKER not in text
assert _CREATE_RESPONSE_MARKER not in text
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert log.create == []
assert log.other_table == []
finally:
cm.__exit__(None, None, None)
@pytest.mark.asyncio
async def test_closed_async_table_fails_with_zero_operation_requests():
log = _RequestLog()
state = {"opened": False}
binding_describe = _describe_body()
open_describe = {"version": 1, "schema": _text_schema_fields()}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _LOOKUP_PATH:
log.note(request.path, body)
_json_response(request, {"function": _sample_function_wire()})
return
if request.path == _DESCRIBE_PATH:
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding_describe)
return
if request.path == _CREATE_PATH:
log.note(request.path, body)
_json_response(request, {"job_id": _JOB_ID_ASYNC})
return
request.send_response(404)
request.end_headers()
async with _mock_remote_db_async(handler) as db:
normalize = await db.functions.get(_CATALOG_NAME)
table = await db.open_table(_TABLE_NAME)
call = normalize(text=col("text"))
# Public close only — do not mutate private implementation fields.
table.close()
log.start()
try:
await table.add_generated_column("normalized_text", call)
except AttributeError:
# Method missing: re-raise so the failure names the public seam.
raise
except Exception as exc:
text = _exception_text(exc)
assert "closed" in text.lower()
else:
pytest.fail("closed AsyncTable must fail before transport")
_assert_no_operation_traffic(log)
def test_rejects_non_authored_call_before_any_operation_request():
db, table, normalize, log, cm = _open_table_and_function()
try:
bad_values = (
normalize, # exact Function handle itself
{"text": "x"},
col("text"), # direct query Expr
object(),
)
for bad in bad_values:
with pytest.raises(TypeError):
table.add_generated_column("normalized_text", bad)
_assert_no_operation_traffic(log)
finally:
cm.__exit__(None, None, None)
def test_native_valid_call_returns_not_supported_without_mutation(tmp_path):
# Immutable Function handle is connection-free; obtain it via remote lookup.
def lookup_only(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
_json_response(request, {"function": _sample_function_wire()})
with _mock_remote_db(lookup_only) as remote_db:
normalize = _lookup_function(remote_db)
db = lancedb.connect(tmp_path)
table = db.create_table(_TABLE_NAME, [{"text": "Hello"}, {"text": "World"}])
assert isinstance(table, LanceTable)
version_before = table.version
schema_before = table.schema
rows_before = table.to_arrow().to_pylist()
call = normalize(text=col("text"))
with pytest.raises(NotImplementedError) as raised:
table.add_generated_column("normalized_text", call)
text = _exception_text(raised.value)
assert "not supported" in text.lower() or "submit_create_generated_column" in text
assert "add_columns" not in text.lower()
assert table.version == version_before
assert table.schema == schema_before
assert "normalized_text" not in table.schema.names
assert table.to_arrow().to_pylist() == rows_before
def test_public_surface_is_minimal_and_private_call_stays_opaque():
for name in _FORBIDDEN_PUBLIC_NAMES:
assert name not in getattr(lancedb, "__all__", [])
assert not hasattr(lancedb, name)
assert not hasattr(lancedb, "_FunctionCall")
authored_type = getattr(_native, "_FunctionCall", None)
assert authored_type is not None
with pytest.raises(TypeError):
authored_type()
# When the public method exists, reject overdesign kwargs and keep the frozen
# positional surface: (self, column_name, call).
if hasattr(Table, "add_generated_column"):
_assert_exact_public_signature(Table.add_generated_column)
for keyword in _FORBIDDEN_METHOD_KWARGS:
assert (
keyword not in inspect.signature(Table.add_generated_column).parameters
)
db, table, normalize, log, cm = _open_table_and_function()
try:
call = normalize(text=col("text"))
for keyword in _FORBIDDEN_METHOD_KWARGS:
with pytest.raises(TypeError):
table.add_generated_column(
"normalized_text",
call,
**{keyword: object()},
)
_assert_no_operation_traffic(log)
finally:
cm.__exit__(None, None, None)
if hasattr(LanceTable, "add_generated_column"):
_assert_exact_public_signature(LanceTable.add_generated_column)
if hasattr(RemoteTable, "add_generated_column"):
_assert_exact_public_signature(RemoteTable.add_generated_column)
if hasattr(AsyncTable, "add_generated_column"):
_assert_exact_public_signature(AsyncTable.add_generated_column)
for keyword in _FORBIDDEN_METHOD_KWARGS:
assert (
keyword
not in inspect.signature(AsyncTable.add_generated_column).parameters
)
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,672 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python ``table.generated_column_status`` (B3d2).
Public user shape under test:
status = table.generated_column_status("complete_col") # "complete" | "incomplete"
These tests exercise the live worktree PyO3 extension and public sync/async
wrappers. While the public methods and hidden native bridge are absent they
fail against that extension; once present they freeze the public contract
below. They must not fake success paths.
"""
from __future__ import annotations
import contextlib
import http.server
import inspect
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable, Literal, get_type_hints
import pytest
import lancedb
import lancedb.table
from lancedb import _lancedb as _native
from lancedb.remote.table import RemoteTable
from lancedb.table import AsyncTable, LanceTable, Table
_TABLE_NAME = "articles"
_DESCRIBE_PATH = f"/v1/table/{_TABLE_NAME}/describe/"
_ORDINARY_FIELD_ID = 1
_COMPLETE_FIELD_ID = 5
_INCOMPLETE_FIELD_ID = 7
_STABLE_FIELD_IDS = [_ORDINARY_FIELD_ID, _COMPLETE_FIELD_ID, _INCOMPLETE_FIELD_ID]
_STATUS_FUNCTION_ID = "fn.exact.status.projection"
_METADATA_KEY = "lancedb::generated_column"
_RAW_METADATA_MARKER = "SENSITIVE_STATUS_METADATA_MARKER_b3d2_py_9f2e"
# Pinned Rust-canonical schema-only Utf8 type IPC (base64), shared with FF-028.
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_EXPECTED_RETURN = Literal["complete", "incomplete"]
_FORBIDDEN_PUBLIC_NAMES = (
"GeneratedColumnStatus",
"GeneratedColumnDefinition",
"GeneratedColumnBindingSnapshot",
"GeneratedColumnBindingEntry",
)
_FORBIDDEN_BRIDGE_KWARGS = (
"epoch",
"dependency_epoch",
"materialized_epoch",
"function_id",
"field_id",
"field_ids",
"version",
"branch",
"wait",
"job",
"request",
"backend",
)
def _definition_metadata_json(
output_field_id: int,
dependency_epoch: int,
materialized_epoch: int,
*,
text_field_id: int = _ORDINARY_FIELD_ID,
) -> str:
"""Exact JSON stored under Arrow field metadata ``lancedb::generated_column``."""
return json.dumps(
{
"format_version": 1,
"output_field_id": output_field_id,
"function_call": {
"function_id": _STATUS_FUNCTION_ID,
"arguments": [
{
"parameter": "text",
"value": {
"kind": "field",
"field_id": text_field_id,
"data_type_ipc": _UTF8_TYPE_IPC_B64,
},
}
],
},
"dependency_epoch": dependency_epoch,
"materialized_epoch": materialized_epoch,
},
separators=(",", ":"),
)
def _field(
name: str,
*,
arrow_type: str = "string",
nullable: bool = True,
metadata: dict[str, str] | None = None,
) -> dict[str, Any]:
body: dict[str, Any] = {
"name": name,
"type": {"type": arrow_type},
"nullable": nullable,
}
if metadata is not None:
body["metadata"] = metadata
return body
def _status_schema_fields(
*,
complete_meta: str | None = None,
incomplete_meta: str | None = None,
bad_name: str | None = None,
bad_meta: str | None = None,
) -> dict[str, Any]:
fields = [
_field("ordinary", arrow_type="string"),
_field(
"complete_col",
arrow_type="int32",
metadata={
_METADATA_KEY: complete_meta
if complete_meta is not None
else _definition_metadata_json(_COMPLETE_FIELD_ID, 3, 3)
},
),
_field(
"incomplete_col",
arrow_type="int32",
metadata={
_METADATA_KEY: incomplete_meta
if incomplete_meta is not None
else _definition_metadata_json(_INCOMPLETE_FIELD_ID, 4, 1)
},
),
]
if bad_name is not None and bad_meta is not None:
fields.append(
_field(
bad_name,
arrow_type="int32",
metadata={_METADATA_KEY: bad_meta},
)
)
return {"fields": fields}
def _describe_body(
*,
version: int = 11,
field_ids: list[int] | None = _STABLE_FIELD_IDS,
schema: dict[str, Any] | None = None,
) -> dict[str, Any]:
body: dict[str, Any] = {
"version": version,
"schema": schema if schema is not None else _status_schema_fields(),
}
if field_ids is not None:
body["field_ids"] = field_ids
return body
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _exception_text(exc: BaseException) -> str:
parts = [str(exc), repr(exc)]
current: BaseException | None = exc
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(f"{type(current).__name__}: {current}")
current = current.__cause__ or current.__context__
return "\n".join(parts)
def _json_response(
request: http.server.BaseHTTPRequestHandler, body: dict[str, Any]
) -> None:
payload = json.dumps(body).encode("utf-8")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(payload)
class _RequestLog:
"""Track post-open describe and any non-describe operation traffic."""
def __init__(self) -> None:
self.describe: list[dict[str, Any]] = []
self.other: list[str] = []
self.recording = False
def start(self) -> None:
self.describe.clear()
self.other.clear()
self.recording = True
def note(self, path: str, body: dict[str, Any] | None = None) -> None:
if not self.recording:
return
if path == _DESCRIBE_PATH:
self.describe.append(body or {})
else:
self.other.append(path)
def _assert_no_operation_traffic(log: _RequestLog) -> None:
assert log.describe == []
assert log.other == []
def _assert_one_status_describe(log: _RequestLog) -> None:
assert len(log.describe) == 1, f"expected one status describe, got {log.describe!r}"
assert log.other == [], f"unexpected non-describe traffic: {log.other!r}"
def _assert_exact_public_signature(method: Any) -> None:
"""Freeze ``(self, column_name)`` with no varargs/kwargs/keyword-only escape."""
params = list(inspect.signature(method).parameters.values())
assert [p.name for p in params] == ["self", "column_name"]
for param in params:
assert param.kind in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
)
assert param.default is inspect.Parameter.empty
assert param.kind is not inspect.Parameter.VAR_POSITIONAL
assert param.kind is not inspect.Parameter.VAR_KEYWORD
assert param.kind is not inspect.Parameter.KEYWORD_ONLY
def _assert_status_string(value: Any, expected: str) -> None:
assert value == expected
assert type(value) is str
assert value in ("complete", "incomplete")
def _open_remote_table(
*,
status_describe: dict[str, Any] | None = None,
):
"""Open sync RemoteTable; return (table, log, cm)."""
log = _RequestLog()
binding = status_describe if status_describe is not None else _describe_body()
open_describe = {
"version": 1,
"schema": {"fields": [_field("ordinary", arrow_type="string")]},
}
state = {"opened": False}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _DESCRIBE_PATH:
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding)
return
log.note(request.path, body)
request.send_response(404)
request.end_headers()
request.wfile.write(b"unexpected path")
cm = _mock_remote_db(handler)
db = cm.__enter__()
table = db.open_table(_TABLE_NAME)
assert isinstance(table, RemoteTable)
log.start()
return table, log, cm
async def _open_remote_table_async(
*,
status_describe: dict[str, Any] | None = None,
):
"""Open async table under a live mock server; return (table, log, cm)."""
log = _RequestLog()
binding = status_describe if status_describe is not None else _describe_body()
open_describe = {
"version": 1,
"schema": {"fields": [_field("ordinary", arrow_type="string")]},
}
state = {"opened": False}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _DESCRIBE_PATH:
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding)
return
log.note(request.path, body)
request.send_response(404)
request.end_headers()
request.wfile.write(b"unexpected path")
cm = _mock_remote_db_async(handler)
db = await cm.__aenter__()
table = await db.open_table(_TABLE_NAME)
assert isinstance(table, AsyncTable)
log.start()
return table, log, cm
def test_no_public_generated_column_status_resource_exported():
"""Baseline: no public status class/enum/resource is exported."""
for mod in (lancedb, lancedb.table, _native):
for name in _FORBIDDEN_PUBLIC_NAMES:
assert not hasattr(mod, name), f"{mod.__name__}.{name} must not be public"
def test_public_surface_signatures_annotations_and_hidden_bridge():
"""Four public methods + hidden native bridge must exist with frozen shape."""
assert hasattr(_native.Table, "_generated_column_status"), (
"native private bridge Table._generated_column_status is missing"
)
assert hasattr(Table, "generated_column_status"), (
"Table.generated_column_status is missing"
)
assert hasattr(LanceTable, "generated_column_status"), (
"LanceTable.generated_column_status is missing"
)
assert hasattr(RemoteTable, "generated_column_status"), (
"RemoteTable.generated_column_status is missing"
)
assert hasattr(AsyncTable, "generated_column_status"), (
"AsyncTable.generated_column_status is missing"
)
bridge = _native.Table._generated_column_status
_assert_exact_public_signature(bridge)
for keyword in _FORBIDDEN_BRIDGE_KWARGS:
assert keyword not in inspect.signature(bridge).parameters
for method in (
Table.generated_column_status,
LanceTable.generated_column_status,
RemoteTable.generated_column_status,
):
_assert_exact_public_signature(method)
assert not inspect.iscoroutinefunction(method)
assert get_type_hints(method)["return"] == _EXPECTED_RETURN
async_method = AsyncTable.generated_column_status
_assert_exact_public_signature(async_method)
assert inspect.iscoroutinefunction(async_method)
assert get_type_hints(async_method)["return"] == _EXPECTED_RETURN
def test_sync_remote_complete_and_incomplete_one_describe_each():
table, log, cm = _open_remote_table()
try:
complete = table.generated_column_status("complete_col")
_assert_status_string(complete, "complete")
_assert_one_status_describe(log)
log.start()
incomplete = table.generated_column_status("incomplete_col")
_assert_status_string(incomplete, "incomplete")
_assert_one_status_describe(log)
finally:
cm.__exit__(None, None, None)
@pytest.mark.asyncio
async def test_async_remote_complete_and_incomplete_one_describe_each():
table, log, cm = await _open_remote_table_async()
try:
complete = await table.generated_column_status("complete_col")
_assert_status_string(complete, "complete")
_assert_one_status_describe(log)
log.start()
incomplete = await table.generated_column_status("incomplete_col")
_assert_status_string(incomplete, "incomplete")
_assert_one_status_describe(log)
finally:
await cm.__aexit__(None, None, None)
@pytest.mark.parametrize(
("column_name", "status_describe", "expected_exc"),
[
(
"missing",
_describe_body(),
ValueError,
),
(
"Complete_Col",
_describe_body(),
ValueError,
),
(
"ordinary",
_describe_body(),
ValueError,
),
(
"complete_col",
_describe_body(
schema=_status_schema_fields(
complete_meta=_definition_metadata_json(
_COMPLETE_FIELD_ID + 1, 3, 3
)
)
),
ValueError,
),
(
"gen_bad",
_describe_body(
field_ids=[*_STABLE_FIELD_IDS, 9],
schema=_status_schema_fields(
bad_name="gen_bad",
bad_meta=(
'{"format_version":1,"output_field_id":9,'
f'"function_call":{_RAW_METADATA_MARKER},'
'"dependency_epoch":1,"materialized_epoch":1}'
),
),
),
ValueError,
),
(
"complete_col",
_describe_body(
schema=_status_schema_fields(
complete_meta=_definition_metadata_json(
_COMPLETE_FIELD_ID, 1, 1
).replace('"format_version":1', '"format_version":2')
)
),
ValueError,
),
(
"incomplete_col",
_describe_body(
schema=_status_schema_fields(
incomplete_meta=_definition_metadata_json(
_INCOMPLETE_FIELD_ID, 1, 2
)
)
),
ValueError,
),
(
"complete_col",
_describe_body(field_ids=None),
NotImplementedError,
),
],
ids=[
"missing",
"case_mismatch",
"ordinary",
"output_id_mismatch",
"malformed_metadata",
"unknown_format_version",
"reversed_epochs",
"old_server_missing_field_ids",
],
)
def test_remote_fail_closed_matrix_one_describe(
column_name: str,
status_describe: dict[str, Any],
expected_exc: type[BaseException],
):
table, log, cm = _open_remote_table(status_describe=status_describe)
try:
with pytest.raises(expected_exc) as raised:
table.generated_column_status(column_name)
text = _exception_text(raised.value)
assert _RAW_METADATA_MARKER not in text
_assert_one_status_describe(log)
finally:
cm.__exit__(None, None, None)
def test_sync_empty_name_zero_post_open_requests():
table, log, cm = _open_remote_table()
try:
with pytest.raises(ValueError):
table.generated_column_status("")
_assert_no_operation_traffic(log)
finally:
cm.__exit__(None, None, None)
@pytest.mark.asyncio
async def test_async_empty_name_zero_post_open_requests():
table, log, cm = await _open_remote_table_async()
try:
with pytest.raises(ValueError):
await table.generated_column_status("")
_assert_no_operation_traffic(log)
finally:
await cm.__aexit__(None, None, None)
@pytest.mark.asyncio
async def test_async_closed_status_empty_validation_wins_and_nonempty_closed():
"""Publicly closed AsyncTable: empty validates first; nonempty is closed."""
table, log, cm = await _open_remote_table_async()
try:
table.close()
log.start()
try:
await table.generated_column_status("complete_col")
except AttributeError:
raise
except Exception as exc:
text = _exception_text(exc)
assert "closed" in text.lower()
else:
pytest.fail("closed AsyncTable must fail before transport")
_assert_no_operation_traffic(log)
log.start()
with pytest.raises(ValueError) as raised:
await table.generated_column_status("")
text = _exception_text(raised.value)
assert "closed" not in text.lower()
_assert_no_operation_traffic(log)
finally:
await cm.__aexit__(None, None, None)
def test_local_sync_ordinary_column_fails_without_side_effects(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table(
"ordinary_only",
[{"ordinary": "alpha", "id": 1}, {"ordinary": "beta", "id": 2}],
)
assert isinstance(table, LanceTable)
version_before = table.version
schema_before = table.schema
data_before = table.to_arrow()
with pytest.raises(ValueError):
table.generated_column_status("ordinary")
assert table.version == version_before
assert table.schema == schema_before
assert table.to_arrow().equals(data_before)
@pytest.mark.asyncio
async def test_local_async_ordinary_column_fails_without_side_effects(tmp_path):
db = await lancedb.connect_async(tmp_path)
table = await db.create_table(
"ordinary_only_async",
[{"ordinary": "alpha", "id": 1}, {"ordinary": "beta", "id": 2}],
)
assert isinstance(table, AsyncTable)
version_before = await table.version()
schema_before = await table.schema()
data_before = await table.to_arrow()
with pytest.raises(ValueError):
await table.generated_column_status("ordinary")
assert await table.version() == version_before
assert await table.schema() == schema_before
assert (await table.to_arrow()).equals(data_before)
+291
View File
@@ -0,0 +1,291 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for the local @udf declaration surface."""
from __future__ import annotations
import importlib
import inspect
import types
import pyarrow as pa
import pytest
import lancedb
from lancedb import Function, Job, udf
from lancedb._udf import _get_udf_config
_REMOVED_AUTHORING_KNOBS = (
"user_version",
"idempotency_key",
"deterministic",
"null_handling",
"null_policy",
"on_error",
"error_policy",
"FunctionVersion",
"artifact",
"digest",
"geneva",
)
def _decorate(fn, **overrides):
kwargs = {
"inputs": {"x": pa.int32()},
"output": pa.int64(),
"python": "3.12",
}
kwargs.update(overrides)
return udf(**kwargs)(fn)
def test_udf_top_level_export_and_identity_metadata_behavior():
assert "udf" in lancedb.__all__
assert udf is lancedb.udf
assert isinstance(importlib.import_module("lancedb._udf"), types.ModuleType)
assert not isinstance(lancedb.udf, types.ModuleType)
def add(x, y=1):
"""Add locally."""
return x + y
original = add
decorated = _decorate(
add,
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
)
assert decorated is original
assert decorated.__name__ == "add"
assert decorated.__doc__ == "Add locally."
assert str(inspect.signature(decorated)) == "(x, y=1)"
assert decorated(2) == 3
assert decorated(2, 5) == 7
assert decorated(x=4, y=6) == 10
def test_udf_config_snapshot_order_defaults_and_immutability():
inputs = {"z": pa.string(), "a": pa.int32()}
packages = ["pkg-b==2", "pkg-a==1"]
def combine(z, a):
return f"{z}:{a}"
decorated = udf(
inputs=inputs,
output=pa.string(),
python="3.11",
packages=packages,
output_nullable=False,
)(combine)
config = _get_udf_config(decorated)
assert config.inputs == (("z", pa.string()), ("a", pa.int32()))
assert isinstance(config.inputs, tuple)
assert config.output == pa.string()
assert config.output_nullable is False
assert config.python == "3.11"
assert config.packages == ("pkg-b==2", "pkg-a==1")
assert isinstance(config.packages, tuple)
inputs["extra"] = pa.bool_()
del inputs["z"]
packages.append("pkg-c==3")
packages[0] = "mutated==0"
assert config.inputs == (("z", pa.string()), ("a", pa.int32()))
assert config.packages == ("pkg-b==2", "pkg-a==1")
for attr in ("inputs", "output", "output_nullable", "python", "packages"):
with pytest.raises(AttributeError):
setattr(config, attr, None)
def defaults_only(x):
return x
defaulted = udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
)(defaults_only)
default_config = _get_udf_config(defaulted)
assert default_config.packages == ()
assert default_config.output_nullable is True
def test_udf_accepts_lambda_and_closure_for_local_declaration():
ambient = "ambient-secret-value-xyz"
lam = udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
)(lambda n: n + 1)
assert lam(3) == 4
assert _get_udf_config(lam).inputs == (("n", pa.int32()),)
def factory(offset):
@udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
packages=["demo==0.1"],
)
def closed(n):
return n + offset + len(ambient)
return closed
closed = factory(10)
assert closed(2) == 12 + len(ambient)
assert _get_udf_config(closed).packages == ("demo==0.1",)
def test_udf_declaration_defers_signature_and_implementation_packaging():
"""Declaration must not validate callable signature or embed implementation."""
def local_add(left, right=1):
return left + right
decorated = udf(
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
python="3.12",
)(local_add)
assert decorated is local_add
assert str(inspect.signature(decorated)) == "(left, right=1)"
assert decorated(2) == 3
assert decorated(2, 5) == 7
config = _get_udf_config(decorated)
assert config.inputs == (("x", pa.int32()), ("y", pa.int32()))
for attr in (
"source",
"module",
"callable",
"function",
"implementation",
"bundle",
"artifact",
"digest",
):
assert not hasattr(config, attr)
def test_udf_lookup_double_decoration_and_non_function_target():
def plain(x):
return x
with pytest.raises((TypeError, ValueError)):
_get_udf_config(plain)
decorated = _decorate(plain)
with pytest.raises((TypeError, ValueError)):
_decorate(decorated)
with pytest.raises(TypeError):
udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)(object())
with pytest.raises(TypeError):
udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)(42)
def test_udf_config_validation_errors():
def target(x):
return x
with pytest.raises(TypeError):
udf({"x": pa.int32()}, pa.int32(), "3.12")(target)
with pytest.raises(TypeError):
_decorate(target, inputs=[("x", pa.int32())])
with pytest.raises(TypeError):
_decorate(target, inputs={1: pa.int32()})
with pytest.raises(ValueError):
_decorate(target, inputs={"": pa.int32()})
with pytest.raises(TypeError):
_decorate(target, inputs={"x": "int32"})
with pytest.raises(TypeError):
_decorate(target, output="int64")
with pytest.raises(TypeError):
_decorate(target, python=3.12)
with pytest.raises(ValueError):
_decorate(target, python="")
with pytest.raises(TypeError):
_decorate(target, packages="pkg==1")
with pytest.raises(ValueError):
_decorate(target, packages=["pkg==1", ""])
with pytest.raises(ValueError):
_decorate(target, packages=["pkg==1", "pkg==1"])
with pytest.raises(TypeError):
_decorate(target, packages=["pkg==1", 2])
with pytest.raises(TypeError):
_decorate(target, output_nullable=1)
with pytest.raises(TypeError):
_decorate(target, output_nullable="true")
def test_udf_rejects_removed_overdesign_and_has_no_durable_side_effects():
params = inspect.signature(udf).parameters
for name in _REMOVED_AUTHORING_KNOBS:
assert name not in params
def score(x):
"""score body marker unique-xyz."""
ambient = "ambient-secret-value-xyz"
return f"{ambient}:{x}"
decorated = _decorate(
score,
packages=["score==1.0"],
output_nullable=True,
)
config = _get_udf_config(decorated)
text = repr(config).lower()
assert "score body marker unique-xyz" not in text
assert "ambient-secret-value-xyz" not in text
for token in (
"user_version",
"idempotency_key",
"deterministic",
"null_policy",
"on_error",
"functionversion",
"artifact",
"digest",
"geneva",
):
assert token not in text
for attr in _REMOVED_AUTHORING_KNOBS:
assert not hasattr(config, attr)
assert not isinstance(decorated, Function)
assert not isinstance(decorated, Job)
for attr in ("id", "function_id", "job", "job_id", "registration"):
assert not hasattr(decorated, attr)
@@ -0,0 +1,490 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for local FunctionCapability authoring and @udf capabilities."""
from __future__ import annotations
import inspect
import pyarrow as pa
import pytest
import lancedb
from lancedb import Function, FunctionCapability, Job, udf
from lancedb._udf import _get_udf_config, _package_udf
_SECRET_REFERENCE = "secret://team/capability-redact-token-xyz"
_SECRET_ENV = "API_TOKEN"
_NETWORK_ORIGIN = "https://api.example.com"
_OVERDESIGN_ATTRS = (
"id",
"function_id",
"FunctionVersion",
"function_version",
"artifact",
"digest",
"authorization",
"authorized",
"value",
"plaintext",
"plaintext_secret",
"secret_value",
"job",
"job_id",
"catalog",
"retry_key",
"user_version",
"idempotency_key",
"deterministic",
"null_handling",
"null_policy",
"on_error",
"error_policy",
"geneva",
)
def _decorate(fn, **overrides):
kwargs = {
"inputs": {"x": pa.int32()},
"output": pa.int64(),
"python": "3.12",
}
kwargs.update(overrides)
return udf(**kwargs)(fn)
def _network(origin: str = _NETWORK_ORIGIN) -> FunctionCapability:
return FunctionCapability.network(origin)
def _secret(
reference: str = _SECRET_REFERENCE,
*,
environment_variable: str = _SECRET_ENV,
) -> FunctionCapability:
return FunctionCapability.secret(
reference,
environment_variable=environment_variable,
)
@udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=False,
)
def packable_without_capabilities(x):
return x + 1
@udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=False,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
],
)
def packable_with_capabilities(x):
return x + 1
def test_function_capability_export_factories_projection_equality_immutability():
assert "FunctionCapability" in lancedb.__all__
assert FunctionCapability is lancedb.FunctionCapability
network = _network()
secret = _secret()
assert network.kind == "network"
assert network.origin == _NETWORK_ORIGIN
assert network.reference is None
assert network.environment_variable is None
assert secret.kind == "secret"
assert secret.reference == _SECRET_REFERENCE
assert secret.environment_variable == _SECRET_ENV
assert secret.origin is None
assert network == FunctionCapability.network(_NETWORK_ORIGIN)
assert secret == FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
)
assert network != secret
assert network != FunctionCapability.network("https://other.example.com")
assert secret != FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable="OTHER_TOKEN",
)
public_attrs = ("kind", "origin", "reference", "environment_variable")
internal_slots = ("_kind", "_origin", "_reference", "_environment_variable")
immutable_attrs = public_attrs + internal_slots
for attr in public_attrs:
with pytest.raises(AttributeError):
setattr(network, attr, None)
with pytest.raises(AttributeError):
setattr(secret, attr, None)
for attr in immutable_attrs:
# Fresh instances per attempt so a RED slot mutation cannot corrupt
# shared fixtures used by later assertions in this test.
fresh_network = _network("https://fresh-immutability.example.com")
fresh_secret = _secret(
"secret://team/fresh-immutability-token",
environment_variable="FRESH_IMMUTABILITY_TOKEN",
)
with pytest.raises(AttributeError):
setattr(fresh_network, attr, None)
with pytest.raises(AttributeError):
setattr(fresh_secret, attr, None)
with pytest.raises(AttributeError):
delattr(fresh_network, attr)
with pytest.raises(AttributeError):
delattr(fresh_secret, attr)
retained_origin = "https://config-retain.example.com"
retained_reference = "secret://team/config-retain-token"
retained_env = "CONFIG_RETAIN_TOKEN"
retained_network = FunctionCapability.network(retained_origin)
retained_secret = FunctionCapability.secret(
retained_reference,
environment_variable=retained_env,
)
expected_capabilities = (
FunctionCapability.network(retained_origin),
FunctionCapability.secret(
retained_reference,
environment_variable=retained_env,
),
)
def retain_target(x):
return x
retained = udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
capabilities=[retained_network, retained_secret],
)(retain_target)
retained_config = _get_udf_config(retained)
assert retained_config.capabilities == expected_capabilities
for attr in immutable_attrs:
with pytest.raises(AttributeError):
setattr(retained_network, attr, "mutated")
with pytest.raises(AttributeError):
setattr(retained_secret, attr, "mutated")
with pytest.raises(AttributeError):
delattr(retained_network, attr)
with pytest.raises(AttributeError):
delattr(retained_secret, attr)
assert retained_config.capabilities == expected_capabilities
assert retained_config.capabilities[0] is retained_network
assert retained_config.capabilities[1] is retained_secret
assert retained_config.capabilities[0].kind == "network"
assert retained_config.capabilities[0].origin == retained_origin
assert retained_config.capabilities[0].reference is None
assert retained_config.capabilities[0].environment_variable is None
assert retained_config.capabilities[1].kind == "secret"
assert retained_config.capabilities[1].reference == retained_reference
assert retained_config.capabilities[1].environment_variable == retained_env
assert retained_config.capabilities[1].origin is None
with pytest.raises(TypeError):
FunctionCapability()
with pytest.raises(TypeError):
FunctionCapability( # type: ignore[call-arg]
kind="network",
origin=_NETWORK_ORIGIN,
)
assert not isinstance(network, Function)
assert not isinstance(secret, Function)
assert not isinstance(network, Job)
assert not isinstance(secret, Job)
for attr in _OVERDESIGN_ATTRS:
assert not hasattr(network, attr)
assert not hasattr(secret, attr)
def test_function_capability_validation_and_secret_redaction():
with pytest.raises(TypeError):
FunctionCapability.network(None) # type: ignore[arg-type]
with pytest.raises(TypeError):
FunctionCapability.network(123) # type: ignore[arg-type]
with pytest.raises(ValueError):
FunctionCapability.network("")
# Backend authorization owns URL/scheme policy; non-empty is enough here.
loose = FunctionCapability.network("example.com")
assert loose.kind == "network"
assert loose.origin == "example.com"
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[misc]
_SECRET_REFERENCE,
_SECRET_ENV,
)
with pytest.raises(TypeError):
FunctionCapability.secret(None, environment_variable=_SECRET_ENV) # type: ignore[arg-type]
with pytest.raises(TypeError):
FunctionCapability.secret(123, environment_variable=_SECRET_ENV) # type: ignore[arg-type]
with pytest.raises(TypeError):
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable=None) # type: ignore[arg-type]
with pytest.raises(TypeError):
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable=1) # type: ignore[arg-type]
with pytest.raises(ValueError) as empty_ref:
FunctionCapability.secret("", environment_variable=_SECRET_ENV)
assert _SECRET_REFERENCE not in str(empty_ref.value)
assert _SECRET_REFERENCE not in repr(empty_ref.value)
with pytest.raises(ValueError) as empty_env:
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable="")
assert _SECRET_REFERENCE not in str(empty_env.value)
assert _SECRET_REFERENCE not in repr(empty_env.value)
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[call-arg]
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
value="super-secret",
)
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[call-arg]
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
plaintext_secret="super-secret",
)
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[call-arg]
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
environment={_SECRET_ENV: "super-secret"},
)
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[call-arg]
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
headers={"Authorization": "Bearer super-secret"},
)
with pytest.raises(TypeError):
FunctionCapability.network( # type: ignore[call-arg]
_NETWORK_ORIGIN,
headers={"X-Trace": "1"},
)
secret = _secret()
assert not hasattr(secret, "value")
assert not hasattr(secret, "plaintext")
assert not hasattr(secret, "plaintext_secret")
assert not hasattr(secret, "secret_value")
secret_text = repr(secret)
assert "secret" in secret_text.lower()
assert _SECRET_ENV in secret_text
assert _SECRET_REFERENCE not in secret_text
assert "super-secret" not in secret_text
network_text = repr(_network())
assert "network" in network_text.lower()
assert _NETWORK_ORIGIN in network_text
def test_udf_capabilities_ordered_immutable_config_default_and_validation():
params = inspect.signature(udf).parameters
assert "capabilities" in params
assert params["capabilities"].kind is inspect.Parameter.KEYWORD_ONLY
assert params["capabilities"].default == ()
def identity_target(x):
"""capabilities identity marker."""
return x + 1
original = identity_target
decorated = _decorate(identity_target)
assert decorated is original
assert decorated.__name__ == "identity_target"
assert decorated.__doc__ == "capabilities identity marker."
assert decorated(2) == 3
assert _get_udf_config(decorated).capabilities == ()
first = _network("https://b.example.com")
second = _network("https://a.example.com")
third = _network("https://b.example.com")
secret = _secret()
capabilities = [first, second, third, secret]
def combine(x):
return x
with_caps = udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
packages=["pkg-b==2", "pkg-a==1"],
capabilities=capabilities,
)(combine)
config = _get_udf_config(with_caps)
assert config.capabilities == (first, second, third, secret)
assert isinstance(config.capabilities, tuple)
assert config.packages == ("pkg-b==2", "pkg-a==1")
assert config.inputs == (("x", pa.int32()),)
capabilities.append(_network("https://mutated.example.com"))
capabilities[0] = _network("https://replaced.example.com")
assert config.capabilities == (first, second, third, secret)
with pytest.raises(AttributeError):
setattr(config, "capabilities", ())
def target(x):
return x
with pytest.raises(TypeError):
_decorate(target, capabilities="https://api.example.com")
with pytest.raises(TypeError):
_decorate(target, capabilities=b"https://api.example.com")
class _BadCapability:
def __repr__(self) -> str:
return "unique-bad-capability-repr-xyz"
with pytest.raises(TypeError) as bad_item:
_decorate(target, capabilities=[_BadCapability()])
assert "unique-bad-capability-repr-xyz" not in str(bad_item.value)
assert "unique-bad-capability-repr-xyz" not in repr(bad_item.value)
with pytest.raises(TypeError) as bad_mixed:
_decorate(
target,
capabilities=[_network(), "unique-bad-capability-string-xyz"],
)
assert "unique-bad-capability-string-xyz" not in str(bad_mixed.value)
assert "unique-bad-capability-string-xyz" not in repr(bad_mixed.value)
def test_udf_capabilities_rejects_function_capability_subclass_before_property_access():
marker = "unique-hostile-capability-subclass-marker-xyz"
class _HostileFunctionCapability(FunctionCapability):
@property
def kind(self) -> str:
raise RuntimeError(marker)
@property
def origin(self) -> str | None:
raise RuntimeError(marker)
@property
def reference(self) -> str | None:
raise RuntimeError(marker)
@property
def environment_variable(self) -> str | None:
raise RuntimeError(marker)
hostile = object.__new__(_HostileFunctionCapability)
assert isinstance(hostile, FunctionCapability)
assert type(hostile) is not FunctionCapability
def target(x):
return x
with pytest.raises(TypeError) as exc_info:
_decorate(target, capabilities=[hostile])
assert marker not in str(exc_info.value)
assert marker not in repr(exc_info.value)
assert _SECRET_REFERENCE not in str(exc_info.value)
assert _SECRET_REFERENCE not in repr(exc_info.value)
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
def test_package_udf_preserves_capabilities_and_redacts_secret_reference():
packaged = _package_udf(packable_with_capabilities)
config = packaged.config
assert packaged.config is _get_udf_config(packable_with_capabilities)
assert config.capabilities == (
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
)
assert config.capabilities[0].kind == "network"
assert config.capabilities[0].origin == _NETWORK_ORIGIN
assert config.capabilities[1].kind == "secret"
assert config.capabilities[1].reference == _SECRET_REFERENCE
assert config.capabilities[1].environment_variable == _SECRET_ENV
assert config.packages == ("pkg-a==1",)
assert config.python == "3.12"
assert config.output_nullable is False
nested = (
f"{packaged!r}\n{config!r}\n{config.capabilities!r}\n{config.capabilities[1]!r}"
)
assert _SECRET_REFERENCE not in nested
assert _SECRET_ENV in repr(config.capabilities[1])
def test_capabilities_are_additive_to_existing_declaration_and_packaging():
def score(x):
return x
decorated = _decorate(
score,
packages=["score==1.0"],
output_nullable=True,
)
config = _get_udf_config(decorated)
assert config.inputs == (("x", pa.int32()),)
assert config.output == pa.int64()
assert config.output_nullable is True
assert config.python == "3.12"
assert config.packages == ("score==1.0",)
assert config.capabilities == ()
assert decorated is score
assert decorated(4) == 4
packaged = _package_udf(packable_without_capabilities)
assert packaged.config is _get_udf_config(packable_without_capabilities)
assert packaged.callable_name == "packable_without_capabilities"
assert packaged.config.capabilities == ()
assert packaged.config.packages == ("pkg-a==1",)
assert packaged.config.output_nullable is False
assert packable_without_capabilities(1) == 2
params = inspect.signature(udf).parameters
for name in (
"user_version",
"idempotency_key",
"deterministic",
"null_handling",
"null_policy",
"on_error",
"error_policy",
"FunctionVersion",
"artifact",
"digest",
"geneva",
):
assert name not in params
@@ -0,0 +1,506 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for the private UDF -> FunctionDefinition bridge."""
from __future__ import annotations
import base64
import io
import json
from pathlib import Path
import pyarrow as pa
import pytest
import lancedb
from lancedb import FunctionCapability, udf
from lancedb import _lancedb as _native
from lancedb import _udf as _udf_mod
_SOURCE_MARKER = "bridge-source-marker-unique-xyz"
_SECRET_REFERENCE = "secret://team/bridge-redact-token-xyz"
_SECRET_ENV = "BRIDGE_API_TOKEN"
_NETWORK_ORIGIN = "https://api.bridge-example.com"
_NETWORK_ORIGIN_B = "https://other.bridge-example.com"
_FORBIDDEN_WIRE_KEYS = (
"id",
"function_id",
"FunctionId",
"catalog",
"catalog_name",
"version",
"function_version",
"FunctionVersion",
"lineage",
"user_version",
"idempotency_key",
"digest",
"artifact",
"artifact_digest",
"storage",
"storage_location",
"location",
"deterministic",
"null_policy",
"nullPolicy",
"timestamp",
"created_at",
"updated_at",
"worker",
"scheduler",
"attempt",
"attempt_id",
"replica",
"placement",
"job",
"job_id",
"retry_key",
"registration",
)
_OVERDESIGN_ATTRS = (
"id",
"function_id",
"FunctionVersion",
"function_version",
"artifact",
"digest",
"job",
"job_id",
"catalog",
"retry_key",
"user_version",
"idempotency_key",
"deterministic",
"null_policy",
"null_handling",
)
@udf(
inputs={"text": pa.string(), "limit": pa.int32()},
output=pa.string(),
python="3.12",
packages=["pkg-b==2", "pkg-a==1"],
output_nullable=True,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
FunctionCapability.network(_NETWORK_ORIGIN_B),
],
)
def packable_bridge_normalize(text, limit):
"""bridge-source-marker-unique-xyz."""
return text[:limit]
def _build_function_definition(fn: object):
return _udf_mod._build_function_definition(fn)
def _function_definition_type():
return _native._FunctionDefinition
def _new_function_definition(**kwargs):
return _native._new_function_definition(**kwargs)
def _json_bytes(definition) -> bytes:
payload = definition._to_json()
if isinstance(payload, bytes):
return payload
assert isinstance(payload, str)
return payload.encode("utf-8")
def _decode_type_ipc(encoded: str) -> pa.DataType:
raw = base64.b64decode(encoded)
reader = pa.ipc.open_file(io.BytesIO(raw))
assert reader.num_record_batches == 0
assert len(reader.schema) == 1
return reader.schema.field(0).type
def _assert_exact_object_keys(value: dict, expected: set[str], *, context: str) -> None:
assert isinstance(value, dict), f"{context} must be an object"
assert set(value) == expected, f"{context} keys must match exactly: {set(value)!r}"
def _assert_forbidden_keys_absent(value: object, *, context: str) -> None:
if isinstance(value, dict):
for key in value:
assert key not in _FORBIDDEN_WIRE_KEYS, (
f"forbidden key {key!r} at {context}: {value!r}"
)
if key == "name" and context in {
"definition",
"signature",
"signature.output",
"implementation",
}:
raise AssertionError(
f"catalog/function identity key `name` must be absent at {context}"
)
child_context = f"{context}.{key}"
if key == "parameters" and context == "signature":
child_context = "signature.parameters"
_assert_forbidden_keys_absent(value[key], context=child_context)
elif isinstance(value, list):
for idx, item in enumerate(value):
item_context = (
f"signature.parameters[{idx}]"
if context == "signature.parameters"
else f"{context}[{idx}]"
)
if context == "signature.parameters":
assert isinstance(item, dict)
assert "name" in item
for key in item:
assert key not in _FORBIDDEN_WIRE_KEYS
assert key != "catalog_name"
_assert_forbidden_keys_absent(
{k: v for k, v in item.items() if k != "name"},
context=item_context,
)
else:
_assert_forbidden_keys_absent(item, context=item_context)
def _assert_sanitized_text(*parts: object) -> None:
combined = "\n".join(str(part) for part in parts)
lowered = combined.lower()
assert _SOURCE_MARKER.lower() not in lowered
assert _SECRET_REFERENCE.lower() not in lowered
assert str(Path(__file__).resolve()).lower() not in lowered
assert Path(__file__).resolve().as_posix().lower() not in lowered
def _assert_clean_validation_error(exc_info) -> None:
_assert_sanitized_text(exc_info.value, repr(exc_info.value))
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
def _valid_builder_kwargs(**overrides):
kwargs = {
"parameters": [("text", pa.string()), ("limit", pa.int32())],
"output_type": pa.string(),
"output_nullable": True,
"module": "bridge_mod",
"callable_name": "normalize",
"source": (
"def normalize(text, limit):\n"
f" # {_SOURCE_MARKER}\n"
" return text[:limit]\n"
),
"python": "3.12",
"packages": ["pkg-b==2", "pkg-a==1"],
"capabilities": [
("network", _NETWORK_ORIGIN, None),
("secret", _SECRET_REFERENCE, _SECRET_ENV),
("network", _NETWORK_ORIGIN_B, None),
],
}
kwargs.update(overrides)
return kwargs
def test_build_function_definition_private_native_immutability_and_export_surface():
assert "_build_function_definition" not in getattr(lancedb, "__all__", [])
assert "_FunctionDefinition" not in lancedb.__all__
assert not hasattr(lancedb, "_FunctionDefinition")
assert not hasattr(lancedb, "_build_function_definition")
assert not hasattr(lancedb, "_new_function_definition")
definition = _build_function_definition(packable_bridge_normalize)
definition_type = _function_definition_type()
assert type(definition) is definition_type
assert definition_type.__module__ == "lancedb._lancedb"
assert definition_type.__name__ == "_FunctionDefinition"
with pytest.raises(TypeError):
definition_type()
for attr in _OVERDESIGN_ATTRS:
assert not hasattr(definition, attr)
for attr in ("signature", "module", "source", "capabilities"):
with pytest.raises(AttributeError):
setattr(definition, attr, None)
def test_build_function_definition_json_wire_ordered_contract_without_identity():
definition = _build_function_definition(packable_bridge_normalize)
encoded_a = _json_bytes(definition)
encoded_b = _json_bytes(definition)
assert encoded_a == encoded_b
wire = json.loads(encoded_a.decode("utf-8"))
_assert_exact_object_keys(
wire,
{"format_version", "signature", "implementation", "capabilities"},
context="definition",
)
assert wire["format_version"] == 1
_assert_forbidden_keys_absent(wire, context="definition")
signature = wire["signature"]
_assert_exact_object_keys(signature, {"parameters", "output"}, context="signature")
parameters = signature["parameters"]
assert [parameter["name"] for parameter in parameters] == ["text", "limit"]
for parameter in parameters:
_assert_exact_object_keys(
parameter, {"name", "data_type_ipc"}, context="parameter"
)
assert isinstance(parameter["data_type_ipc"], str)
assert parameter["data_type_ipc"]
assert _decode_type_ipc(parameters[0]["data_type_ipc"]) == pa.string()
assert _decode_type_ipc(parameters[1]["data_type_ipc"]) == pa.int32()
output = signature["output"]
_assert_exact_object_keys(
output, {"data_type_ipc", "nullable"}, context="signature.output"
)
assert output["nullable"] is True
assert _decode_type_ipc(output["data_type_ipc"]) == pa.string()
implementation = wire["implementation"]
_assert_exact_object_keys(
implementation,
{"kind", "module", "callable", "source", "python", "packages"},
context="implementation",
)
assert implementation["kind"] == "python"
assert implementation["module"] == __name__
assert implementation["callable"] == "packable_bridge_normalize"
assert implementation["source"] == Path(__file__).read_text(encoding="utf-8")
assert _SOURCE_MARKER in implementation["source"]
assert implementation["python"] == "3.12"
assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"]
capabilities = wire["capabilities"]
assert capabilities == [
{"kind": "network", "origin": _NETWORK_ORIGIN},
{
"kind": "secret",
"reference": _SECRET_REFERENCE,
"environment_variable": _SECRET_ENV,
},
{"kind": "network", "origin": _NETWORK_ORIGIN_B},
]
for capability in capabilities:
assert "value" not in capability
assert "plaintext" not in capability
assert "plaintext_secret" not in capability
assert "secret_value" not in capability
def test_native_definition_repr_includes_safe_structure_and_redacts_sensitive_text():
definition = _build_function_definition(packable_bridge_normalize)
rendered = repr(definition)
assert "_FunctionDefinition" in rendered or "FunctionDefinition" in rendered
assert __name__ in rendered
assert "packable_bridge_normalize" in rendered
assert "3.12" in rendered
_assert_sanitized_text(rendered)
def test_new_function_definition_builder_preserves_normalized_wire():
definition = _new_function_definition(**_valid_builder_kwargs())
assert type(definition) is _function_definition_type()
encoded_a = _json_bytes(definition)
encoded_b = _json_bytes(definition)
assert encoded_a == encoded_b
wire = json.loads(encoded_a.decode("utf-8"))
assert wire["format_version"] == 1
assert [parameter["name"] for parameter in wire["signature"]["parameters"]] == [
"text",
"limit",
]
assert _decode_type_ipc(wire["signature"]["parameters"][0]["data_type_ipc"]) == (
pa.string()
)
assert _decode_type_ipc(wire["signature"]["parameters"][1]["data_type_ipc"]) == (
pa.int32()
)
assert wire["signature"]["output"]["nullable"] is True
assert _decode_type_ipc(wire["signature"]["output"]["data_type_ipc"]) == pa.string()
implementation = wire["implementation"]
assert implementation["kind"] == "python"
assert implementation["module"] == "bridge_mod"
assert implementation["callable"] == "normalize"
assert implementation["source"] == _valid_builder_kwargs()["source"]
assert _SOURCE_MARKER in implementation["source"]
assert implementation["python"] == "3.12"
assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"]
assert wire["capabilities"] == [
{"kind": "network", "origin": _NETWORK_ORIGIN},
{
"kind": "secret",
"reference": _SECRET_REFERENCE,
"environment_variable": _SECRET_ENV,
},
{"kind": "network", "origin": _NETWORK_ORIGIN_B},
]
_assert_forbidden_keys_absent(wire, context="definition")
@pytest.mark.parametrize(
("overrides",),
[
({"parameters": [("text", pa.string()), ("text", pa.int32())]},),
({"parameters": [("", pa.string())]},),
({"module": ""},),
({"callable_name": ""},),
({"source": ""},),
({"python": ""},),
({"packages": ["pkg-a==1", ""]},),
({"packages": ["pkg-a==1", "pkg-a==1"]},),
({"capabilities": [("filesystem", _NETWORK_ORIGIN, None)]},),
({"capabilities": [("network", _NETWORK_ORIGIN, _SECRET_ENV)]},),
({"capabilities": [("secret", _SECRET_REFERENCE, None)]},),
({"capabilities": [("secret", _SECRET_REFERENCE, "")]},),
({"capabilities": [("network", "", None)]},),
({"capabilities": [("secret", "", _SECRET_ENV)]},),
],
)
def test_new_function_definition_strict_validation_rejections(overrides):
kwargs = _valid_builder_kwargs(**overrides)
with pytest.raises(ValueError) as exc_info:
_new_function_definition(**kwargs)
_assert_clean_validation_error(exc_info)
def test_new_function_definition_validation_does_not_echo_secret_or_source_marker():
with pytest.raises(ValueError) as exc_info:
_new_function_definition(**_valid_builder_kwargs(module=""))
_assert_clean_validation_error(exc_info)
with pytest.raises(ValueError) as exc_info:
_new_function_definition(
**_valid_builder_kwargs(packages=["pkg-a==1", "pkg-a==1"])
)
_assert_clean_validation_error(exc_info)
with pytest.raises(ValueError) as exc_info:
_new_function_definition(
**_valid_builder_kwargs(
capabilities=[("secret", _SECRET_REFERENCE, None)],
)
)
_assert_clean_validation_error(exc_info)
@pytest.mark.parametrize(
("overrides",),
[
({"parameters": [("text", "not-a-datatype")]},),
({"parameters": [(123, pa.string())]},),
({"output_type": "not-a-datatype"},),
({"output_type": None},),
({"output_nullable": "yes"},),
({"packages": "pkg-a==1"},),
({"capabilities": "network"},),
({"capabilities": [("network", _NETWORK_ORIGIN)]},),
({"capabilities": [("network", _NETWORK_ORIGIN, None, "extra")]},),
],
)
def test_new_function_definition_wrong_pyarrow_and_shape_values_fail_closed(overrides):
kwargs = _valid_builder_kwargs(**overrides)
with pytest.raises((TypeError, ValueError)) as exc_info:
_new_function_definition(**kwargs)
_assert_clean_validation_error(exc_info)
class _HostileRaisingIterable:
def __iter__(self):
raise RuntimeError(f"{_SECRET_REFERENCE} {_SOURCE_MARKER}")
@pytest.mark.parametrize(
("overrides",),
[
({"parameters": _HostileRaisingIterable()},),
({"packages": _HostileRaisingIterable()},),
({"capabilities": _HostileRaisingIterable()},),
(
{
"capabilities": [
("network", _NETWORK_ORIGIN, None),
_HostileRaisingIterable(),
("network", _NETWORK_ORIGIN_B, None),
]
},
),
],
)
def test_new_function_definition_hostile_iterable_iter_raises_fail_closed(overrides):
kwargs = _valid_builder_kwargs(**overrides)
with pytest.raises((TypeError, ValueError)) as exc_info:
_new_function_definition(**kwargs)
_assert_clean_validation_error(exc_info)
@udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=False,
)
def packable_bridge_capability_exact_type(x):
return x + 1
def test_build_function_definition_rejects_forged_function_capability_subclass():
marker = f"{_SECRET_REFERENCE} {_SOURCE_MARKER}"
class _HostileFunctionCapability(FunctionCapability):
@property
def kind(self) -> str:
raise RuntimeError(marker)
@property
def origin(self) -> str | None:
raise RuntimeError(marker)
@property
def reference(self) -> str | None:
raise RuntimeError(marker)
@property
def environment_variable(self) -> str | None:
raise RuntimeError(marker)
hostile = object.__new__(_HostileFunctionCapability)
assert isinstance(hostile, FunctionCapability)
assert type(hostile) is not FunctionCapability
config_attr = _udf_mod._CONFIG_ATTR
original = getattr(packable_bridge_capability_exact_type, config_attr)
forged = _udf_mod._UdfConfig(
inputs=original.inputs,
output=original.output,
output_nullable=original.output_nullable,
python=original.python,
packages=original.packages,
capabilities=(hostile,),
)
setattr(packable_bridge_capability_exact_type, config_attr, forged)
try:
with pytest.raises((TypeError, ValueError)) as exc_info:
_build_function_definition(packable_bridge_capability_exact_type)
_assert_clean_validation_error(exc_info)
assert marker not in str(exc_info.value)
assert marker not in repr(exc_info.value)
finally:
setattr(packable_bridge_capability_exact_type, config_attr, original)
@@ -0,0 +1,486 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for private UDF packaging validation."""
from __future__ import annotations
import importlib
import inspect
import json
import sys
from collections.abc import Iterator
from contextlib import contextmanager
from pathlib import Path
import pyarrow as pa
import pytest
from lancedb import Function, Job, udf
from lancedb._udf import _get_udf_config, _package_udf
_BODY_MARKER = "packaging body marker unique-xyz"
_AMBIENT_SECRET = "ambient-secret-value-xyz"
_BUILTIN_SHADOW_SECRET = "builtin-shadow-secret-xyz"
_SOURCE_MISMATCH_SECRET = "source-mismatch-secret-xyz"
_INVALID_UTF8_SECRET = "invalid-utf8-secret-xyz"
_OVERDESIGN_ATTRS = (
"user_version",
"idempotency_key",
"deterministic",
"null_handling",
"null_policy",
"on_error",
"error_policy",
"FunctionVersion",
"function_version",
"artifact",
"digest",
"geneva",
"id",
"function_id",
"job",
"job_id",
"registration",
"catalog",
"retry_key",
"source_path",
"path",
"function",
)
_PACKAGING_CONSTANT = 41
def _packaging_helper(value: int) -> int:
return value + _PACKAGING_CONSTANT
@udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=False,
)
def packable_add(x):
"""packaging body marker unique-xyz."""
return _packaging_helper(x) + len(json.dumps({"k": 1}))
@udf(
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
python="3.12",
)
def packable_kwonly(x, *, y=2):
return x + y
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def packable_rebind_target(x):
return x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def uses_injected_ambient(x):
return x + len(INJECTED_AMBIENT_GLOBAL) # noqa: F821
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def uses_shadowed_builtin_len(x):
return x + len((1, 2, 3))
@udf(
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
python="3.12",
)
def mismatch_names(left, right):
return left + right
@udf(
inputs={"y": pa.int32(), "x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def mismatch_order(x, y):
return x + y
@udf(
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
python="3.12",
)
def positional_only(x, /, y):
return x + y
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def varargs_fn(x, *args):
return x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def kwargs_fn(x, **kwargs):
return x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
async def async_fn(x):
return x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
async def async_gen_fn(x):
yield x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def generator_fn(x):
yield x
def _assert_sanitized_text(*parts: object, secret: str = _AMBIENT_SECRET) -> None:
combined = "\n".join(str(part) for part in parts)
lowered = combined.lower()
assert _BODY_MARKER.lower() not in lowered
assert secret.lower() not in lowered
assert str(Path(__file__).resolve()).lower() not in lowered
assert Path(__file__).resolve().as_posix().lower() not in lowered
def _assert_packaging_rejection(exc_info, *, secret: str = _AMBIENT_SECRET) -> None:
_assert_sanitized_text(exc_info.value, repr(exc_info.value), secret=secret)
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
@contextmanager
def _temporary_imported_module(
directory: Path, module_name: str, source: str
) -> Iterator[tuple[Path, object]]:
path = directory / f"{module_name}.py"
path.write_text(source, encoding="utf-8")
inserted = str(directory)
sys.path.insert(0, inserted)
try:
sys.modules.pop(module_name, None)
module = importlib.import_module(module_name)
yield path, module
finally:
sys.modules.pop(module_name, None)
try:
sys.path.remove(inserted)
except ValueError:
pass
def _temp_udf_module_source(*, body: str, secret: str | None = None) -> str:
secret_line = f"_SECRET = {secret!r}\n" if secret is not None else ""
return (
"import pyarrow as pa\n"
"from lancedb import udf\n"
f"{secret_line}\n"
"@udf(\n"
' inputs={"x": pa.int32()},\n'
" output=pa.int32(),\n"
' python="3.12",\n'
")\n"
"def temp_pack_target(x):\n"
f" {body}\n"
)
def test_package_udf_success_snapshot_source_module_callable_config_and_repr():
packaged = _package_udf(packable_add)
source = Path(__file__).read_text(encoding="utf-8")
assert packaged.source == source
assert packaged.module == __name__
assert packaged.module != "__main__"
assert packaged.callable_name == "packable_add"
assert packable_add.__qualname__ == "packable_add"
assert packaged.config is _get_udf_config(packable_add)
assert packaged.config.inputs == (("x", pa.int32()),)
assert packaged.config.output == pa.int64()
assert packaged.config.output_nullable is False
assert packaged.config.python == "3.12"
assert packaged.config.packages == ("pkg-a==1",)
for attr in ("source", "module", "callable_name", "config"):
with pytest.raises(AttributeError):
setattr(packaged, attr, None)
text = repr(packaged)
_assert_sanitized_text(text)
assert _BODY_MARKER not in text
def test_package_udf_allows_source_bound_import_constant_and_helper():
packaged = _package_udf(packable_add)
assert packaged.callable_name == "packable_add"
assert "import json" in packaged.source
assert "_PACKAGING_CONSTANT" in packaged.source
assert "_packaging_helper" in packaged.source
assert packable_add(1) == _packaging_helper(1) + len(json.dumps({"k": 1}))
def test_package_udf_accepts_positional_or_keyword_and_keyword_only_defaults():
packaged = _package_udf(packable_kwonly)
assert packaged.callable_name == "packable_kwonly"
assert packaged.config.inputs == (("x", pa.int32()), ("y", pa.int32()))
assert str(inspect.signature(packable_kwonly)) == "(x, *, y=2)"
assert packable_kwonly(3) == 5
assert packable_kwonly(3, y=7) == 10
def test_package_udf_rejects_lambda_and_closure():
lam = udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
)(lambda n: n + 1)
with pytest.raises(ValueError) as exc_info:
_package_udf(lam)
_assert_packaging_rejection(exc_info)
ambient = _AMBIENT_SECRET
def factory(offset):
@udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
)
def closed(n):
return n + offset + len(ambient)
return closed
closed = factory(10)
with pytest.raises(ValueError) as exc_info:
_package_udf(closed)
_assert_packaging_rejection(exc_info)
def outer():
total = 0
@udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
)
def nested(n):
nonlocal total
total += n
return total
return nested
with pytest.raises(ValueError) as exc_info:
_package_udf(outer())
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_signature_mismatches_and_unsupported_parameter_kinds():
for target in (
mismatch_names,
mismatch_order,
positional_only,
varargs_fn,
kwargs_fn,
):
with pytest.raises(ValueError) as exc_info:
_package_udf(target)
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_async_and_generator_functions():
for target in (async_fn, async_gen_fn, generator_fn):
with pytest.raises(ValueError) as exc_info:
_package_udf(target)
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_dynamic_exec_source():
namespace: dict[str, object] = {}
exec(
"def dynamic_pack_target(x):\n return x + 1\n",
namespace,
)
dynamic = udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)(namespace["dynamic_pack_target"])
with pytest.raises(ValueError) as exc_info:
_package_udf(dynamic)
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_undecorated_and_wrong_input_types():
def plain(x):
return x
with pytest.raises(TypeError) as exc_info:
_package_udf(plain)
_assert_packaging_rejection(exc_info)
with pytest.raises(TypeError) as exc_info:
_package_udf(object())
_assert_packaging_rejection(exc_info)
with pytest.raises(TypeError) as exc_info:
_package_udf(42)
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_rebound_module_attribute():
module = sys.modules[__name__]
original = module.packable_rebind_target
replacement = udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)(lambda x: x)
module.packable_rebind_target = replacement
try:
with pytest.raises(ValueError) as exc_info:
_package_udf(original)
_assert_packaging_rejection(exc_info)
finally:
module.packable_rebind_target = original
def test_package_udf_rejects_injected_ambient_global():
module = sys.modules[__name__]
secret = _AMBIENT_SECRET
module.INJECTED_AMBIENT_GLOBAL = secret
try:
assert uses_injected_ambient(3) == 3 + len(secret)
with pytest.raises(ValueError) as exc_info:
_package_udf(uses_injected_ambient)
_assert_packaging_rejection(exc_info, secret=secret)
finally:
delattr(module, "INJECTED_AMBIENT_GLOBAL")
def test_package_udf_rejects_builtin_shadow_injection():
module = sys.modules[__name__]
secret = _BUILTIN_SHADOW_SECRET
assert not hasattr(module, "len")
module.len = secret
try:
with pytest.raises(ValueError) as exc_info:
_package_udf(uses_shadowed_builtin_len)
_assert_packaging_rejection(exc_info, secret=secret)
finally:
delattr(module, "len")
def test_package_udf_rejects_loaded_code_source_mismatch(tmp_path: Path):
secret = _SOURCE_MISMATCH_SECRET
module_name = "udf_pkg_source_mismatch_mod"
original = _temp_udf_module_source(body="return x + 1")
replacement = _temp_udf_module_source(
body=f"return x + 99 # {secret}",
secret=secret,
)
with _temporary_imported_module(tmp_path, module_name, original) as (
path,
module,
):
target = module.temp_pack_target
assert target(1) == 2
path.write_text(replacement, encoding="utf-8")
with pytest.raises(ValueError) as exc_info:
_package_udf(target)
_assert_packaging_rejection(exc_info, secret=secret)
err_text = f"{exc_info.value}\n{exc_info.value!r}"
assert str(path.resolve()) not in err_text
assert path.resolve().as_posix() not in err_text
def test_package_udf_rejects_invalid_utf8_after_import(tmp_path: Path):
secret = _INVALID_UTF8_SECRET
module_name = "udf_pkg_invalid_utf8_mod"
original = _temp_udf_module_source(body="return x + 1")
with _temporary_imported_module(tmp_path, module_name, original) as (
path,
module,
):
target = module.temp_pack_target
assert target(1) == 2
path.write_bytes(secret.encode("utf-8") + b"\xff\xfe invalid-bytes")
with pytest.raises(ValueError) as exc_info:
_package_udf(target)
assert type(exc_info.value) is ValueError
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
err_text = f"{exc_info.value}\n{exc_info.value!r}"
assert secret not in err_text
assert "b'" not in err_text
assert r"\xff" not in err_text
assert str(path.resolve()) not in err_text
assert path.resolve().as_posix() not in err_text
def test_package_udf_snapshot_has_no_durable_overdesign_and_is_not_function_or_job():
packaged = _package_udf(packable_add)
assert not isinstance(packaged, Function)
assert not isinstance(packaged, Job)
for attr in _OVERDESIGN_ATTRS:
assert not hasattr(packaged, attr)
text = repr(packaged).lower()
for token in (
"user_version",
"idempotency_key",
"deterministic",
"null_policy",
"on_error",
"functionversion",
"artifact",
"digest",
"geneva",
"retry_key",
):
assert token not in text
_assert_sanitized_text(text)
-25
View File
@@ -203,31 +203,6 @@ async def test_async_hybrid_query_default_limit(table: AsyncTable):
assert texts.count("a") == 1
def test_hybrid_query_offset(sync_table: Table):
# The offset window of a hybrid query must be a suffix of the same query
# run without an offset -- it must not be silently ignored.
full = (
sync_table.search(query_type="hybrid")
.vector([0.0, 0.4])
.text("dog")
.limit(4)
.with_row_id(True)
.to_arrow()
)
assert len(full) == 4
offset_result = (
sync_table.search(query_type="hybrid")
.vector([0.0, 0.4])
.text("dog")
.offset(2)
.limit(2)
.with_row_id(True)
.to_arrow()
)
assert offset_result["_rowid"].to_pylist() == full["_rowid"].to_pylist()[2:]
def test_hybrid_query_minimum_nprobes_zero_raises(sync_table: Table):
# minimum_nprobes(0) must raise the same validation error a plain vector
# query raises, not silently no-op because 0 is falsy.
+225 -125
View File
@@ -1133,131 +1133,6 @@ def test_stats():
assert res == stats
@contextlib.contextmanager
def lsm_test_table(lsm_handler):
"""A remote table whose LSM routes are served by ``lsm_handler``.
``lsm_handler(request, route)`` is called for ``/v1/table/test/<route>/``
where route is one of flush_lsm, compact_lsm, get_lsm_stats, and is
responsible for writing the response.
"""
routes = ("flush_lsm", "compact_lsm", "get_lsm_stats")
def handler(request):
match = re.fullmatch(r"/v1/table/test/(\w+)/", request.path)
route = match.group(1) if match else None
if route in routes:
lsm_handler(request, route)
elif route == "describe":
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(b'{"version": 1, "schema": {"fields": []}}')
else:
request.send_response(404)
request.end_headers()
with mock_lancedb_connection(handler) as db:
yield db.open_table("test")
def read_json_body(request):
content_len = int(request.headers.get("Content-Length"))
return json.loads(request.rfile.read(content_len))
def send_json(request, payload, status=200):
request.send_response(status)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(payload).encode())
def test_get_lsm_stats_sync():
"""The sync wrapper round-trips the server payload into a dict."""
bucket = {
"shard_id": "b0",
"status": "Active",
"writer_epoch": 3,
"manifest_version": 12,
"current_generation": 6,
"replay_after_wal_entry_position": 40,
"wal_entry_position_last_seen": 42,
"generations": [{"generation": 5, "bytes": 1024, "rows": 7}],
"compacting": False,
"memtables": [
{
"generation": 6,
"rows": 2,
"bytes": 64,
"batches": 1,
"indexes": ["vec_idx"],
}
],
}
seen_bodies = []
def lsm_handler(request, route):
assert route == "get_lsm_stats"
seen_bodies.append(read_json_body(request))
send_json(request, {"lsm_stats": {"buckets": [bucket]}})
with lsm_test_table(lsm_handler) as table:
assert table.get_lsm_stats() == {"buckets": [bucket]}
# Off by default, and forwarded when asked for.
assert seen_bodies == [{"include_generation_rows": False}]
table.get_lsm_stats(include_generation_rows=True)
assert seen_bodies[-1] == {"include_generation_rows": True}
def test_get_lsm_stats_sync_returns_none_when_lsm_disabled():
"""A null envelope means the LSM write path is not enabled, not an error."""
def lsm_handler(request, route):
send_json(request, {"lsm_stats": None})
with lsm_test_table(lsm_handler) as table:
assert table.get_lsm_stats() is None
def test_flush_and_compact_lsm_sync():
"""Both are one-shot POSTs answered 202 with no body."""
called = []
def lsm_handler(request, route):
called.append(route)
request.send_response(202)
request.end_headers()
with lsm_test_table(lsm_handler) as table:
assert table.flush_lsm() is None
assert table.compact_lsm() is None
assert called == ["flush_lsm", "compact_lsm"]
def test_checkpoint_lsm_sync():
"""Seal, read the watermark, and return once L0 holds nothing.
The convergence loop itself is covered in Rust; this pins the sync
binding to the endpoints it drives.
"""
called = []
def lsm_handler(request, route):
called.append(route)
if route == "get_lsm_stats":
# An empty L0 yields no target watermark, so the loop is done
# after the seal without ever polling compaction.
send_json(request, {"lsm_stats": {"buckets": []}})
else:
request.send_response(202)
request.end_headers()
with lsm_test_table(lsm_handler) as table:
assert table.checkpoint_lsm() is None
assert called == ["flush_lsm", "get_lsm_stats"]
@contextlib.contextmanager
def query_test_table(query_handler, *, server_version=Version("0.1.0")):
def handler(request):
@@ -2431,3 +2306,228 @@ def test_remote_connection_jobs_surface():
assert job.status() == "failed"
with pytest.raises(JobFailedError, match="worker died"):
job.wait(timeout=timedelta(seconds=5))
# Pinned Rust-canonical schema-only type IPC (base64). PyArrow's schema-only
# FileWriter bytes are not byte-identical to the Arrow Rust FileWriter used by
# the strict Function decoder, so these fixtures are derived from Rust serde.
_FIRST_CLASS_FUNCTION_JOB_RESULT_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_FIRST_CLASS_FUNCTION_JOB_RESULT_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID = "fn.exact.python-job-result"
_FIRST_CLASS_FUNCTION_JOB_RESULT_ABSENT = object()
_FIRST_CLASS_FUNCTION_JOB_RESULT_NULL = object()
def _first_class_function_job_result_function_wire():
int32_ipc = _FIRST_CLASS_FUNCTION_JOB_RESULT_INT32_TYPE_IPC_B64
utf8_ipc = _FIRST_CLASS_FUNCTION_JOB_RESULT_UTF8_TYPE_IPC_B64
return {
"kind": "function",
"format_version": 1,
"function": {
"format_version": 1,
"id": _FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "x", "data_type_ipc": int32_ipc},
{"name": "label", "data_type_ipc": utf8_ipc},
],
"output": {
"data_type_ipc": int32_ipc,
"nullable": True,
},
},
},
}
def _first_class_function_job_result_none_wire():
return {"kind": "none", "format_version": 1}
def _first_class_function_job_result_describe_body(
job_id, job_type, result=_FIRST_CLASS_FUNCTION_JOB_RESULT_ABSENT
):
body = {
"job_id": job_id,
"job_state": "DONE",
"job_type": job_type,
"creation_ms": 1,
"spec": {},
}
if result is _FIRST_CLASS_FUNCTION_JOB_RESULT_NULL:
body["result"] = None
elif result is not _FIRST_CLASS_FUNCTION_JOB_RESULT_ABSENT:
body["result"] = result
return body
def _first_class_function_job_result_describe_handler(bodies_by_job_id):
def handler(request):
content_len = int(request.headers.get("Content-Length", 0))
body = request.rfile.read(content_len) if content_len > 0 else b""
payload = json.loads(body) if body else {}
if request.path != "/v1/jobs/describe":
request.send_response(404)
request.end_headers()
return
job_id = payload["job_id"]
if job_id not in bodies_by_job_id:
request.send_response(404)
request.end_headers()
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(bodies_by_job_id[job_id]).encode())
return handler
def _assert_exact_first_class_function_job_result(function):
assert isinstance(function, lancedb.Function)
assert function is not None
assert not isinstance(function, dict)
assert function.id == _FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID
assert function.parameters == (("x", pa.int32()), ("label", pa.utf8()))
assert function.output_type == pa.int32()
assert function.output_nullable is True
text = repr(function)
assert "Function" in text
assert _FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID in text
for token in ("definition", "source", "packages", "artifact", "digest", "secret"):
assert token not in text.lower()
def test_first_class_function_job_result_sync_wait_returns_exact_function():
bodies = {
"job-register": _first_class_function_job_result_describe_body(
"job-register",
"register_function",
_first_class_function_job_result_function_wire(),
)
}
with mock_lancedb_connection(
_first_class_function_job_result_describe_handler(bodies)
) as db:
result = db.job("job-register").wait()
_assert_exact_first_class_function_job_result(result)
timed_out = db.job("job-register").wait(timeout=timedelta(seconds=5))
_assert_exact_first_class_function_job_result(timed_out)
with pytest.raises(TypeError):
lancedb.Function()
with pytest.raises(AttributeError):
result.id = "mutated"
with pytest.raises(AttributeError):
result.parameters = ()
with pytest.raises(AttributeError):
result.output_type = pa.int64()
with pytest.raises(AttributeError):
result.output_nullable = False
@pytest.mark.asyncio
async def test_first_class_function_job_result_async_wait_returns_exact_function():
bodies = {
"job-register": _first_class_function_job_result_describe_body(
"job-register",
"register_function",
_first_class_function_job_result_function_wire(),
)
}
async with mock_lancedb_connection_async(
_first_class_function_job_result_describe_handler(bodies)
) as db:
result = await db.job("job-register").wait()
_assert_exact_first_class_function_job_result(result)
timed_out = await db.job("job-register").wait(timeout=timedelta(seconds=5))
_assert_exact_first_class_function_job_result(timed_out)
def test_first_class_function_job_result_no_result_wait_returns_none():
bodies = {
"job-index-absent": _first_class_function_job_result_describe_body(
"job-index-absent", "create_index"
),
"job-index-explicit": _first_class_function_job_result_describe_body(
"job-index-explicit",
"create_index",
_first_class_function_job_result_none_wire(),
),
}
with mock_lancedb_connection(
_first_class_function_job_result_describe_handler(bodies)
) as db:
assert db.job("job-index-absent").wait() is None
assert db.job("job-index-explicit").wait(timeout=timedelta(seconds=5)) is None
@pytest.mark.asyncio
async def test_first_class_function_job_result_async_no_result_wait_returns_none():
bodies = {
"job-index-absent": _first_class_function_job_result_describe_body(
"job-index-absent", "create_index"
),
"job-index-explicit": _first_class_function_job_result_describe_body(
"job-index-explicit",
"create_index",
_first_class_function_job_result_none_wire(),
),
}
async with mock_lancedb_connection_async(
_first_class_function_job_result_describe_handler(bodies)
) as db:
assert await db.job("job-index-absent").wait() is None
assert (
await db.job("job-index-explicit").wait(timeout=timedelta(seconds=5))
is None
)
def test_first_class_function_job_result_get_job_result_projection():
bodies = {
"job-register": _first_class_function_job_result_describe_body(
"job-register",
"register_function",
_first_class_function_job_result_function_wire(),
),
"job-absent": _first_class_function_job_result_describe_body(
"job-absent", "create_index"
),
"job-null": _first_class_function_job_result_describe_body(
"job-null",
"create_index",
_FIRST_CLASS_FUNCTION_JOB_RESULT_NULL,
),
"job-explicit-none": _first_class_function_job_result_describe_body(
"job-explicit-none",
"create_index",
_first_class_function_job_result_none_wire(),
),
}
with mock_lancedb_connection(
_first_class_function_job_result_describe_handler(bodies)
) as db:
register_description = db.get_job("job-register")
_assert_exact_first_class_function_job_result(register_description.result)
assert db.get_job("job-absent").result is None
assert db.get_job("job-null").result is None
assert db.get_job("job-explicit-none").result is None
-62
View File
@@ -3854,65 +3854,3 @@ async def test_async_search_runs_embedding_on_dedicated_executor(
assert all(name.startswith("lancedb-embedding") for name in captured_threads), (
f"embedding ran off the dedicated executor: {captured_threads}"
)
def test_computed_column_declare_and_refresh(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table("computed", [{"x": 1}, {"x": 2}])
table.add_columns(computed={"doubled": "x * 2"})
assert table.to_arrow()["doubled"].to_pylist() == [None, None]
result = table.refresh_column("doubled")
assert result.rows_filled == 2
assert sorted(table.to_arrow()["doubled"].to_pylist()) == [2, 4]
table.add([{"x": 5}])
assert table.refresh_column("doubled").rows_filled == 1
assert sorted(table.to_arrow()["doubled"].to_pylist()) == [2, 4, 10]
def test_computed_column_rejects_transforms_and_computed_together(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table("computed_mixed", [{"x": 1}])
with pytest.raises(ValueError):
table.add_columns({"a": "x + 1"}, computed={"b": "x * 2"})
@pytest.mark.asyncio
async def test_computed_column_async(tmp_path):
db = await lancedb.connect_async(tmp_path)
table = await db.create_table("computed_async", [{"x": 3}])
await table.add_columns(computed={"tripled": "x * 3"})
await table.refresh_column("tripled")
assert (await table.to_arrow())["tripled"].to_pylist() == [9]
def test_refresh_column_async_returns_job(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table("computed_job", [{"x": 1}, {"x": 2}])
table.add_columns(computed={"doubled": "x * 2"})
job = table.refresh_column_async("doubled")
assert job.id is None # in-process jobs have no server id
job.wait()
assert job.status() == "finished"
assert sorted(table.to_arrow()["doubled"].to_pylist()) == [2, 4]
# Bad input raises at the call, not through the job.
with pytest.raises(Exception, match="not a computed column"):
table.refresh_column_async("x")
@pytest.mark.asyncio
async def test_refresh_column_async_job_async_table(tmp_path):
db = await lancedb.connect_async(tmp_path)
table = await db.create_table("computed_job_async", [{"x": 3}])
await table.add_columns(computed={"tripled": "x * 3"})
job = await table.refresh_column_async("tripled")
await job.wait()
assert await job.status() == "finished"
assert (await table.to_arrow())["tripled"].to_pylist() == [9]
+124 -30
View File
@@ -23,6 +23,7 @@ use lancedb::{
connection::NamespaceClientPushdownOperation,
database::namespace::LanceNamespaceDatabase,
database::{CreateTableMode, Database, ReadConsistency},
function::{FunctionId, RegisterFunctionJobSpec},
};
use pyo3::{
Bound, FromPyObject, Py, PyAny, PyRef, PyResult, Python,
@@ -121,8 +122,6 @@ impl Connection {
}
#[pyo3(signature = (namespace_path=None, start_after=None, limit=None))]
// Deprecated in favour of `list_tables`, but still exposed to Python.
#[allow(deprecated)]
pub fn table_names(
self_: PyRef<'_, Self>,
namespace_path: Option<Vec<String>>,
@@ -348,23 +347,6 @@ impl Connection {
})
}
#[pyo3(signature = (name, namespace_path=None))]
pub fn drop_table_async(
self_: PyRef<'_, Self>,
name: String,
namespace_path: Option<Vec<String>>,
) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.get_inner()?.clone();
let ns_path = namespace_path.unwrap_or_default();
future_into_py(self_.py(), async move {
inner
.drop_table_async(name, &ns_path)
.await
.infer_error()
.map(crate::job::Job::new)
})
}
#[pyo3(signature = (namespace_path=None,))]
pub fn drop_all_tables(
self_: PyRef<'_, Self>,
@@ -524,17 +506,14 @@ impl Connection {
let inner = self_.get_inner()?.clone();
let py = self_.py();
future_into_py(py, async move {
let mut request = inner.list_tables();
if let Some(namespace_path) = namespace_path {
request = request.namespace(namespace_path);
}
if let Some(page_token) = page_token {
request = request.page_token(page_token);
}
if let Some(limit) = limit {
request = request.limit(limit);
}
let response = request.execute().await.infer_error()?;
use lance_namespace::models::ListTablesRequest;
let request = ListTablesRequest {
id: namespace_path,
page_token,
limit: limit.map(|l| l as i32),
..Default::default()
};
let response = inner.list_tables(request).await.infer_error()?;
Python::attach(|py| -> PyResult<Py<PyDict>> {
let dict = PyDict::new(py);
dict.set_item("tables", response.tables)?;
@@ -611,6 +590,121 @@ impl Connection {
})
})
}
/// Submit a first-class Function registration job.
///
/// Accepts the exact private [`crate::function::PyFunctionDefinition`] and
/// builds [`RegisterFunctionJobSpec`] with `expected_current_function_id =
/// None` (create-if-absent). Does not JSON round-trip the definition.
pub fn _register_function<'py>(
self_: PyRef<'py, Self>,
name: String,
definition: Bound<'_, crate::function::PyFunctionDefinition>,
) -> PyResult<Bound<'py, PyAny>> {
let inner = self_.get_inner()?.clone();
let definition = definition.get().inner().clone();
future_into_py(self_.py(), async move {
let spec = RegisterFunctionJobSpec::try_new(name, definition, None).infer_error()?;
let job = inner.register_function(spec).await.infer_error()?;
Ok(crate::job::Job::new(job))
})
}
/// Submit a first-class Function conditional replace job.
///
/// Accepts the observed native [`crate::function::Function`] handle and the
/// exact private [`crate::function::PyFunctionDefinition`], then builds
/// [`RegisterFunctionJobSpec`] with `expected_current_function_id =
/// Some(current.id)`. Reads only `current.inner().id().clone()`. Does not
/// JSON round-trip the definition.
pub fn _replace_function<'py>(
self_: PyRef<'py, Self>,
name: String,
current: Bound<'_, crate::function::Function>,
definition: Bound<'_, crate::function::PyFunctionDefinition>,
) -> PyResult<Bound<'py, PyAny>> {
let inner = self_.get_inner()?.clone();
let definition = definition.get().inner().clone();
let current_id = current.get().inner().id().clone();
future_into_py(self_.py(), async move {
let spec = RegisterFunctionJobSpec::try_new(name, definition, Some(current_id))
.infer_error()?;
let job = inner.register_function(spec).await.infer_error()?;
Ok(crate::job::Job::new(job))
})
}
/// Look up the Function currently bound to a database-scoped name.
///
/// Wraps the exact Rust [`lancedb::function::Function`] once. Empty names
/// fail as [`PyValueError`] before transport via the Rust connection.
pub fn _lookup_function_by_name<'py>(
self_: PyRef<'py, Self>,
name: String,
) -> PyResult<Bound<'py, PyAny>> {
let inner = self_.get_inner()?.clone();
future_into_py(self_.py(), async move {
let function = inner.lookup_function_by_name(&name).await.infer_error()?;
Ok(crate::function::Function::new(function))
})
}
/// Look up an immutable Function by exact opaque Function ID string.
///
/// Constructs [`FunctionId`] with [`FunctionId::try_new`] before dispatch so
/// empty IDs fail as [`PyValueError`] before transport.
pub fn _lookup_function_by_id<'py>(
self_: PyRef<'py, Self>,
function_id: String,
) -> PyResult<Bound<'py, PyAny>> {
let inner = self_.get_inner()?.clone();
future_into_py(self_.py(), async move {
let id = FunctionId::try_new(function_id).infer_error()?;
let function = inner.lookup_function_by_id(&id).await.infer_error()?;
Ok(crate::function::Function::new(function))
})
}
/// Conditionally remove a database-scoped Function catalog name.
///
/// Clones the observed native [`crate::function::Function`] once and
/// delegates to Rust [`lancedb::Connection::remove_function_name`]. Empty
/// names fail as [`PyValueError`] before transport via the Rust connection.
pub fn _remove_function_name<'py>(
self_: PyRef<'py, Self>,
name: String,
current: Bound<'_, crate::function::Function>,
) -> PyResult<Bound<'py, PyAny>> {
let inner = self_.get_inner()?.clone();
let current = current.get().inner().clone();
future_into_py(self_.py(), async move {
inner
.remove_function_name(&name, &current)
.await
.infer_error()?;
// `()` maps to an empty Python tuple via IntoPyObject; return Option
// so the async bridge yields exact Python None.
Ok(None::<()>)
})
}
/// Revoke an exact immutable Function by administrator set-bit.
///
/// Clones the observed native [`crate::function::Function`] once and
/// delegates to Rust [`lancedb::Connection::revoke_function`].
pub fn _revoke_function<'py>(
self_: PyRef<'py, Self>,
function: Bound<'_, crate::function::Function>,
) -> PyResult<Bound<'py, PyAny>> {
let inner = self_.get_inner()?.clone();
let function = function.get().inner().clone();
future_into_py(self_.py(), async move {
inner.revoke_function(&function).await.infer_error()?;
// `()` maps to an empty Python tuple via IntoPyObject; return Option
// so the async bridge yields exact Python None.
Ok(None::<()>)
})
}
}
#[pyfunction]
+15 -2
View File
@@ -102,11 +102,14 @@ impl<T> PythonErrorExt<T> for std::result::Result<T, LanceError> {
err.setattr(intern!(py, "__cause__"), cause_err)?;
Err(PyErr::from_value(err))
}),
LanceError::JobFailed { .. } => Python::attach(|py| {
LanceError::JobFailed { failure, .. } => Python::attach(|py| {
let cls = py
.import(intern!(py, "lancedb.exceptions"))?
.getattr(intern!(py, "JobFailedError"))?;
Err(PyErr::from_value(cls.call1((err.to_string(),))?))
// Structural projection only: failure.error_code.as_str().
// Never infer a code from message, phase, retryable, or source.
let error_code = failure.error_code.as_ref().map(|code| code.as_str());
Err(PyErr::from_value(cls.call1((err.to_string(), error_code))?))
}),
LanceError::JobCancelled { .. } => Python::attach(|py| {
let cls = py
@@ -114,6 +117,16 @@ impl<T> PythonErrorExt<T> for std::result::Result<T, LanceError> {
.getattr(intern!(py, "JobCancelledError"))?;
Err(PyErr::from_value(cls.call1((err.to_string(),))?))
}),
LanceError::Function { code, message } => Python::attach(|py| {
let cls = py
.import(intern!(py, "lancedb.exceptions"))?
.getattr(intern!(py, "FunctionError"))?;
// Structural projection only: code.as_str() + sanitized message.
// Never infer a code from HTTP status or diagnostic text.
Err(PyErr::from_value(
cls.call1((message.as_str(), code.as_str()))?,
))
}),
_ => self.runtime_error(),
},
}
+29 -21
View File
@@ -10,7 +10,7 @@
use std::ops::{Add, Div, Mul, Not, Sub};
use arrow::{datatypes::DataType, pyarrow::PyArrowType};
use datafusion_common::ScalarValue;
use datafusion_common::{Column, ScalarValue};
use lancedb::expr::{
DfExpr, col as ldb_col, contains, expr_cast, is_in, lit as df_lit, lower, upper,
};
@@ -27,6 +27,33 @@ use pyo3::{Bound, PyAny, PyResult, exceptions::PyValueError, prelude::*, pyfunct
#[derive(Clone)]
pub struct PyExpr(pub DfExpr);
/// Crate-private inspection result for Function call authoring (FF-028).
#[derive(Debug, Clone)]
pub(crate) enum DirectExprView<'a> {
/// Direct unqualified DataFusion Column; name is case-sensitive.
UnqualifiedColumn(&'a str),
/// Direct Literal scalar; Arrow type is owned by the scalar value.
Literal(&'a ScalarValue),
}
impl PyExpr {
/// Inspect a direct Column/Literal node for Function call authoring.
///
/// Returns `None` for every other expression shape (arithmetic, cast,
/// scalar function, predicate, alias, qualified column, etc.).
pub(crate) fn as_direct_column_or_literal(&self) -> Option<DirectExprView<'_>> {
match &self.0 {
DfExpr::Column(Column {
relation: None,
name,
..
}) => Some(DirectExprView::UnqualifiedColumn(name.as_str())),
DfExpr::Literal(value, _) => Some(DirectExprView::Literal(value)),
_ => None,
}
}
}
#[pymethods]
impl PyExpr {
// ── comparisons ──────────────────────────────────────────────────────────
@@ -191,27 +218,8 @@ pub fn expr_lit(value: Bound<'_, PyAny>) -> PyResult<PyExpr> {
}
// datetime.datetime is a subclass of datetime.date, so it must be checked first.
//
// Python's datetime.timestamp() treats *naive* datetimes as local wall time.
// PyArrow (and therefore Lance table storage) encodes naive timestamps as
// UTC wall-clock microseconds. Using .timestamp() for naive values therefore
// shifts the literal by the local UTC offset on non-UTC machines, so
// `col("ts") == lit(naive_dt)` fails against a table that holds the same
// naive value. Fix: treat naive datetimes as UTC wall clock (match Arrow);
// keep aware datetimes on the real .timestamp() path (correct epoch).
if let Ok(dt) = value.cast::<PyDateTime>() {
let ts: f64 = if dt.getattr("tzinfo")?.is_none() {
// Force UTC interpretation of the naive wall clock.
let utc = pyo3::types::PyModule::import(value.py(), "datetime")?
.getattr("timezone")?
.getattr("utc")?;
let kwargs = pyo3::types::PyDict::new(value.py());
kwargs.set_item("tzinfo", utc)?;
let aware = dt.call_method("replace", (), Some(&kwargs))?;
aware.call_method0("timestamp")?.extract()?
} else {
dt.call_method0("timestamp")?.extract()?
};
let ts: f64 = dt.call_method0("timestamp")?.extract()?;
let micros = (ts * 1_000_000.0).round() as i64;
return Ok(PyExpr(df_lit(ScalarValue::TimestampMicrosecond(
Some(micros),
File diff suppressed because it is too large Load Diff
+31 -4
View File
@@ -3,6 +3,7 @@
use std::sync::Arc;
use crate::function::Function;
use crate::runtime::future_into_py;
use pyo3::{Bound, PyAny, PyRef, PyResult, pyclass, pymethods};
@@ -21,6 +22,23 @@ impl Job {
}
}
/// Project a Rust [`lancedb::JobResult`] onto the Python success surface.
///
/// Delegates variant interpretation to [`lancedb::JobResult::into_function`]:
/// no nested Function collapses to Python `None`; an exact Function becomes
/// the corresponding [`Function`] handle.
fn project_wait_result(result: lancedb::JobResult) -> Option<Function> {
result.into_function().map(Function::new)
}
/// Project a describe `result` onto Python `Optional[Function]`.
///
/// Rust `None`, `Some(JobResult::None)`, and JSON null all become Python
/// `None`. Only `Some(JobResult::Function)` becomes a [`Function`] handle.
fn project_description_result(result: Option<lancedb::JobResult>) -> Option<Function> {
result.and_then(project_wait_result)
}
#[pymethods]
impl Job {
#[getter]
@@ -39,8 +57,8 @@ impl Job {
pub fn wait(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.inner.clone();
future_into_py(self_.py(), async move {
inner.wait().await.infer_error()?;
Ok(())
let result = inner.wait().await.infer_error()?;
Ok(project_wait_result(result))
})
}
@@ -93,14 +111,16 @@ pub struct JobFailureInfo {
phase: Option<String>,
message: Option<String>,
retryable: Option<bool>,
/// Exact wire `error_code` string when Rust decoded one; never inferred.
error_code: Option<String>,
}
#[pymethods]
impl JobFailureInfo {
fn __repr__(&self) -> String {
format!(
"JobFailureInfo(phase={:?}, message={:?}, retryable={:?})",
self.phase, self.message, self.retryable
"JobFailureInfo(phase={:?}, message={:?}, retryable={:?}, error_code={:?})",
self.phase, self.message, self.retryable, self.error_code
)
}
}
@@ -115,6 +135,7 @@ pub struct JobDescription {
creation_ms: i64,
spec_json: Option<String>,
failure: Option<JobFailureInfo>,
result: Option<Function>,
}
#[pymethods]
@@ -139,7 +160,13 @@ impl From<lancedb::database::JobDescription> for JobDescription {
phase: failure.phase,
message: failure.message,
retryable: failure.retryable,
// Structural projection only: exact as_str(); never infer.
error_code: failure
.error_code
.as_ref()
.map(|code| code.as_str().to_string()),
}),
result: project_description_result(description.result),
}
}
}
+9 -3
View File
@@ -16,14 +16,14 @@ use query::{FTSQuery, HybridQuery, Query, VectorQuery};
use session::Session;
use table::{
AddColumnsResult, AddResult, AlterColumnsResult, DeleteResult, DropColumnsResult, FtsToken,
LsmWriteSpec, MergeResult, PyBlobFile, RefreshColumnResult, Table, UpdateFieldMetadataResult,
UpdateResult,
LsmWriteSpec, MergeResult, PyBlobFile, Table, UpdateFieldMetadataResult, UpdateResult,
};
pub mod arrow;
pub mod connection;
pub mod error;
pub mod expr;
pub mod function;
pub mod header;
pub mod index;
pub mod job;
@@ -46,6 +46,9 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<Connection>()?;
m.add_class::<Session>()?;
m.add_class::<Table>()?;
m.add_class::<crate::function::Function>()?;
m.add_class::<crate::function::PyFunctionDefinition>()?;
m.add_class::<crate::function::AuthoredFunctionCall>()?;
m.add_class::<crate::job::Job>()?;
m.add_class::<crate::job::JobInfo>()?;
m.add_class::<crate::job::JobDescription>()?;
@@ -58,7 +61,6 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<VectorQuery>()?;
m.add_class::<RecordBatchStream>()?;
m.add_class::<AddColumnsResult>()?;
m.add_class::<RefreshColumnResult>()?;
m.add_class::<AlterColumnsResult>()?;
m.add_class::<UpdateFieldMetadataResult>()?;
m.add_class::<AddResult>()?;
@@ -90,6 +92,10 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(wrap_pyfunction!(expr_col, m)?)?;
m.add_function(wrap_pyfunction!(expr_lit, m)?)?;
m.add_function(wrap_pyfunction!(expr_func, m)?)?;
m.add_function(wrap_pyfunction!(
crate::function::_new_function_definition,
m
)?)?;
m.add("__version__", env!("CARGO_PKG_VERSION"))?;
Ok(())
}
+141 -61
View File
@@ -26,7 +26,7 @@ use lancedb::table::{
use lancedb::tokenize as lancedb_tokenize;
use pyo3::{
Bound, FromPyObject, Py, PyAny, PyRef, PyResult, Python,
exceptions::{PyRuntimeError, PyValueError},
exceptions::{PyNotImplementedError, PyRuntimeError, PyValueError},
pyclass, pyfunction, pymethods,
types::{IntoPyDict, PyAnyMethods, PyBytes, PyDict, PyDictMethods, PyList, PyListMethods},
};
@@ -415,32 +415,6 @@ pub struct AddColumnsResult {
pub version: u64,
}
#[pyclass(get_all, from_py_object)]
#[derive(Clone, Debug)]
pub struct RefreshColumnResult {
pub rows_filled: u64,
pub version: u64,
}
#[pymethods]
impl RefreshColumnResult {
pub fn __repr__(&self) -> String {
format!(
"RefreshColumnResult(rows_filled={}, version={})",
self.rows_filled, self.version
)
}
}
impl From<lancedb::table::RefreshColumnResult> for RefreshColumnResult {
fn from(result: lancedb::table::RefreshColumnResult) -> Self {
Self {
rows_filled: result.rows_filled,
version: result.version,
}
}
}
#[pymethods]
impl AddColumnsResult {
pub fn __repr__(&self) -> String {
@@ -956,6 +930,146 @@ impl Table {
})
}
/// Hidden bridge: bind an authored Function call once and submit create.
///
/// Private native path for Python ``table.add_generated_column``. Rejects an
/// empty ``column_name`` before reading the table handle. Does not expose
/// source version, stable field IDs, the operation spec, or request envelope.
#[doc(hidden)]
pub fn _add_generated_column<'a>(
self_: PyRef<'a, Self>,
column_name: String,
call: Bound<'_, crate::function::AuthoredFunctionCall>,
) -> PyResult<Bound<'a, PyAny>> {
if column_name.is_empty() {
return Err(PyValueError::new_err("column_name must be non-empty"));
}
let inner = self_.inner_ref()?.clone();
let authored = call.get().clone();
future_into_py(self_.py(), async move {
let (source_table_version, bound_call) =
authored.bind_to_table(&inner).await.infer_error()?;
let spec = lancedb::function::CreateGeneratedColumnJobSpec::try_new(
column_name,
authored.function(),
bound_call,
)
.infer_error()?;
let job = inner
.submit_create_generated_column(source_table_version, spec)
.await
.infer_error()?;
Ok(crate::job::Job::new(job))
})
}
/// Hidden bridge: project generated-column completeness for one column name.
///
/// Private native path for Python ``table.generated_column_status``. Rejects
/// an empty ``column_name`` before reading the table handle. Maps only the
/// known Rust status variants to ``"complete"`` / ``"incomplete"``.
#[doc(hidden)]
pub fn _generated_column_status<'a>(
self_: PyRef<'a, Self>,
column_name: String,
) -> PyResult<Bound<'a, PyAny>> {
if column_name.is_empty() {
return Err(PyValueError::new_err("column_name must be non-empty"));
}
let inner = self_.inner_ref()?.clone();
future_into_py(self_.py(), async move {
let status = inner
.generated_column_status(column_name)
.await
.infer_error()?;
match status {
lancedb::function::GeneratedColumnStatus::Complete => Ok("complete"),
lancedb::function::GeneratedColumnStatus::Incomplete => Ok("incomplete"),
_ => Err(PyNotImplementedError::new_err(
"unsupported generated column status",
)),
}
})
}
/// Hidden bridge: load exact definition, resolve Function by ID, submit refresh.
///
/// Private native path for Python ``table.refresh_generated_column``. Rejects
/// an empty ``column_name`` before reading the table handle. Does not expose
/// source version, Function, field IDs, epochs, specs, or request envelope.
#[doc(hidden)]
pub fn _refresh_generated_column<'a>(
self_: PyRef<'a, Self>,
column_name: String,
) -> PyResult<Bound<'a, PyAny>> {
if column_name.is_empty() {
return Err(PyValueError::new_err("column_name must be non-empty"));
}
let inner = self_.inner_ref()?.clone();
future_into_py(self_.py(), async move {
let (source_table_version, definition) = inner
.generated_column_definition_snapshot(column_name)
.await
.infer_error()?;
let function_id = definition.function_call().function_id().clone();
let function = inner
.resolve_function_for_generated_column(&function_id)
.await
.infer_error()?;
let spec =
lancedb::function::RefreshGeneratedColumnJobSpec::try_new(&function, definition)
.infer_error()?;
let job = inner
.submit_refresh_generated_column(source_table_version, spec)
.await
.infer_error()?;
Ok(crate::job::Job::new(job))
})
}
/// Hidden bridge: one binding snapshot, bind new call, submit change.
///
/// Private native path for Python ``table.alter_generated_column``. Rejects
/// an empty ``column_name`` before reading the table handle. Fetches exactly
/// one binding snapshot, loads the expected definition from that same
/// object, binds the authored call against it, and submits change. Does not
/// expose source version, Function handles, field IDs, epochs, specs, or
/// request envelope.
#[doc(hidden)]
pub fn _alter_generated_column<'a>(
self_: PyRef<'a, Self>,
column_name: String,
new_call: Bound<'_, crate::function::AuthoredFunctionCall>,
) -> PyResult<Bound<'a, PyAny>> {
if column_name.is_empty() {
return Err(PyValueError::new_err("column_name must be non-empty"));
}
let inner = self_.inner_ref()?.clone();
let authored = new_call.get().clone();
future_into_py(self_.py(), async move {
let snapshot = inner
.generated_column_binding_snapshot()
.await
.infer_error()?;
let expected_definition = snapshot
.generated_column_definition(&column_name)
.infer_error()?;
let (source_table_version, bound_new_call) =
authored.bind_against_snapshot(&snapshot).infer_error()?;
let spec = lancedb::function::ChangeGeneratedColumnJobSpec::try_new(
expected_definition,
authored.function(),
bound_new_call,
)
.infer_error()?;
let job = inner
.submit_change_generated_column(source_table_version, spec)
.await
.infer_error()?;
Ok(crate::job::Job::new(job))
})
}
pub fn drop_index(self_: PyRef<'_, Self>, index_name: String) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.inner_ref()?.clone();
future_into_py(self_.py(), async move {
@@ -1536,40 +1650,6 @@ impl Table {
})
}
pub fn add_computed_columns(
self_: PyRef<'_, Self>,
columns: Vec<(String, String)>,
) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.inner_ref()?.clone();
future_into_py(self_.py(), async move {
let mut builder = inner.add_columns();
for (name, expression) in columns {
builder = builder.computed(name, expression);
}
let result = builder.execute().await.infer_error()?;
Ok(AddColumnsResult::from(result))
})
}
pub fn refresh_column(self_: PyRef<'_, Self>, column: String) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.inner_ref()?.clone();
future_into_py(self_.py(), async move {
let result = inner.refresh_column(column).await.infer_error()?;
Ok(RefreshColumnResult::from(result))
})
}
pub fn refresh_column_async(
self_: PyRef<'_, Self>,
column: String,
) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.inner_ref()?.clone();
future_into_py(self_.py(), async move {
let job = inner.refresh_column_async(column).await.infer_error()?;
Ok(crate::job::Job::new(job))
})
}
pub fn add_columns_with_schema(
self_: PyRef<'_, Self>,
schema: PyArrowType<Schema>,
+2 -4
View File
@@ -1,6 +1,6 @@
[package]
name = "lancedb"
version = "0.38.0-beta.2"
version = "0.37.1-beta.1"
edition.workspace = true
description = "LanceDB: A serverless, low-latency vector database for AI applications"
license.workspace = true
@@ -12,6 +12,7 @@ rust-version.workspace = true
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
[dependencies]
ahash = { workspace = true }
base64 = "0.22"
arrow = { workspace = true }
arrow-array = { workspace = true }
arrow-buffer = { workspace = true }
@@ -188,9 +189,6 @@ required-features = ["bedrock"]
[[example]]
name = "bench_streaming_dataloader"
[[example]]
name = "bench_open_missing_table"
[[example]]
name = "simple"
@@ -1,150 +0,0 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
// Release benchmark for opening a missing table as sibling-table cardinality grows.
//
// The fixture uses real `.lance` directories and marker files. Fixture creation is
// outside the timed section. Defaults intentionally cover 1k, 10k, and 100k siblings
// with 10 warmups and 100 distinct missing-table opens per scale:
//
// ```text
// cargo run --release -p lancedb --example bench_open_missing_table
// ```
//
// `BENCH_SIBLINGS`, `BENCH_WARMUPS`, and `BENCH_TRIALS` override those defaults.
// Reduced settings are useful only as a smoke test. Performance comparisons require
// the same machine, filesystem, fixture sizes, settings, lockfile, and alternating
// baseline/candidate execution order.
use std::time::{Duration, Instant};
use anyhow::{Context, Result, bail};
use lancedb::connection::Connection;
use lancedb::{Error, connect};
use object_store::ObjectStoreExt as _;
use object_store::path::Path;
const MAX_SIBLINGS: usize = 1_000_000;
const MAX_WARMUPS: usize = 10_000;
const MAX_TRIALS: usize = 100_000;
fn env_usize(key: &str, default: usize, max: usize) -> Result<usize> {
let value = match std::env::var(key) {
Ok(value) => value
.parse()
.with_context(|| format!("invalid {key} value: {value}"))?,
Err(std::env::VarError::NotPresent) => default,
Err(error) => return Err(error).with_context(|| format!("reading {key}")),
};
if value == 0 || value > max {
bail!("{key} must be between 1 and {max}");
}
Ok(value)
}
fn sibling_counts() -> Result<Vec<usize>> {
let raw = std::env::var("BENCH_SIBLINGS").unwrap_or_else(|_| "1000,10000,100000".into());
let mut counts = raw
.split(',')
.map(|value| {
value
.trim()
.parse::<usize>()
.with_context(|| format!("invalid BENCH_SIBLINGS value: {value}"))
})
.collect::<Result<Vec<_>>>()?;
counts.sort_unstable();
counts.dedup();
if counts.is_empty() || counts[0] == 0 || counts[counts.len() - 1] > MAX_SIBLINGS {
bail!("BENCH_SIBLINGS values must be between 1 and {MAX_SIBLINGS}");
}
Ok(counts)
}
async fn add_siblings(
store: &object_store::local::LocalFileSystem,
start: usize,
end: usize,
) -> Result<()> {
for index in start..end {
let marker = Path::from(format!("sibling_{index:06}.lance/_marker"));
store
.put(&marker, bytes::Bytes::new().into())
.await
.with_context(|| format!("creating benchmark marker {marker}"))?;
}
Ok(())
}
async fn time_missing_open(db: &Connection, name: &str) -> Result<Duration> {
let started = Instant::now();
let result = db.open_table(name).execute().await;
let elapsed = started.elapsed();
match result {
Err(Error::TableNotFound { .. }) => Ok(elapsed),
Err(error) => bail!("expected TableNotFound for {name}, got {error:?}"),
Ok(_) => bail!("benchmark missing-table name unexpectedly exists: {name}"),
}
}
fn percentile(sorted: &[Duration], percentile: usize) -> Duration {
let rank = (sorted.len() * percentile).div_ceil(100).saturating_sub(1);
sorted[rank]
}
#[tokio::main]
async fn main() -> Result<()> {
let counts = sibling_counts()?;
let warmups = env_usize("BENCH_WARMUPS", 10, MAX_WARMUPS)?;
let trials = env_usize("BENCH_TRIALS", 100, MAX_TRIALS)?;
let fixture = tempfile::tempdir().context("creating benchmark fixture")?;
let database_path = fixture.path();
let fixture_store = object_store::local::LocalFileSystem::new_with_prefix(database_path)
.context("creating benchmark object store")?;
let db = connect(database_path.to_str().context("non-UTF-8 fixture path")?)
.execute()
.await?;
println!(
"config: siblings={counts:?} warmups={warmups} trials={trials} profile={} os={} arch={}",
if cfg!(debug_assertions) {
"debug"
} else {
"release"
},
std::env::consts::OS,
std::env::consts::ARCH,
);
println!("lower is better; fixture setup and teardown are excluded");
println!("| siblings | samples | p50 | p95 | max |");
println!("| ---: | ---: | ---: | ---: | ---: |");
let mut created = 0;
for sibling_count in counts {
add_siblings(&fixture_store, created, sibling_count).await?;
created = sibling_count;
for index in 0..warmups {
let name = format!("__missing_warmup_{sibling_count}_{index}");
let _ = time_missing_open(&db, &name).await?;
}
let mut samples = Vec::with_capacity(trials);
for index in 0..trials {
let name = format!("__missing_trial_{sibling_count}_{index}");
samples.push(time_missing_open(&db, &name).await?);
}
samples.sort_unstable();
println!(
"| {sibling_count} | {} | {:?} | {:?} | {:?} |",
samples.len(),
percentile(&samples, 50),
percentile(&samples, 95),
samples[samples.len() - 1],
);
}
Ok(())
}
+1 -1
View File
@@ -27,7 +27,7 @@ async fn main() -> Result<()> {
// --8<-- [end:connect]
// --8<-- [start:list_names]
println!("{:?}", db.list_tables().execute().await?.tables);
println!("{:?}", db.table_names().execute().await?);
// --8<-- [end:list_names]
let tbl = create_table(&db).await?;
create_index(&tbl).await?;
+13 -8
View File
@@ -333,11 +333,13 @@ pub(crate) fn ensure_blob_storage_version(schema: &Schema, params: &mut WritePar
.data_storage_version
.unwrap_or(LanceFileVersion::Stable)
.resolve();
if matches!(
resolved,
ConcreteFileVersion::V1 | ConcreteFileVersion::V2_0 | ConcreteFileVersion::V2_1
) {
params.data_storage_version = Some(LanceFileVersion::V2_2);
// Exact formats deliberately have no Ord: capability is not implied by
// release order. Enumerate every current concrete variant explicitly.
match resolved {
ConcreteFileVersion::V1 | ConcreteFileVersion::V2_0 | ConcreteFileVersion::V2_1 => {
params.data_storage_version = Some(LanceFileVersion::V2_2);
}
ConcreteFileVersion::V2_2 | ConcreteFileVersion::V2_3 => {}
}
}
@@ -502,7 +504,7 @@ mod tests {
ensure_blob_storage_version(&blob_schema(), &mut params);
assert_eq!(
params.data_storage_version.unwrap().resolve(),
ConcreteFileVersion::V2_2
LanceFileVersion::V2_2.resolve()
);
}
@@ -515,7 +517,7 @@ mod tests {
ensure_blob_storage_version(&blob_schema(), &mut params);
assert_eq!(
params.data_storage_version.unwrap().resolve(),
ConcreteFileVersion::V2_2
LanceFileVersion::V2_2.resolve()
);
}
@@ -526,7 +528,10 @@ mod tests {
..Default::default()
};
ensure_blob_storage_version(&blob_schema(), &mut params);
assert_eq!(params.data_storage_version.unwrap(), LanceFileVersion::V2_3);
assert_eq!(
params.data_storage_version.unwrap().resolve(),
LanceFileVersion::V2_3.resolve()
);
}
#[test]
+91 -176
View File
@@ -28,6 +28,7 @@ use crate::database::{
};
use crate::embeddings::{EmbeddingRegistry, MemoryRegistry};
use crate::error::{Error, Result};
use crate::function::{Function, FunctionId, RegisterFunctionJobSpec};
#[cfg(feature = "remote")]
use crate::remote::{
client::ClientConfig,
@@ -73,13 +74,11 @@ fn set_storage_options_provider(
}
/// A builder for configuring a [`Connection::table_names`] operation
#[deprecated(note = "Use Connection::list_tables instead")]
pub struct TableNamesBuilder {
parent: Arc<dyn Database>,
request: TableNamesRequest,
}
#[allow(deprecated)]
impl TableNamesBuilder {
fn new(parent: Arc<dyn Database>) -> Self {
Self {
@@ -117,57 +116,6 @@ impl TableNamesBuilder {
}
}
/// A builder for configuring a [`Connection::list_tables`] operation
pub struct ListTablesBuilder {
parent: Arc<dyn Database>,
request: ListTablesRequest,
}
impl ListTablesBuilder {
fn new(parent: Arc<dyn Database>) -> Self {
Self {
parent,
request: ListTablesRequest {
// The root namespace is an empty path, not an absent one: a
// namespace-backed database rejects a request that names no namespace.
id: Some(Vec::new()),
..Default::default()
},
}
}
/// Resume listing from a previous page.
///
/// Pass the `page_token` from the previous [`ListTablesResponse`]. The token is
/// opaque: it carries whatever the database needs to resume, and callers should
/// not construct or interpret one. A response whose token is `None` or empty is
/// the end of the listing.
pub fn page_token(mut self, page_token: impl Into<String>) -> Self {
self.request.page_token = Some(page_token.into());
self
}
/// An upper bound on how many tables to return.
///
/// A page may hold fewer than this and still not be the last one, so continue
/// while the response carries a page token rather than while pages are full.
pub fn limit(mut self, limit: u32) -> Self {
self.request.limit = Some(i32::try_from(limit).unwrap_or(i32::MAX));
self
}
/// Set the namespace path to list tables from. Defaults to the root namespace.
pub fn namespace(mut self, namespace_path: Vec<String>) -> Self {
self.request.id = Some(namespace_path);
self
}
/// Execute the list tables operation
pub async fn execute(self) -> Result<ListTablesResponse> {
self.parent.clone().list_tables(self.request).await
}
}
#[derive(Clone, Debug)]
pub struct OpenTableBuilder {
parent: Arc<dyn Database>,
@@ -462,14 +410,7 @@ impl Connection {
///
/// The names will be returned in lexicographical order (ascending)
///
/// Listing databases discover physical `*.lance` entries without opening every
/// dataset. The result is a point-in-time discovery snapshot: an entry may still be
/// under creation, may contain only uncommitted storage, or may be concurrently
/// dropped before it is opened.
///
/// The parameters `start_after` and `limit` can be used to paginate the results
#[deprecated(note = "Use Connection::list_tables instead")]
#[allow(deprecated)]
/// The parameters `page_token` and `limit` can be used to paginate the results
pub fn table_names(&self) -> TableNamesBuilder {
TableNamesBuilder::new(self.internal.clone())
}
@@ -516,9 +457,10 @@ impl Connection {
///
/// # Returns
/// Created [`TableRef`], or [`Error::TableNotFound`] if the table does not exist.
/// On listing databases, a committed Lance manifest is authoritative for table
/// existence. Uncommitted files or a physical `<name>.lance` directory alone do not
/// make a table openable.
/// If the table's storage is present but holds no readable dataset (for example a
/// `<name>.lance` directory left behind by an interrupted drop and re-create, which
/// [`Self::table_names`] still lists) this returns [`Error::TableCorrupted`]
/// instead.
pub fn open_table(&self, name: impl Into<String>) -> OpenTableBuilder {
OpenTableBuilder::new(
self.internal.clone(),
@@ -609,6 +551,88 @@ impl Connection {
self.internal.job_history(job_id).await
}
/// Submit a first-class Function registration job.
///
/// Returns a [`crate::job::Job`] handle for the accepted server-side job.
/// Only remote databases support registration; local databases return
/// [`Error::NotSupported`].
pub async fn register_function(
&self,
spec: RegisterFunctionJobSpec,
) -> Result<crate::job::Job> {
self.internal.register_function(spec).await
}
/// Look up the Function currently bound to a database-scoped name.
///
/// The name is lookup indirection only and is never part of the returned
/// [`Function`]. Empty names return [`Error::InvalidInput`] before backend
/// dispatch. Only remote databases support enterprise catalog lookup;
/// nonempty local lookups return [`Error::NotSupported`].
pub async fn lookup_function_by_name(&self, name: impl AsRef<str>) -> Result<Function> {
let name = name.as_ref();
// Public nonempty invariant: validate before any Database backend sees
// the call so local and remote Connections agree on InvalidInput.
if name.is_empty() {
return Err(Error::InvalidInput {
message: "function lookup name must be non-empty".into(),
});
}
self.internal.lookup_function_by_name(name).await
}
/// Look up an immutable Function by exact opaque [`FunctionId`].
///
/// Exact-ID lookup is independent of later catalog name changes. Only
/// remote databases support enterprise catalog lookup; local databases
/// return [`Error::NotSupported`].
pub async fn lookup_function_by_id(&self, function_id: &FunctionId) -> Result<Function> {
self.internal.lookup_function_by_id(function_id).await
}
/// Conditionally remove a database-scoped Function catalog name.
///
/// This is a direct synchronous catalog compare-and-swap (CAS), not a
/// [`crate::job::Job`], not physical [`Function`] deletion, and not
/// revocation. The caller supplies an observed immutable [`Function`]
/// handle; only [`Function::id`] is authority for the CAS precondition.
///
/// Empty names return [`Error::InvalidInput`] before backend dispatch.
/// Nonempty names on local/default backends return [`Error::NotSupported`].
/// Remote backends complete only when the server reports durable CAS
/// success for the `(name, current.id)` pair.
pub async fn remove_function_name(
&self,
name: impl AsRef<str>,
current: &Function,
) -> Result<()> {
let name = name.as_ref();
// Public nonempty invariant: validate before any Database backend sees
// the call so local and remote Connections agree on InvalidInput.
if name.is_empty() {
return Err(Error::InvalidInput {
message: "function name removal name must be non-empty".into(),
});
}
self.internal.remove_function_name(name, current).await
}
/// Revoke an exact immutable [`Function`] by opaque id.
///
/// This is a direct synchronous administrator catalog set-bit, not a
/// [`crate::job::Job`], not catalog name removal, not physical deletion,
/// and not [`Function`] or generated-column mutation. The caller supplies
/// an already-validated exact [`Function`] handle; only [`Function::id`]
/// is sent on the wire.
///
/// Local/default backends return [`Error::NotSupported`]. Remote backends
/// complete only when the server reports durable success for that exact
/// id. Repeated logical calls that each receive success succeed; there is
/// no client-side already-revoked branch.
pub async fn revoke_function(&self, function: &Function) -> Result<()> {
self.internal.revoke_function(function).await
}
/// Drop a table in the database.
///
/// # Arguments
@@ -620,21 +644,6 @@ impl Connection {
.await
}
/// Start dropping a table and return a handle to the cleanup job.
///
/// The table may become unavailable before its physical data is removed.
/// Call [`crate::job::Job::wait`] to wait for cleanup to finish. Local
/// backends may complete the drop before returning the handle.
pub async fn drop_table_async(
&self,
name: impl AsRef<str>,
namespace_path: &[String],
) -> Result<crate::job::Job> {
self.internal
.drop_table_async(name.as_ref(), namespace_path)
.await
}
/// Drop the database
///
/// This is the same as dropping all of the tables
@@ -707,32 +716,9 @@ impl Connection {
self.internal.namespace_client_config().await
}
/// List the tables in the database, a page at a time
///
/// ```
/// # use lancedb::Connection;
/// # async fn list_all(conn: &Connection) -> Result<Vec<String>, lancedb::Error> {
/// let mut names = Vec::new();
/// let mut token = None;
/// loop {
/// let mut request = conn.list_tables().limit(100);
/// if let Some(token) = token {
/// request = request.page_token(token);
/// }
/// let page = request.execute().await?;
/// names.extend(page.tables);
/// // A page may be short without being the last one, so the token is what ends
/// // the walk.
/// token = page.page_token.filter(|token| !token.is_empty());
/// if token.is_none() {
/// break;
/// }
/// }
/// # Ok(names)
/// # }
/// ```
pub fn list_tables(&self) -> ListTablesBuilder {
ListTablesBuilder::new(self.internal.clone())
/// List tables with pagination support
pub async fn list_tables(&self, request: ListTablesRequest) -> Result<ListTablesResponse> {
self.internal.list_tables(request).await
}
/// Get the in-memory embedding registry.
@@ -1388,8 +1374,6 @@ mod test_utils {
}
#[cfg(test)]
// `table_names` is deprecated but still supported, so its tests still call it.
#[allow(deprecated)]
mod tests {
use arrow_schema::{DataType, Field, Schema};
use lance_testing::datagen::{BatchGenerator, IncrementingInt32};
@@ -1732,75 +1716,6 @@ mod tests {
assert_eq!(tables, names[..7]);
}
#[tokio::test]
async fn test_list_tables_paginates() {
let tc = new_test_connection().await.unwrap();
if tc.is_remote {
// What resumes a page is the server's to decide, and asserting it here would be
// asserting the server's contract rather than this one.
return;
}
let db = tc.connection;
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int32, false)]));
let mut names = Vec::with_capacity(25);
for _ in 0..25 {
let name = uuid::Uuid::new_v4().to_string();
names.push(name.clone());
db.create_empty_table(name, schema.clone())
.execute()
.await
.unwrap();
}
names.sort();
let page = db.list_tables().limit(10).execute().await.unwrap();
assert_eq!(page.tables, names[..10]);
// The token is opaque and is not a table name: it is whatever resumes the store
// the database sits on, so a caller checks that there is one, not what it says.
assert!(page.page_token.is_some());
// Walking in pages has to reach every table exactly once, with nothing lost
// at a page boundary.
let mut seen = Vec::with_capacity(names.len());
let mut page_token = None;
loop {
let mut request = db.list_tables().limit(10);
if let Some(token) = page_token {
request = request.page_token(token);
}
let page = request.execute().await.unwrap();
seen.extend(page.tables);
page_token = page.page_token.filter(|token| !token.is_empty());
if page_token.is_none() {
break;
}
}
assert_eq!(seen, names);
}
#[tokio::test]
async fn test_list_tables_exhausted_has_no_token() {
let tc = new_test_connection().await.unwrap();
let db = tc.connection;
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int32, false)]));
for i in 0..3 {
db.create_empty_table(format!("table{i}"), schema.clone())
.execute()
.await
.unwrap();
}
// A limit the listing does not fill leaves no token behind.
let page = db.list_tables().limit(10).execute().await.unwrap();
assert_eq!(page.tables.len(), 3);
assert_eq!(page.page_token, None);
// Neither does one that exactly exhausts it.
let page = db.list_tables().limit(3).execute().await.unwrap();
assert_eq!(page.tables.len(), 3);
assert_eq!(page.page_token, None);
}
#[tokio::test]
async fn test_open_table() {
let tc = new_test_connection().await.unwrap();
+1 -1
View File
@@ -439,7 +439,7 @@ mod tests {
.unwrap()
.data_storage_format
.lance_file_format();
// Compare resolved versions since Stable/Next are aliases that resolve at storage time
// Compare concrete stored format to the resolved requested alias.
assert_eq!(storage_format, data_storage_version.resolve());
}
+68 -12
View File
@@ -30,12 +30,16 @@ use lance_namespace::models::{
use crate::data::scannable::Scannable;
use crate::error::Result;
use crate::function::{Function, FunctionId, RegisterFunctionJobSpec};
use crate::table::{BaseTable, WriteOptions};
pub mod listing;
pub mod namespace;
pub(crate) mod read_freshness;
#[cfg(test)]
mod create_table_generated_column_schema_admission_contract;
pub trait DatabaseOptions {
fn serialize_into_map(&self, map: &mut HashMap<String, String>);
}
@@ -230,6 +234,12 @@ pub struct JobDescription {
pub creation_ms: i64,
/// The job-type-specific specification. Null when the server omits it.
pub spec: serde_json::Value,
/// Explicit success result from the describe envelope, when present.
///
/// Missing or JSON `null` wire `result` is [`None`]. An explicit
/// [`crate::JobResult::None`] object is `Some(JobResult::None)`. An exact
/// Function result is `Some(JobResult::Function(...))`.
pub result: Option<crate::job::JobResult>,
/// Why the job failed, when the job is failed and the server reports a
/// reason.
pub failure: Option<crate::error::JobFailure>,
@@ -311,6 +321,64 @@ pub trait Database:
async fn job_history(&self, _job_id: Option<&str>) -> Result<Vec<RecordBatch>> {
job_op_not_supported("job_history")
}
/// Submit a first-class Function registration job.
///
/// Returns a [`crate::job::Job`] handle for the accepted server-side job.
/// Local databases do not support registration.
async fn register_function(&self, _spec: RegisterFunctionJobSpec) -> Result<crate::job::Job> {
job_op_not_supported("register_function")
}
/// Look up the Function currently bound to a database-scoped name.
///
/// The name is lookup indirection only and is never part of the returned
/// [`Function`]. Empty names return [`crate::Error::InvalidInput`] before
/// the unsupported fallback so local and remote backends agree. Nonempty
/// names on databases without enterprise catalog lookup return
/// [`crate::Error::NotSupported`].
async fn lookup_function_by_name(&self, name: &str) -> Result<Function> {
// Public nonempty invariant on the Database trait seam itself:
// Connection::database() exposes Arc<dyn Database>, so empty-name
// rejection cannot rely solely on Connection prevalidation.
if name.is_empty() {
return Err(crate::error::Error::InvalidInput {
message: "function lookup name must be non-empty".into(),
});
}
job_op_not_supported("lookup_function_by_name")
}
/// Look up an immutable Function by exact opaque [`FunctionId`].
///
/// Exact-ID lookup is independent of later catalog name changes. Local
/// databases do not support enterprise catalog lookup.
async fn lookup_function_by_id(&self, _function_id: &FunctionId) -> Result<Function> {
job_op_not_supported("lookup_function_by_id")
}
/// Conditionally remove a database-scoped Function catalog name.
///
/// Direct synchronous catalog CAS, not a Job and not physical Function
/// deletion. Empty names return [`crate::Error::InvalidInput`] before the
/// unsupported fallback so local and remote backends agree. Nonempty names
/// on databases without enterprise catalog mutation return
/// [`crate::Error::NotSupported`].
async fn remove_function_name(&self, name: &str, _current: &Function) -> Result<()> {
// Public nonempty invariant on the Database trait seam itself:
// Connection::database() exposes Arc<dyn Database>, so empty-name
// rejection cannot rely solely on Connection prevalidation.
if name.is_empty() {
return Err(crate::error::Error::InvalidInput {
message: "function name removal name must be non-empty".into(),
});
}
job_op_not_supported("remove_function_name")
}
/// Revoke an exact immutable [`Function`] by opaque id.
///
/// Direct synchronous administrator catalog set-bit, not a Job, not name
/// removal, and not physical Function deletion. Databases without
/// enterprise catalog mutation return [`crate::Error::NotSupported`].
async fn revoke_function(&self, _function: &Function) -> Result<()> {
job_op_not_supported("revoke_function")
}
/// Open a table in the database
async fn open_table(&self, request: OpenTableRequest) -> Result<Arc<dyn BaseTable>>;
/// Rename a table in the database
@@ -323,18 +391,6 @@ pub trait Database:
) -> Result<()>;
/// Drop a table in the database
async fn drop_table(&self, name: &str, namespace_path: &[String]) -> Result<()>;
/// Start dropping a table and return a handle to the cleanup job.
///
/// Backends without asynchronous cleanup complete the drop before
/// returning an already-finished job.
async fn drop_table_async(
&self,
name: &str,
namespace_path: &[String],
) -> Result<crate::job::Job> {
self.drop_table(name, namespace_path).await?;
Ok(crate::job::Job::new_done())
}
/// Drop all tables in the database
async fn drop_all_tables(&self, namespace_path: &[String]) -> Result<()>;
fn as_any(&self) -> &dyn std::any::Any;
@@ -0,0 +1,788 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! RED runtime contract tests for create-table schema admission (B4g).
//!
//! Caller-authored Arrow field metadata under
//! [`crate::function::GENERATED_COLUMN_METADATA_KEY`] must not enter table
//! schema state through general-purpose `Database::create_table`. Generated
//! definitions are Job-owned. This module proves the missing admission guard
//! on Native listing, Native namespace, and Remote create paths.
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use arrow_array::{Int32Array, RecordBatch, StringArray};
use arrow_schema::{DataType, Field, Schema, SchemaRef};
use tempfile::TempDir;
use crate::arrow::SendableRecordBatchStream;
use crate::data::scannable::Scannable;
use crate::database::listing::ListingDatabase;
use crate::database::{CreateTableMode, CreateTableRequest, Database, TableNamesRequest};
use crate::error::Error;
use crate::function::{
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition,
};
const ID: &str = "id";
const ORDINARY: &str = "ordinary";
const GEN_OUT: &str = "gen_out";
const ORDINARY_META_KEY: &str = "unit";
const ORDINARY_META_VALUE: &str = "label";
const FN_ID: &str = "fn.exact.b4g.create_table.literal";
const MALFORMED_MARKER: &str = "SENSITIVE_B4G_CREATE_TABLE_METADATA_MARKER_9d2e_a7c1";
/// Counts [`Scannable::scan_as_stream`] calls. [`Scannable::schema`] is free.
struct ObservableScannable {
batch: RecordBatch,
scan_calls: Arc<AtomicUsize>,
}
impl ObservableScannable {
fn new(batch: RecordBatch, scan_calls: Arc<AtomicUsize>) -> Self {
Self { batch, scan_calls }
}
}
impl Scannable for ObservableScannable {
fn schema(&self) -> SchemaRef {
self.batch.schema()
}
fn scan_as_stream(&mut self) -> SendableRecordBatchStream {
self.scan_calls.fetch_add(1, Ordering::SeqCst);
self.batch.scan_as_stream()
}
fn num_rows(&self) -> Option<usize> {
Some(self.batch.num_rows())
}
fn rescannable(&self) -> bool {
true
}
}
fn literal_definition(output_field_id: i32) -> GeneratedColumnDefinition {
let function = Function::new(
FunctionId::try_new(FN_ID).unwrap(),
FunctionSignature::try_new(
vec![FunctionParameter::new("label", DataType::Utf8)],
FunctionOutput::new(DataType::Int32, true),
)
.unwrap(),
);
let call = FunctionCall::try_new(
&function,
vec![(
"label".to_string(),
FunctionArgument::try_literal(
Arc::new(StringArray::from(vec![Some("literal-only")])) as arrow_array::ArrayRef
)
.unwrap(),
)],
)
.unwrap();
GeneratedColumnDefinition::try_new(output_field_id, call, 1, 1).unwrap()
}
fn valid_reserved_payload() -> String {
literal_definition(1).to_metadata_json().unwrap()
}
fn malformed_reserved_payload() -> String {
format!(
r#"{{"format_version":1,"output_field_id":1,"function_call":"{MALFORMED_MARKER}","dependency_epoch":1,"materialized_epoch":1}}"#
)
}
fn batch_with_field_metadata(metadata: HashMap<String, String>) -> RecordBatch {
let gen_field = Field::new(GEN_OUT, DataType::Int32, true).with_metadata(metadata);
let schema = Arc::new(Schema::new(vec![
Field::new(ID, DataType::Int32, false),
Field::new(ORDINARY, DataType::Utf8, true),
gen_field,
]));
RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from(vec![1])),
Arc::new(StringArray::from(vec![Some("seed")])),
Arc::new(Int32Array::from(vec![10])),
],
)
.unwrap()
}
fn reserved_batch(payload: &str) -> RecordBatch {
batch_with_field_metadata(
[(
GENERATED_COLUMN_METADATA_KEY.to_string(),
payload.to_string(),
)]
.into(),
)
}
fn ordinary_metadata_batch() -> RecordBatch {
batch_with_field_metadata(
[(
ORDINARY_META_KEY.to_string(),
ORDINARY_META_VALUE.to_string(),
)]
.into(),
)
}
fn plain_seed_batch() -> RecordBatch {
batch_with_field_metadata(HashMap::new())
}
fn assert_not_supported_redacted(err: &Error, label: &str, forbidden_substrings: &[&str]) {
match err {
Error::NotSupported { message } => {
let rendered = format!("{err}\n{err:?}\n{message}");
assert!(
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
"{label}: leaked metadata wire key: {rendered}"
);
assert!(
!rendered.contains(FN_ID),
"{label}: leaked Function ID: {rendered}"
);
assert!(
!rendered.contains(MALFORMED_MARKER),
"{label}: leaked malformed marker: {rendered}"
);
for needle in forbidden_substrings {
assert!(
!rendered.contains(needle),
"{label}: leaked forbidden substring `{needle}`: {rendered}"
);
}
assert!(
message.to_lowercase().contains("generated")
|| message.to_lowercase().contains("job"),
"{label}: message must describe Job-owned generated-column boundary: {message}"
);
}
other => panic!("{label}: expected Error::NotSupported, got {other:?}"),
}
}
async fn listing_db() -> (TempDir, ListingDatabase) {
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap();
let request = crate::connection::ConnectRequest {
uri: uri.to_string(),
#[cfg(feature = "remote")]
client_config: Default::default(),
options: Default::default(),
namespace_client_properties: Default::default(),
manifest_enabled: false,
read_consistency_interval: None,
session: None,
};
let db = ListingDatabase::connect_with_options(&request)
.await
.unwrap();
(tmp, db)
}
fn listing_table_dir(tmp: &TempDir, name: &str) -> std::path::PathBuf {
tmp.path().join(format!("{name}.lance"))
}
async fn listing_create(
db: &ListingDatabase,
name: &str,
data: Box<dyn Scannable>,
mode: CreateTableMode,
) -> crate::Result<Arc<dyn crate::table::BaseTable>> {
db.create_table(CreateTableRequest {
name: name.to_string(),
namespace_path: vec![],
data,
mode,
write_options: Default::default(),
location: None,
namespace_client: None,
})
.await
}
async fn assert_listing_absent(db: &ListingDatabase, tmp: &TempDir, name: &str) {
#[allow(deprecated)]
let names = db.table_names(TableNamesRequest::default()).await.unwrap();
assert!(
!names.contains(&name.to_string()),
"rejected create must leave no listed table `{name}`; got {names:?}"
);
assert!(
!listing_table_dir(tmp, name).exists(),
"rejected create must leave no storage directory for `{name}`"
);
}
#[tokio::test]
async fn listing_create_rejects_reserved_generated_column_metadata_before_scan() {
let (tmp, db) = listing_db().await;
let payload = valid_reserved_payload();
let scan_calls = Arc::new(AtomicUsize::new(0));
let data = Box::new(ObservableScannable::new(
reserved_batch(&payload),
scan_calls.clone(),
)) as Box<dyn Scannable>;
let err = listing_create(&db, "b4g_listing_create", data, CreateTableMode::Create)
.await
.expect_err("listing Create must reject reserved generated-column metadata");
assert_not_supported_redacted(
&err,
"listing Create reserved admission",
&[payload.as_str()],
);
assert_eq!(
scan_calls.load(Ordering::SeqCst),
0,
"rejection must occur before Scannable::scan_as_stream"
);
assert_listing_absent(&db, &tmp, "b4g_listing_create").await;
}
#[tokio::test]
async fn listing_overwrite_rejects_reserved_generated_column_metadata_and_preserves_table() {
let (tmp, db) = listing_db().await;
let seed = listing_create(
&db,
"b4g_listing_overwrite",
Box::new(plain_seed_batch()) as Box<dyn Scannable>,
CreateTableMode::Create,
)
.await
.unwrap();
let version_before = seed.version().await.unwrap();
let schema_before = seed.schema().await.unwrap();
assert!(
!schema_before
.field_with_name(GEN_OUT)
.unwrap()
.metadata()
.contains_key(GENERATED_COLUMN_METADATA_KEY)
);
let payload = malformed_reserved_payload();
assert!(payload.contains(MALFORMED_MARKER));
let scan_calls = Arc::new(AtomicUsize::new(0));
let data = Box::new(ObservableScannable::new(
reserved_batch(&payload),
scan_calls.clone(),
)) as Box<dyn Scannable>;
let err = listing_create(
&db,
"b4g_listing_overwrite",
data,
CreateTableMode::Overwrite,
)
.await
.expect_err("listing Overwrite must reject reserved generated-column metadata");
assert_not_supported_redacted(
&err,
"listing Overwrite reserved admission",
&[payload.as_str()],
);
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
let reopened = db
.open_table(crate::database::OpenTableRequest {
name: "b4g_listing_overwrite".to_string(),
namespace_path: vec![],
index_cache_size: None,
lance_read_params: None,
location: None,
namespace_client: None,
managed_versioning: None,
})
.await
.unwrap();
assert_eq!(reopened.version().await.unwrap(), version_before);
let schema_after = reopened.schema().await.unwrap();
assert_eq!(schema_after.as_ref(), schema_before.as_ref());
assert!(
!schema_after
.field_with_name(GEN_OUT)
.unwrap()
.metadata()
.contains_key(GENERATED_COLUMN_METADATA_KEY)
);
assert!(listing_table_dir(&tmp, "b4g_listing_overwrite").exists());
}
#[tokio::test]
async fn listing_exist_ok_absent_rejects_reserved_generated_column_metadata_before_scan() {
let (tmp, db) = listing_db().await;
let payload = valid_reserved_payload();
let scan_calls = Arc::new(AtomicUsize::new(0));
let data = Box::new(ObservableScannable::new(
reserved_batch(&payload),
scan_calls.clone(),
)) as Box<dyn Scannable>;
let err = listing_create(
&db,
"b4g_listing_exist_ok",
data,
CreateTableMode::exist_ok(|req| req),
)
.await
.expect_err("listing ExistOk (absent) must reject reserved generated-column metadata");
assert_not_supported_redacted(
&err,
"listing ExistOk absent reserved admission",
&[payload.as_str()],
);
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
assert_listing_absent(&db, &tmp, "b4g_listing_exist_ok").await;
}
#[tokio::test]
async fn listing_ordinary_field_metadata_is_accepted_and_preserved() {
let (_tmp, db) = listing_db().await;
let scan_calls = Arc::new(AtomicUsize::new(0));
let data = Box::new(ObservableScannable::new(
ordinary_metadata_batch(),
scan_calls.clone(),
)) as Box<dyn Scannable>;
let table = listing_create(&db, "b4g_listing_ordinary", data, CreateTableMode::Create)
.await
.expect("ordinary field metadata must remain accepted");
assert!(
scan_calls.load(Ordering::SeqCst) > 0,
"successful create may consume the Scannable"
);
let schema = table.schema().await.unwrap();
let md = schema.field_with_name(GEN_OUT).unwrap().metadata();
assert_eq!(
md.get(ORDINARY_META_KEY).map(String::as_str),
Some(ORDINARY_META_VALUE)
);
assert!(!md.contains_key(GENERATED_COLUMN_METADATA_KEY));
}
#[cfg(not(windows))] // directory namespace tests are unix-only in this crate
mod namespace_admission {
use super::*;
use crate::connect_namespace;
use lance_namespace::models::{CreateNamespaceRequest, DescribeTableRequest};
async fn namespace_conn() -> (TempDir, crate::Connection) {
let tmp = tempfile::tempdir().unwrap();
let root = tmp.path().to_str().unwrap().to_string();
let mut properties = HashMap::new();
properties.insert("root".to_string(), root);
let conn = connect_namespace("dir", properties)
.execute()
.await
.unwrap();
conn.create_namespace(CreateNamespaceRequest {
id: Some(vec!["b4g_ns".into()]),
..Default::default()
})
.await
.unwrap();
(tmp, conn)
}
async fn assert_namespace_undeclared(conn: &crate::Connection, name: &str) {
let names = conn
.table_names()
.namespace(vec!["b4g_ns".into()])
.execute()
.await
.unwrap();
assert!(
!names.contains(&name.to_string()),
"rejected namespace create must leave no declared/listed table `{name}`; got {names:?}"
);
let ns = conn.namespace_client().await.unwrap();
let describe = ns
.describe_table(DescribeTableRequest {
id: Some(vec!["b4g_ns".into(), name.into()]),
..Default::default()
})
.await;
assert!(
describe.is_err(),
"rejected namespace create must leave no describable table `{name}`"
);
}
async fn namespace_create(
conn: &crate::Connection,
name: &str,
data: Box<dyn Scannable>,
mode: CreateTableMode,
) -> crate::Result<Arc<dyn crate::table::BaseTable>> {
conn.database()
.create_table(CreateTableRequest {
name: name.to_string(),
namespace_path: vec!["b4g_ns".into()],
data,
mode,
write_options: Default::default(),
location: None,
namespace_client: None,
})
.await
}
#[tokio::test]
async fn namespace_create_rejects_reserved_before_declare_describe_or_storage() {
let (_tmp, conn) = namespace_conn().await;
let payload = valid_reserved_payload();
let scan_calls = Arc::new(AtomicUsize::new(0));
let data = Box::new(ObservableScannable::new(
reserved_batch(&payload),
scan_calls.clone(),
)) as Box<dyn Scannable>;
let err = namespace_create(&conn, "b4g_ns_create", data, CreateTableMode::Create)
.await
.expect_err("namespace Create must reject reserved generated-column metadata");
assert_not_supported_redacted(
&err,
"namespace Create reserved admission",
&[payload.as_str()],
);
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
assert_namespace_undeclared(&conn, "b4g_ns_create").await;
}
#[tokio::test]
async fn namespace_overwrite_rejects_reserved_before_declare_describe_or_storage() {
let (_tmp, conn) = namespace_conn().await;
let seed = namespace_create(
&conn,
"b4g_ns_overwrite",
Box::new(plain_seed_batch()) as Box<dyn Scannable>,
CreateTableMode::Create,
)
.await
.unwrap();
let version_before = seed.version().await.unwrap();
let schema_before = seed.schema().await.unwrap();
let payload = malformed_reserved_payload();
let scan_calls = Arc::new(AtomicUsize::new(0));
let data = Box::new(ObservableScannable::new(
reserved_batch(&payload),
scan_calls.clone(),
)) as Box<dyn Scannable>;
let err = namespace_create(&conn, "b4g_ns_overwrite", data, CreateTableMode::Overwrite)
.await
.expect_err("namespace Overwrite must reject reserved generated-column metadata");
assert_not_supported_redacted(
&err,
"namespace Overwrite reserved admission",
&[payload.as_str()],
);
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
let reopened = conn
.database()
.open_table(crate::database::OpenTableRequest {
name: "b4g_ns_overwrite".to_string(),
namespace_path: vec!["b4g_ns".into()],
index_cache_size: None,
lance_read_params: None,
location: None,
namespace_client: None,
managed_versioning: None,
})
.await
.unwrap();
assert_eq!(reopened.version().await.unwrap(), version_before);
assert_eq!(
reopened.schema().await.unwrap().as_ref(),
schema_before.as_ref()
);
assert!(
!reopened
.schema()
.await
.unwrap()
.field_with_name(GEN_OUT)
.unwrap()
.metadata()
.contains_key(GENERATED_COLUMN_METADATA_KEY)
);
}
#[tokio::test]
async fn namespace_exist_ok_absent_rejects_reserved_before_declare_describe_or_storage() {
let (_tmp, conn) = namespace_conn().await;
let payload = valid_reserved_payload();
let scan_calls = Arc::new(AtomicUsize::new(0));
let data = Box::new(ObservableScannable::new(
reserved_batch(&payload),
scan_calls.clone(),
)) as Box<dyn Scannable>;
let err = namespace_create(
&conn,
"b4g_ns_exist_ok",
data,
CreateTableMode::exist_ok(|req| req),
)
.await
.expect_err("namespace ExistOk (absent) must reject reserved generated-column metadata");
assert_not_supported_redacted(
&err,
"namespace ExistOk absent reserved admission",
&[payload.as_str()],
);
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
assert_namespace_undeclared(&conn, "b4g_ns_exist_ok").await;
}
#[tokio::test]
async fn namespace_exist_ok_existing_rejects_reserved_even_when_mode_would_ignore_data() {
let (_tmp, conn) = namespace_conn().await;
let seed = namespace_create(
&conn,
"b4g_ns_exist_ok_existing",
Box::new(plain_seed_batch()) as Box<dyn Scannable>,
CreateTableMode::Create,
)
.await
.unwrap();
let version_before = seed.version().await.unwrap();
let schema_before = seed.schema().await.unwrap();
let payload = valid_reserved_payload();
let scan_calls = Arc::new(AtomicUsize::new(0));
let data = Box::new(ObservableScannable::new(
reserved_batch(&payload),
scan_calls.clone(),
)) as Box<dyn Scannable>;
let err = namespace_create(
&conn,
"b4g_ns_exist_ok_existing",
data,
CreateTableMode::exist_ok(|req| req),
)
.await
.expect_err(
"namespace ExistOk must not accept reserved metadata merely because data is ignored",
);
assert_not_supported_redacted(
&err,
"namespace ExistOk existing reserved admission",
&[payload.as_str()],
);
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
let reopened = conn
.database()
.open_table(crate::database::OpenTableRequest {
name: "b4g_ns_exist_ok_existing".to_string(),
namespace_path: vec!["b4g_ns".into()],
index_cache_size: None,
lance_read_params: None,
location: None,
namespace_client: None,
managed_versioning: None,
})
.await
.unwrap();
assert_eq!(reopened.version().await.unwrap(), version_before);
assert_eq!(
reopened.schema().await.unwrap().as_ref(),
schema_before.as_ref()
);
}
}
#[cfg(feature = "remote")]
mod remote_admission {
use super::*;
use std::io::Cursor;
use arrow_ipc::reader::StreamReader;
use async_trait::async_trait;
use crate::Connection;
use crate::remote::{ClientConfig, HeaderProvider};
#[derive(Debug)]
struct CountingHeaderProvider {
calls: Arc<AtomicUsize>,
}
#[async_trait]
impl HeaderProvider for CountingHeaderProvider {
async fn get_headers(&self) -> crate::Result<HashMap<String, String>> {
self.calls.fetch_add(1, Ordering::SeqCst);
Ok(HashMap::from([(
"X-B4g-Test".to_string(),
"must-not-be-requested".to_string(),
)]))
}
}
fn counting_handler(
calls: Arc<AtomicUsize>,
) -> impl Fn(reqwest::Request) -> http::Response<String> + Clone + Send + Sync + 'static {
move |_request| {
calls.fetch_add(1, Ordering::SeqCst);
http::Response::builder()
.status(200)
.body(String::new())
.unwrap()
}
}
async fn remote_create(
conn: &Connection,
name: &str,
data: Box<dyn Scannable>,
mode: CreateTableMode,
) -> crate::Result<Arc<dyn crate::table::BaseTable>> {
// Direct Database trait path used by Connection::create_table.
conn.database()
.create_table(CreateTableRequest {
name: name.to_string(),
namespace_path: vec![],
data,
mode,
write_options: Default::default(),
location: None,
namespace_client: None,
})
.await
}
async fn assert_remote_rejects(
mode: CreateTableMode,
table_name: &str,
payload: &str,
label: &str,
) {
let handler_calls = Arc::new(AtomicUsize::new(0));
let header_calls = Arc::new(AtomicUsize::new(0));
let scan_calls = Arc::new(AtomicUsize::new(0));
let config = ClientConfig {
header_provider: Some(Arc::new(CountingHeaderProvider {
calls: header_calls.clone(),
}) as Arc<dyn HeaderProvider>),
..Default::default()
};
let conn = Connection::new_with_handler_and_config(
counting_handler(handler_calls.clone()),
config,
);
let data = Box::new(ObservableScannable::new(
reserved_batch(payload),
scan_calls.clone(),
)) as Box<dyn Scannable>;
let err = remote_create(&conn, table_name, data, mode)
.await
.expect_err("remote create must reject reserved generated-column metadata");
assert_not_supported_redacted(&err, label, &[payload]);
assert_eq!(
scan_calls.load(Ordering::SeqCst),
0,
"{label}: rejection must occur before scan_as_stream"
);
assert_eq!(
header_calls.load(Ordering::SeqCst),
0,
"{label}: rejection must occur before header-provider invocation"
);
assert_eq!(
handler_calls.load(Ordering::SeqCst),
0,
"{label}: rejection must occur before HTTP handler"
);
}
#[tokio::test]
async fn remote_create_rejects_reserved_before_scan_headers_and_http() {
assert_remote_rejects(
CreateTableMode::Create,
"b4g_remote_create",
&valid_reserved_payload(),
"remote Create reserved admission",
)
.await;
}
#[tokio::test]
async fn remote_overwrite_rejects_reserved_before_scan_headers_and_http() {
assert_remote_rejects(
CreateTableMode::Overwrite,
"b4g_remote_overwrite",
&malformed_reserved_payload(),
"remote Overwrite reserved admission",
)
.await;
}
#[tokio::test]
async fn remote_exist_ok_rejects_reserved_before_scan_headers_and_http() {
assert_remote_rejects(
CreateTableMode::exist_ok(|req| req),
"b4g_remote_exist_ok",
&valid_reserved_payload(),
"remote ExistOk reserved admission",
)
.await;
}
#[tokio::test]
async fn remote_ordinary_field_metadata_is_transmitted_unchanged() {
let conn = Connection::new_with_handler(|request| {
assert_eq!(request.method(), &reqwest::Method::POST);
assert_eq!(
request.url().path(),
"/v1/table/b4g_remote_ordinary/create/"
);
let body = request
.body()
.and_then(|b| b.as_bytes())
.expect("ordinary create must send an Arrow IPC body");
let reader = StreamReader::try_new(Cursor::new(body), None).unwrap();
let schema = reader.schema();
let md = schema.field_with_name(GEN_OUT).unwrap().metadata();
assert_eq!(
md.get(ORDINARY_META_KEY).map(String::as_str),
Some(ORDINARY_META_VALUE),
"ordinary field metadata must be transmitted unchanged"
);
assert!(!md.contains_key(GENERATED_COLUMN_METADATA_KEY));
// Consume stream to completion for a well-formed IPC body.
for batch in reader {
batch.unwrap();
}
http::Response::builder()
.status(200)
.body(String::new())
.unwrap()
});
conn.create_table("b4g_remote_ordinary", ordinary_metadata_batch())
.mode(CreateTableMode::Create)
.execute()
.await
.expect("ordinary field metadata must remain accepted on remote create");
}
}

Some files were not shown because too many files have changed in this diff Show More