mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-30 18:08:24 +00:00
Compare commits
29 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| b7dbef971e | |||
| 71202dd3e6 | |||
| 0ab435f2d7 | |||
| fcc6a89b92 | |||
| 27cea03b7d | |||
| f1c4967eeb | |||
| 11c1d81638 | |||
| f6efdc9e9f | |||
| cdebea118d | |||
| 76942306b7 | |||
| d742b174c4 | |||
| a075aa62f8 | |||
| 040a4120c8 | |||
| 928c3dde2d | |||
| 980818df26 | |||
| c429863122 | |||
| fc0d917d32 | |||
| def869bb78 | |||
| 9e4d8bd1c7 | |||
| 4148dfef72 | |||
| 0ac70a8b9f | |||
| 91c5f344d2 | |||
| ffd35c1a8f | |||
| 790d0c684c | |||
| 251f194696 | |||
| 4b7325bd74 | |||
| 1d75638dea | |||
| 031c3585a8 | |||
| 6fb976cf89 |
+1
-1
@@ -1,5 +1,5 @@
|
||||
[tool.bumpversion]
|
||||
current_version = "0.37.1-beta.1"
|
||||
current_version = "0.38.0-beta.2"
|
||||
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. If omitted, the skill will use the latest Lance release that needs an update."
|
||||
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."
|
||||
required: false
|
||||
default: ""
|
||||
type: string
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
tag:
|
||||
description: "Tag name from Lance. Leave empty to use the latest Lance release that needs an update."
|
||||
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."
|
||||
required: false
|
||||
default: ""
|
||||
type: string
|
||||
|
||||
@@ -69,6 +69,16 @@ 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
+54
-68
@@ -959,7 +959,7 @@ dependencies = [
|
||||
"aws-smithy-runtime-api",
|
||||
"aws-smithy-types",
|
||||
"h2 0.3.27",
|
||||
"h2 0.4.14",
|
||||
"h2 0.4.16",
|
||||
"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.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"rand 0.9.5",
|
||||
@@ -3877,9 +3877,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "h2"
|
||||
version = "0.4.14"
|
||||
version = "0.4.16"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "171fefbc92fe4a4de27e0698d6a5b392d6a0e333506bc49133760b3bcf948733"
|
||||
checksum = "a9f37a958b41b3b19ee2707c06439c0e9e547e847223eb791ecb0cb821c65e27"
|
||||
dependencies = [
|
||||
"atomic-waker",
|
||||
"bytes",
|
||||
@@ -4188,7 +4188,7 @@ dependencies = [
|
||||
"bytes",
|
||||
"futures-channel",
|
||||
"futures-core",
|
||||
"h2 0.4.14",
|
||||
"h2 0.4.16",
|
||||
"http 1.5.0",
|
||||
"http-body 1.1.0",
|
||||
"httparse",
|
||||
@@ -4815,8 +4815,8 @@ checksum = "e037a2e1d8d5fdbd49b16a4ea09d5d6401c1f29eca5ff29d03d3824dba16256a"
|
||||
|
||||
[[package]]
|
||||
name = "lance"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arc-swap",
|
||||
"arrow",
|
||||
@@ -4832,7 +4832,6 @@ dependencies = [
|
||||
"async-recursion",
|
||||
"async-trait",
|
||||
"async_cell",
|
||||
"aws-credential-types",
|
||||
"aws-sdk-dynamodb",
|
||||
"byteorder",
|
||||
"bytes",
|
||||
@@ -4848,7 +4847,6 @@ dependencies = [
|
||||
"either",
|
||||
"fst",
|
||||
"futures",
|
||||
"half",
|
||||
"humantime",
|
||||
"itertools 0.14.0",
|
||||
"lance-arrow",
|
||||
@@ -4890,8 +4888,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-arrow"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -4913,7 +4911,7 @@ dependencies = [
|
||||
[[package]]
|
||||
name = "lance-arrow-scalar"
|
||||
version = "58.0.0"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -4927,7 +4925,7 @@ dependencies = [
|
||||
[[package]]
|
||||
name = "lance-arrow-stats"
|
||||
version = "58.0.0"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-schema",
|
||||
@@ -4936,8 +4934,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-bitpacking"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrayref",
|
||||
"crunchy",
|
||||
@@ -4947,8 +4945,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-core"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -4956,12 +4954,10 @@ dependencies = [
|
||||
"arrow-schema",
|
||||
"async-trait",
|
||||
"blake3",
|
||||
"byteorder",
|
||||
"bytes",
|
||||
"datafusion-common",
|
||||
"datafusion-sql",
|
||||
"futures",
|
||||
"itertools 0.14.0",
|
||||
"lance-arrow",
|
||||
"lance-derive",
|
||||
"libc",
|
||||
@@ -4979,7 +4975,6 @@ dependencies = [
|
||||
"snafu 0.9.0",
|
||||
"tempfile",
|
||||
"tokio",
|
||||
"tokio-stream",
|
||||
"tokio-util",
|
||||
"tracing",
|
||||
"twox-hash",
|
||||
@@ -4988,8 +4983,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-datafusion"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-array",
|
||||
@@ -5008,7 +5003,6 @@ dependencies = [
|
||||
"jsonb",
|
||||
"lance-arrow",
|
||||
"lance-core",
|
||||
"lance-datagen",
|
||||
"log",
|
||||
"pin-project",
|
||||
"prost",
|
||||
@@ -5019,8 +5013,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-datagen"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-array",
|
||||
@@ -5037,8 +5031,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-derive"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
@@ -5047,8 +5041,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-encoding"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow-arith",
|
||||
"arrow-array",
|
||||
@@ -5073,7 +5067,6 @@ dependencies = [
|
||||
"num-traits",
|
||||
"prost",
|
||||
"prost-build",
|
||||
"rand 0.9.5",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"xxhash-rust",
|
||||
@@ -5082,8 +5075,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-file"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow-arith",
|
||||
"arrow-array",
|
||||
@@ -5114,8 +5107,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-index"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arc-swap",
|
||||
"arrow",
|
||||
@@ -5130,7 +5123,6 @@ dependencies = [
|
||||
"async-trait",
|
||||
"bitvec",
|
||||
"bytes",
|
||||
"chrono",
|
||||
"crossbeam-queue",
|
||||
"datafusion",
|
||||
"datafusion-common",
|
||||
@@ -5148,7 +5140,6 @@ dependencies = [
|
||||
"lance-bitpacking",
|
||||
"lance-core",
|
||||
"lance-datafusion",
|
||||
"lance-datagen",
|
||||
"lance-encoding",
|
||||
"lance-file",
|
||||
"lance-index-core",
|
||||
@@ -5177,13 +5168,12 @@ dependencies = [
|
||||
"tempfile",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "lance-index-core"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-schema",
|
||||
@@ -5205,8 +5195,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-io"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-array",
|
||||
@@ -5220,7 +5210,6 @@ dependencies = [
|
||||
"futures",
|
||||
"http 1.5.0",
|
||||
"io-uring",
|
||||
"lance-arrow",
|
||||
"lance-core",
|
||||
"lance-namespace",
|
||||
"log",
|
||||
@@ -5238,29 +5227,28 @@ dependencies = [
|
||||
"tokio",
|
||||
"tracing",
|
||||
"url",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "lance-linalg"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
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.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"async-trait",
|
||||
@@ -5272,8 +5260,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-namespace-impls"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-ipc",
|
||||
@@ -5312,9 +5300,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-namespace-reqwest-client"
|
||||
version = "0.8.6"
|
||||
version = "0.11.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ba3f0a235e3ed5f8805205649ccc7d7d0f3df23ce1294242c9265ad488d7f19d"
|
||||
checksum = "0a030196da1c994b63a96a4f0bf5b0cfa459fe6dadc9e962320246ca328da22a"
|
||||
dependencies = [
|
||||
"reqwest 0.12.28",
|
||||
"serde",
|
||||
@@ -5326,14 +5314,13 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-select"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
"arrow-schema",
|
||||
"byteorder",
|
||||
"bytes",
|
||||
"itertools 0.14.0",
|
||||
"lance-core",
|
||||
"roaring",
|
||||
@@ -5342,8 +5329,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-table"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-array",
|
||||
@@ -5383,8 +5370,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-testing"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-schema",
|
||||
@@ -5397,8 +5384,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-tokenizer"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
version = "11.0.0-beta.15"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=d7b1d570461c6d2adde8f3a84ae88db4823c726f#d7b1d570461c6d2adde8f3a84ae88db4823c726f"
|
||||
dependencies = [
|
||||
"frostem",
|
||||
"icu_segmenter",
|
||||
@@ -5411,7 +5398,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lancedb"
|
||||
version = "0.37.1-beta.1"
|
||||
version = "0.38.0-beta.2"
|
||||
dependencies = [
|
||||
"ahash",
|
||||
"anyhow",
|
||||
@@ -5432,7 +5419,6 @@ dependencies = [
|
||||
"aws-sdk-kms",
|
||||
"aws-sdk-s3",
|
||||
"aws-smithy-runtime",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"candle-core",
|
||||
"candle-nn",
|
||||
@@ -5500,7 +5486,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lancedb-nodejs"
|
||||
version = "0.37.1-beta.1"
|
||||
version = "0.38.0-beta.2"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -5525,7 +5511,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lancedb-python"
|
||||
version = "0.37.1-beta.1"
|
||||
version = "0.38.0-beta.2"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"async-trait",
|
||||
@@ -8440,7 +8426,7 @@ dependencies = [
|
||||
"encoding_rs",
|
||||
"futures-core",
|
||||
"futures-util",
|
||||
"h2 0.4.14",
|
||||
"h2 0.4.16",
|
||||
"http 1.5.0",
|
||||
"http-body 1.1.0",
|
||||
"http-body-util",
|
||||
@@ -10096,7 +10082,7 @@ dependencies = [
|
||||
"async-trait",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"h2 0.4.14",
|
||||
"h2 0.4.16",
|
||||
"http 1.5.0",
|
||||
"http-body 1.1.0",
|
||||
"http-body-util",
|
||||
|
||||
+17
-14
@@ -13,20 +13,23 @@ categories = ["database-implementations"]
|
||||
rust-version = "1.91.0"
|
||||
|
||||
[workspace.dependencies]
|
||||
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" }
|
||||
# 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" }
|
||||
ahash = "0.8"
|
||||
# Note that this one does not include pyarrow
|
||||
arrow = { version = "58.0.0", optional = false }
|
||||
|
||||
@@ -101,6 +101,19 @@ 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" },
|
||||
]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
+33
-1
@@ -14,7 +14,7 @@ Add the following dependency to your `pom.xml`:
|
||||
<dependency>
|
||||
<groupId>com.lancedb</groupId>
|
||||
<artifactId>lancedb-core</artifactId>
|
||||
<version>0.37.1-beta.1</version>
|
||||
<version>0.38.0-beta.2</version>
|
||||
</dependency>
|
||||
```
|
||||
|
||||
@@ -55,6 +55,38 @@ 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
|
||||
|
||||
@@ -386,6 +386,29 @@ 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`<[`Job`](Job.md)>
|
||||
|
||||
***
|
||||
|
||||
### getJob()
|
||||
|
||||
```ts
|
||||
@@ -506,6 +529,71 @@ 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`<[`ListTablesOptions`](../interfaces/ListTablesOptions.md)>
|
||||
Pagination options
|
||||
(`pageToken`, `limit`).
|
||||
|
||||
##### Returns
|
||||
|
||||
`Promise`<[`ListTablesResponse`](../interfaces/ListTablesResponse.md)>
|
||||
|
||||
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`<[`ListTablesOptions`](../interfaces/ListTablesOptions.md)>
|
||||
Pagination options
|
||||
(`pageToken`, `limit`).
|
||||
|
||||
##### Returns
|
||||
|
||||
`Promise`<[`ListTablesResponse`](../interfaces/ListTablesResponse.md)>
|
||||
|
||||
Table names and an optional token
|
||||
for fetching the next page.
|
||||
|
||||
***
|
||||
|
||||
### openTable()
|
||||
|
||||
```ts
|
||||
@@ -567,7 +655,7 @@ a "not supported" error.
|
||||
|
||||
***
|
||||
|
||||
### tableNames()
|
||||
### ~~tableNames()~~
|
||||
|
||||
#### tableNames(options)
|
||||
|
||||
@@ -589,6 +677,10 @@ Tables will be returned in lexicographical order.
|
||||
|
||||
`Promise`<`string`[]>
|
||||
|
||||
##### Deprecated
|
||||
|
||||
Use [Connection.listTables](Connection.md#listtables) instead.
|
||||
|
||||
#### tableNames(namespacePath, options)
|
||||
|
||||
```ts
|
||||
@@ -611,3 +703,7 @@ Tables will be returned in lexicographical order.
|
||||
##### Returns
|
||||
|
||||
`Promise`<`string`[]>
|
||||
|
||||
##### Deprecated
|
||||
|
||||
Use [Connection.listTables](Connection.md#listtables) instead.
|
||||
|
||||
@@ -69,14 +69,34 @@ 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`<`any`> \| `Field`<`any`>[] \| `Schema`<`any`> \| [`AddColumnsSql`](../interfaces/AddColumnsSql.md)[]
|
||||
* **newColumnTransforms**:
|
||||
\| `Field`<`any`>
|
||||
\| `Field`<`any`>[]
|
||||
\| `Schema`<`any`>
|
||||
\| [`AddColumnsSql`](../interfaces/AddColumnsSql.md)[]
|
||||
\| `object`
|
||||
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
|
||||
|
||||
@@ -85,6 +105,13 @@ Add new columns with defined values.
|
||||
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()
|
||||
@@ -186,6 +213,39 @@ 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`<`void`>
|
||||
|
||||
#### Example
|
||||
|
||||
```ts
|
||||
const before = await table.getLsmStats();
|
||||
await table.checkpointLsm();
|
||||
const after = await table.getLsmStats();
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### close()
|
||||
|
||||
```ts
|
||||
@@ -223,6 +283,24 @@ 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`<`void`>
|
||||
|
||||
***
|
||||
|
||||
### countRows()
|
||||
|
||||
```ts
|
||||
@@ -421,6 +499,48 @@ 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`<`void`>
|
||||
|
||||
***
|
||||
|
||||
### 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`<`undefined` \| [`LsmStats`](../interfaces/LsmStats.md)>
|
||||
|
||||
***
|
||||
|
||||
### getLsmWriteSpec()
|
||||
|
||||
```ts
|
||||
@@ -718,6 +838,67 @@ 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`<[`RefreshColumnResult`](../interfaces/RefreshColumnResult.md)>
|
||||
|
||||
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`<[`Job`](Job.md)>
|
||||
|
||||
#### Example
|
||||
|
||||
```ts
|
||||
const job = await table.refreshColumnAsync("doubled");
|
||||
await job.wait();
|
||||
console.log(await job.status()); // "finished"
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### restore()
|
||||
|
||||
```ts
|
||||
|
||||
@@ -58,6 +58,7 @@
|
||||
- [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)
|
||||
@@ -81,6 +82,7 @@
|
||||
- [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)
|
||||
@@ -94,7 +96,11 @@
|
||||
- [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)
|
||||
@@ -105,6 +111,7 @@
|
||||
- [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)
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
[**@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.
|
||||
@@ -0,0 +1,40 @@
|
||||
[**@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.
|
||||
@@ -0,0 +1,33 @@
|
||||
[**@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.
|
||||
@@ -0,0 +1,23 @@
|
||||
[**@lancedb/lancedb**](../README.md) • **Docs**
|
||||
|
||||
***
|
||||
|
||||
[@lancedb/lancedb](../globals.md) / ListTablesResponse
|
||||
|
||||
# Interface: ListTablesResponse
|
||||
|
||||
## Properties
|
||||
|
||||
### pageToken?
|
||||
|
||||
```ts
|
||||
optional pageToken: string;
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### tables
|
||||
|
||||
```ts
|
||||
tables: string[];
|
||||
```
|
||||
@@ -0,0 +1,22 @@
|
||||
[**@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.
|
||||
@@ -0,0 +1,60 @@
|
||||
[**@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.
|
||||
@@ -0,0 +1,23 @@
|
||||
[**@lancedb/lancedb**](../README.md) • **Docs**
|
||||
|
||||
***
|
||||
|
||||
[@lancedb/lancedb](../globals.md) / RefreshColumnResult
|
||||
|
||||
# Interface: RefreshColumnResult
|
||||
|
||||
## Properties
|
||||
|
||||
### rowsFilled
|
||||
|
||||
```ts
|
||||
rowsFilled: number;
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### version
|
||||
|
||||
```ts
|
||||
version: number;
|
||||
```
|
||||
@@ -4,11 +4,16 @@
|
||||
|
||||
[@lancedb/lancedb](../globals.md) / TableNamesOptions
|
||||
|
||||
# Interface: TableNamesOptions
|
||||
# Interface: ~~TableNamesOptions~~
|
||||
|
||||
## Deprecated
|
||||
|
||||
Use [ListTablesOptions](ListTablesOptions.md) with [Connection.listTables](../classes/Connection.md#listtables)
|
||||
instead.
|
||||
|
||||
## Properties
|
||||
|
||||
### limit?
|
||||
### ~~limit?~~
|
||||
|
||||
```ts
|
||||
optional limit: number;
|
||||
@@ -18,7 +23,7 @@ An optional limit to the number of results to return.
|
||||
|
||||
***
|
||||
|
||||
### startAfter?
|
||||
### ~~startAfter?~~
|
||||
|
||||
```ts
|
||||
optional startAfter: string;
|
||||
|
||||
@@ -52,6 +52,8 @@ listing a storage directory.
|
||||
|
||||
::: lancedb.table.Branches
|
||||
|
||||
::: lancedb.LsmWriteSpec
|
||||
|
||||
## Expressions
|
||||
|
||||
Type-safe expression builder for filters and projections. Use these instead
|
||||
|
||||
@@ -29,6 +29,48 @@ 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:
|
||||
|
||||
@@ -8,7 +8,7 @@
|
||||
<parent>
|
||||
<groupId>com.lancedb</groupId>
|
||||
<artifactId>lancedb-parent</artifactId>
|
||||
<version>0.37.1-beta.1</version>
|
||||
<version>0.38.0-beta.2</version>
|
||||
<relativePath>../pom.xml</relativePath>
|
||||
</parent>
|
||||
|
||||
@@ -33,6 +33,20 @@
|
||||
<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>
|
||||
|
||||
@@ -0,0 +1,194 @@
|
||||
/*
|
||||
* 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
|
||||
+ "}";
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,64 @@
|
||||
/*
|
||||
* 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 + "}";
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
/*
|
||||
* 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,29 +136,48 @@ public class LanceDbNamespaceClientBuilder {
|
||||
* @throws IllegalStateException if required parameters are missing
|
||||
*/
|
||||
public LanceNamespace build() {
|
||||
// Validate required fields
|
||||
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() {
|
||||
if (apiKey == null) {
|
||||
throw new IllegalStateException("API key is required");
|
||||
}
|
||||
if (database == null) {
|
||||
throw new IllegalStateException("Database is required");
|
||||
}
|
||||
}
|
||||
|
||||
// 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;
|
||||
/** The custom endpoint when set, else the LanceDB Cloud URL for this database and region. */
|
||||
private String resolveUri() {
|
||||
if (endpoint.isPresent()) {
|
||||
uri = endpoint.get();
|
||||
} else {
|
||||
String effectiveRegion = region.orElse(DEFAULT_REGION);
|
||||
uri = String.format(CLOUD_URL_PATTERN, database, effectiveRegion);
|
||||
return endpoint.get();
|
||||
}
|
||||
config.put("uri", uri);
|
||||
|
||||
return LanceNamespace.connect("rest", config, null);
|
||||
return String.format(CLOUD_URL_PATTERN, database, region.orElse(DEFAULT_REGION));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
/*
|
||||
* 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;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,394 @@
|
||||
/*
|
||||
* 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();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
/*
|
||||
* 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 + "}";
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,260 @@
|
||||
/*
|
||||
* 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
|
||||
+ "}";
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
/*
|
||||
* 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
|
||||
+ "}";
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,570 @@
|
||||
/*
|
||||
* 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
@@ -6,7 +6,7 @@
|
||||
|
||||
<groupId>com.lancedb</groupId>
|
||||
<artifactId>lancedb-parent</artifactId>
|
||||
<version>0.37.1-beta.1</version>
|
||||
<version>0.38.0-beta.2</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.3</lance-core.version>
|
||||
<lance-core.version>11.0.0-beta.15</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
@@ -1,7 +1,7 @@
|
||||
[package]
|
||||
name = "lancedb-nodejs"
|
||||
edition.workspace = true
|
||||
version = "0.37.1-beta.1"
|
||||
version = "0.38.0-beta.2"
|
||||
publish = false
|
||||
license.workspace = true
|
||||
description.workspace = true
|
||||
|
||||
@@ -4,7 +4,13 @@
|
||||
import { readdirSync } from "fs";
|
||||
import { Field, Float64, Schema } from "apache-arrow";
|
||||
import * as tmp from "tmp";
|
||||
import { Connection, Table, connect, connectNamespace } from "../lancedb";
|
||||
import {
|
||||
Connection,
|
||||
ListTablesResponse,
|
||||
Table,
|
||||
connect,
|
||||
connectNamespace,
|
||||
} from "../lancedb";
|
||||
import { LocalTable } from "../lancedb/table";
|
||||
|
||||
describe("when connecting", () => {
|
||||
@@ -89,6 +95,16 @@ 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);
|
||||
@@ -119,6 +135,56 @@ 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 }));
|
||||
|
||||
@@ -3340,3 +3340,120 @@ 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]);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -25,12 +25,14 @@ import type {
|
||||
JobDescription,
|
||||
JobInfo,
|
||||
ListNamespacesResponse,
|
||||
ListTablesResponse,
|
||||
} from "./native";
|
||||
export type {
|
||||
CreateNamespaceResponse,
|
||||
DescribeNamespaceResponse,
|
||||
DropNamespaceResponse,
|
||||
ListNamespacesResponse,
|
||||
ListTablesResponse,
|
||||
};
|
||||
import { sanitizeTable } from "./sanitize";
|
||||
import { LocalTable, Table } from "./table";
|
||||
@@ -128,6 +130,10 @@ export interface OpenTableOptions {
|
||||
indexCacheSize?: number;
|
||||
}
|
||||
|
||||
/**
|
||||
* @deprecated Use {@link ListTablesOptions} with {@link Connection.listTables}
|
||||
* instead.
|
||||
*/
|
||||
export interface TableNamesOptions {
|
||||
/**
|
||||
* If present, only return names that come lexicographically after the
|
||||
@@ -141,6 +147,23 @@ 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;
|
||||
@@ -225,6 +248,7 @@ 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[]>;
|
||||
/**
|
||||
@@ -235,12 +259,54 @@ 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
|
||||
@@ -327,6 +393,14 @@ 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).
|
||||
@@ -523,6 +597,29 @@ 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[],
|
||||
@@ -705,6 +802,10 @@ 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 ?? []);
|
||||
}
|
||||
|
||||
@@ -50,6 +50,7 @@ export {
|
||||
MergeResult,
|
||||
AddResult,
|
||||
AddColumnsResult,
|
||||
RefreshColumnResult,
|
||||
AlterColumnsResult,
|
||||
UpdateFieldMetadataResult,
|
||||
DeleteResult,
|
||||
@@ -74,11 +75,13 @@ export {
|
||||
Connection,
|
||||
CreateTableOptions,
|
||||
TableNamesOptions,
|
||||
ListTablesOptions,
|
||||
OpenTableOptions,
|
||||
ListNamespacesOptions,
|
||||
CreateNamespaceOptions,
|
||||
DropNamespaceOptions,
|
||||
ListNamespacesResponse,
|
||||
ListTablesResponse,
|
||||
CreateNamespaceResponse,
|
||||
DropNamespaceResponse,
|
||||
DescribeNamespaceResponse,
|
||||
@@ -146,6 +149,10 @@ export {
|
||||
FtsToken,
|
||||
TokenizeTableOptions,
|
||||
LsmWriteSpec,
|
||||
LsmStats,
|
||||
BucketStats,
|
||||
GenerationStats,
|
||||
MemtableStats,
|
||||
ColumnAlteration,
|
||||
FieldMetadataUpdate,
|
||||
} from "./table";
|
||||
|
||||
+160
-2
@@ -31,8 +31,10 @@ import {
|
||||
IndexConfig,
|
||||
IndexStatistics,
|
||||
Job,
|
||||
LsmStats,
|
||||
Branches as NativeBranches,
|
||||
OptimizeStats,
|
||||
RefreshColumnResult,
|
||||
TableStatistics,
|
||||
Tags,
|
||||
UpdateFieldMetadataResult,
|
||||
@@ -49,6 +51,12 @@ 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`
|
||||
@@ -525,18 +533,75 @@ 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,
|
||||
newColumnTransforms:
|
||||
| AddColumnsSql[]
|
||||
| Field
|
||||
| Field[]
|
||||
| Schema
|
||||
| { computed: AddColumnsSql[] },
|
||||
): 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
|
||||
@@ -648,6 +713,59 @@ 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>;
|
||||
@@ -1088,8 +1206,22 @@ export class LocalTable extends Table {
|
||||
// TODO: Support BatchUDF
|
||||
|
||||
async addColumns(
|
||||
newColumnTransforms: AddColumnsSql[] | Field | Field[] | Schema,
|
||||
newColumnTransforms:
|
||||
| AddColumnsSql[]
|
||||
| Field
|
||||
| Field[]
|
||||
| Schema
|
||||
| { computed: AddColumnsSql[] },
|
||||
): 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];
|
||||
@@ -1124,6 +1256,14 @@ 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> {
|
||||
@@ -1186,6 +1326,24 @@ 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,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-darwin-arm64",
|
||||
"version": "0.37.1-beta.1",
|
||||
"version": "0.38.0-beta.2",
|
||||
"os": ["darwin"],
|
||||
"cpu": ["arm64"],
|
||||
"main": "lancedb.darwin-arm64.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-linux-arm64-gnu",
|
||||
"version": "0.37.1-beta.1",
|
||||
"version": "0.38.0-beta.2",
|
||||
"os": ["linux"],
|
||||
"cpu": ["arm64"],
|
||||
"main": "lancedb.linux-arm64-gnu.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-linux-arm64-musl",
|
||||
"version": "0.37.1-beta.1",
|
||||
"version": "0.38.0-beta.2",
|
||||
"os": ["linux"],
|
||||
"cpu": ["arm64"],
|
||||
"main": "lancedb.linux-arm64-musl.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-linux-x64-gnu",
|
||||
"version": "0.37.1-beta.1",
|
||||
"version": "0.38.0-beta.2",
|
||||
"os": ["linux"],
|
||||
"cpu": ["x64"],
|
||||
"main": "lancedb.linux-x64-gnu.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-linux-x64-musl",
|
||||
"version": "0.37.1-beta.1",
|
||||
"version": "0.38.0-beta.2",
|
||||
"os": ["linux"],
|
||||
"cpu": ["x64"],
|
||||
"main": "lancedb.linux-x64-musl.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-win32-arm64-msvc",
|
||||
"version": "0.37.1-beta.1",
|
||||
"version": "0.38.0-beta.2",
|
||||
"os": [
|
||||
"win32"
|
||||
],
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-win32-x64-msvc",
|
||||
"version": "0.37.1-beta.1",
|
||||
"version": "0.38.0-beta.2",
|
||||
"os": ["win32"],
|
||||
"cpu": ["x64"],
|
||||
"main": "lancedb.win32-x64-msvc.node",
|
||||
|
||||
Generated
+2
-2
@@ -1,12 +1,12 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb",
|
||||
"version": "0.37.1-beta.1",
|
||||
"version": "0.38.0-beta.2",
|
||||
"lockfileVersion": 3,
|
||||
"requires": true,
|
||||
"packages": {
|
||||
"": {
|
||||
"name": "@lancedb/lancedb",
|
||||
"version": "0.37.1-beta.1",
|
||||
"version": "0.38.0-beta.2",
|
||||
"cpu": [
|
||||
"x64",
|
||||
"arm64"
|
||||
|
||||
+1
-1
@@ -11,7 +11,7 @@
|
||||
"ann"
|
||||
],
|
||||
"private": false,
|
||||
"version": "0.37.1-beta.1",
|
||||
"version": "0.38.0-beta.2",
|
||||
"main": "dist/index.js",
|
||||
"exports": {
|
||||
".": "./dist/index.js",
|
||||
|
||||
@@ -36,6 +36,12 @@ 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>>,
|
||||
@@ -189,6 +195,8 @@ 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>>,
|
||||
@@ -206,6 +214,29 @@ 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:
|
||||
@@ -334,6 +365,22 @@ 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();
|
||||
|
||||
+1
-11
@@ -42,19 +42,9 @@ 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<()> {
|
||||
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(),
|
||||
)),
|
||||
}
|
||||
self.inner.wait().await.default_error()
|
||||
}
|
||||
|
||||
/// Request cancellation. Cancelling a finished operation is a no-op.
|
||||
|
||||
@@ -347,6 +347,40 @@ 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,
|
||||
@@ -463,6 +497,34 @@ 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()?
|
||||
@@ -855,6 +917,129 @@ 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)]
|
||||
@@ -1196,6 +1381,21 @@ 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
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "lancedb-python"
|
||||
version = "0.37.1-beta.1"
|
||||
version = "0.38.0-beta.2"
|
||||
publish = false
|
||||
edition.workspace = true
|
||||
description = "Python bindings for LanceDB"
|
||||
|
||||
@@ -12,7 +12,7 @@ __version__ = importlib.metadata.version("lancedb")
|
||||
|
||||
from ._lancedb import connect as lancedb_connect
|
||||
from ._lancedb import FtsToken
|
||||
from ._lancedb import Function
|
||||
from ._lancedb import LsmWriteSpec
|
||||
from ._lancedb import tokenize as _tokenize
|
||||
from .common import URI, sanitize_uri
|
||||
from urllib.parse import urlparse
|
||||
@@ -24,7 +24,6 @@ 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,
|
||||
@@ -509,8 +508,6 @@ __all__ = [
|
||||
"FtsToken",
|
||||
"col",
|
||||
"Expr",
|
||||
"Function",
|
||||
"FunctionCapability",
|
||||
"func",
|
||||
"lit",
|
||||
"URI",
|
||||
@@ -522,9 +519,9 @@ __all__ = [
|
||||
"Job",
|
||||
"LanceDBConnection",
|
||||
"LanceNamespaceDBConnection",
|
||||
"LsmWriteSpec",
|
||||
"RemoteDBConnection",
|
||||
"Session",
|
||||
"Table",
|
||||
"udf",
|
||||
"__version__",
|
||||
]
|
||||
|
||||
@@ -1,108 +0,0 @@
|
||||
# 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)
|
||||
@@ -153,16 +153,6 @@ 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,
|
||||
@@ -208,6 +198,9 @@ 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: ...
|
||||
@@ -226,45 +219,11 @@ 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) -> Optional[Function]: ...
|
||||
async def wait(self) -> None: ...
|
||||
async def cancel(self) -> None: ...
|
||||
|
||||
class JobInfo:
|
||||
@@ -286,8 +245,6 @@ class JobFailureInfo:
|
||||
def message(self) -> Optional[str]: ...
|
||||
@property
|
||||
def retryable(self) -> Optional[bool]: ...
|
||||
@property
|
||||
def error_code(self) -> Optional[str]: ...
|
||||
|
||||
class JobDescription:
|
||||
@property
|
||||
@@ -302,8 +259,6 @@ 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: ...
|
||||
@@ -366,16 +321,6 @@ 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]): ...
|
||||
@@ -393,6 +338,11 @@ 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]]
|
||||
@@ -738,6 +688,10 @@ class LsmWriteSpec:
|
||||
class AddColumnsResult:
|
||||
version: int
|
||||
|
||||
class RefreshColumnResult:
|
||||
rows_filled: int
|
||||
version: int
|
||||
|
||||
class AlterColumnsResult:
|
||||
version: int
|
||||
|
||||
|
||||
@@ -1,538 +0,0 @@
|
||||
# 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,
|
||||
)
|
||||
+37
-126
@@ -63,12 +63,8 @@ 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
|
||||
@@ -528,6 +524,12 @@ 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,
|
||||
@@ -654,71 +656,6 @@ 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):
|
||||
"""
|
||||
@@ -1255,6 +1192,20 @@ 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:
|
||||
@@ -1336,34 +1287,6 @@ 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.
|
||||
@@ -2060,6 +1983,23 @@ 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.
|
||||
|
||||
@@ -2110,35 +2050,6 @@ 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.
|
||||
|
||||
|
||||
@@ -3,8 +3,6 @@
|
||||
|
||||
"""Custom exception handling"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class MissingValueError(ValueError):
|
||||
"""Exception raised when a required value is missing."""
|
||||
@@ -28,47 +26,12 @@ 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."""
|
||||
|
||||
``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
|
||||
pass
|
||||
|
||||
|
||||
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
|
||||
|
||||
@@ -10,7 +10,6 @@ from typing import Optional
|
||||
from lancedb.background_loop import LOOP
|
||||
|
||||
from . import _lancedb
|
||||
from ._lancedb import Function
|
||||
|
||||
|
||||
class AsyncJob:
|
||||
@@ -45,22 +44,18 @@ class AsyncJob:
|
||||
return "finished"
|
||||
return await self._inner.status()
|
||||
|
||||
async def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
|
||||
async def wait(self, timeout: Optional[timedelta] = None):
|
||||
"""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 None
|
||||
return
|
||||
if timeout is None:
|
||||
return await self._inner.wait()
|
||||
await self._inner.wait()
|
||||
else:
|
||||
return await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
|
||||
await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
|
||||
|
||||
async def cancel(self):
|
||||
"""Request cancellation. Cancelling a finished operation is a no-op."""
|
||||
@@ -93,19 +88,15 @@ class Job:
|
||||
return "finished"
|
||||
return LOOP.run(self._inner.status())
|
||||
|
||||
def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
|
||||
def wait(self, timeout: Optional[timedelta] = None):
|
||||
"""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 None
|
||||
return LOOP.run(self._inner.wait(timeout))
|
||||
return
|
||||
LOOP.run(self._inner.wait(timeout))
|
||||
|
||||
def cancel(self):
|
||||
"""Request cancellation. Cancelling a finished operation is a no-op."""
|
||||
|
||||
@@ -49,6 +49,7 @@ 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,
|
||||
@@ -624,6 +625,18 @@ 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,
|
||||
@@ -1134,6 +1147,14 @@ 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,
|
||||
|
||||
@@ -2235,6 +2235,7 @@ class LanceHybridQueryBuilder(LanceQueryBuilder):
|
||||
reranker=self._reranker,
|
||||
limit=self._limit,
|
||||
with_row_ids=True,
|
||||
offset=self._offset,
|
||||
)
|
||||
return self._finish_hybrid_results(results)
|
||||
|
||||
@@ -2256,6 +2257,7 @@ 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")
|
||||
@@ -2332,7 +2334,7 @@ class LanceHybridQueryBuilder(LanceQueryBuilder):
|
||||
score_i = results.column_names.index("_score")
|
||||
results = results.set_column(score_i, "_score", original_scores)
|
||||
|
||||
results = results.slice(length=limit)
|
||||
results = results.slice(offset=offset or 0, length=limit)
|
||||
|
||||
if not with_row_ids:
|
||||
results = results.drop(["_rowid"])
|
||||
@@ -2679,8 +2681,12 @@ class LanceHybridQueryBuilder(LanceQueryBuilder):
|
||||
|
||||
# Apply common configurations
|
||||
if self._limit:
|
||||
self._vector_query.limit(self._limit)
|
||||
self._fts_query.limit(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)
|
||||
if self._columns:
|
||||
self._vector_query.select(self._columns)
|
||||
self._fts_query.select(self._columns)
|
||||
|
||||
@@ -23,12 +23,10 @@ import pyarrow as pa
|
||||
|
||||
from ..common import DATA
|
||||
from ..db import DBConnection, LOOP
|
||||
from ..job import Job
|
||||
from ..job import AsyncJob, Job
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .._lancedb import Function
|
||||
from .._lancedb import Job as NativeJob
|
||||
from .._lancedb import JobDescription, JobInfo, _FunctionDefinition
|
||||
from .._lancedb import JobDescription, JobInfo
|
||||
from ..embeddings import EmbeddingFunctionConfig
|
||||
from lance_namespace import (
|
||||
LanceNamespace,
|
||||
@@ -665,6 +663,16 @@ 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,
|
||||
@@ -736,34 +744,6 @@ 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.
|
||||
|
||||
@@ -7,7 +7,6 @@ import logging
|
||||
from functools import cached_property
|
||||
import os
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
@@ -68,9 +67,6 @@ 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__(
|
||||
@@ -574,45 +570,6 @@ 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,
|
||||
@@ -1001,8 +958,19 @@ 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]) -> AddColumnsResult:
|
||||
return LOOP.run(self._table.add_columns(transforms))
|
||||
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 alter_columns(
|
||||
self, *alterations: Iterable[Dict[str, str]]
|
||||
@@ -1022,17 +990,39 @@ class RemoteTable(Table):
|
||||
return LOOP.run(self._table.set_unenforced_primary_key(columns))
|
||||
|
||||
def set_lsm_write_spec(self, spec: "LsmWriteSpec") -> None:
|
||||
"""Not supported on LanceDB Cloud."""
|
||||
"""Install an LsmWriteSpec."""
|
||||
return LOOP.run(self._table.set_lsm_write_spec(spec))
|
||||
|
||||
def unset_lsm_write_spec(self) -> None:
|
||||
"""Not supported on LanceDB Cloud."""
|
||||
"""Remove the LsmWriteSpec."""
|
||||
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())
|
||||
|
||||
+202
-127
@@ -176,6 +176,7 @@ if TYPE_CHECKING:
|
||||
CompactionStats,
|
||||
Tag,
|
||||
AddColumnsResult,
|
||||
RefreshColumnResult,
|
||||
AddResult,
|
||||
AlterColumnsResult,
|
||||
UpdateFieldMetadataResult,
|
||||
@@ -185,7 +186,6 @@ if TYPE_CHECKING:
|
||||
LsmWriteSpec,
|
||||
MergeResult,
|
||||
UpdateResult,
|
||||
_FunctionCall,
|
||||
)
|
||||
from .index import IndexConfig
|
||||
import pandas
|
||||
@@ -1008,43 +1008,6 @@ 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.
|
||||
@@ -1954,7 +1917,14 @@ class Table(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def add_columns(
|
||||
self, transforms: Dict[str, str] | pa.Field | List[pa.Field] | pa.Schema
|
||||
self,
|
||||
transforms: Dict[str, str]
|
||||
| pa.Field
|
||||
| List[pa.Field]
|
||||
| pa.Schema
|
||||
| None = None,
|
||||
*,
|
||||
computed: Dict[str, str] | None = None,
|
||||
):
|
||||
"""
|
||||
Add new columns with defined values.
|
||||
@@ -1968,11 +1938,95 @@ 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
|
||||
@@ -2887,43 +2941,6 @@ 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,
|
||||
@@ -4014,9 +4031,28 @@ 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
|
||||
self,
|
||||
transforms: Dict[str, str]
|
||||
| pa.field
|
||||
| List[pa.field]
|
||||
| pa.Schema
|
||||
| None = None,
|
||||
*,
|
||||
computed: Dict[str, str] | None = None,
|
||||
) -> AddColumnsResult:
|
||||
return LOOP.run(self._table.add_columns(transforms))
|
||||
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)))
|
||||
|
||||
def alter_columns(
|
||||
self, *alterations: Iterable[Dict[str, str]]
|
||||
@@ -4765,7 +4801,7 @@ class AsyncTable:
|
||||
|
||||
Examples
|
||||
--------
|
||||
>>> from lancedb._lancedb import LsmWriteSpec
|
||||
>>> from lancedb import LsmWriteSpec
|
||||
>>> # table.set_unenforced_primary_key("id")
|
||||
>>> # table.set_lsm_write_spec(LsmWriteSpec.bucket("id", 16))
|
||||
"""
|
||||
@@ -4810,7 +4846,7 @@ class AsyncTable:
|
||||
``asyncio.wait_for`` for a wall-clock bound; abandoning it partway
|
||||
costs nothing.
|
||||
"""
|
||||
return await self._inner.checkpoint_lsm()
|
||||
await self._inner.checkpoint_lsm()
|
||||
|
||||
async def flush_lsm(self) -> None:
|
||||
"""Seal every bucket's active memtable into L0.
|
||||
@@ -4819,7 +4855,7 @@ class AsyncTable:
|
||||
`compact_lsm`. On a node that has not claimed this table, this claims
|
||||
it and replays its WAL log first.
|
||||
"""
|
||||
return await self._inner.flush_lsm()
|
||||
await self._inner.flush_lsm()
|
||||
|
||||
async def compact_lsm(self) -> None:
|
||||
"""Trigger a background L0 to base compaction pass per bucket.
|
||||
@@ -4828,7 +4864,7 @@ class AsyncTable:
|
||||
``get_lsm_stats`` for progress, or use ``checkpoint_lsm`` to loop
|
||||
until the current L0 has reached base.
|
||||
"""
|
||||
return await self._inner.compact_lsm()
|
||||
await self._inner.compact_lsm()
|
||||
|
||||
async def get_lsm_stats(
|
||||
self, *, include_generation_rows: bool = False
|
||||
@@ -5141,50 +5177,6 @@ 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.
|
||||
@@ -5975,7 +5967,14 @@ 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
|
||||
self,
|
||||
transforms: dict[str, str]
|
||||
| pa.field
|
||||
| List[pa.field]
|
||||
| pa.Schema
|
||||
| None = None,
|
||||
*,
|
||||
computed: dict[str, str] | None = None,
|
||||
) -> AddColumnsResult:
|
||||
"""
|
||||
Add new columns with defined values.
|
||||
@@ -5988,6 +5987,22 @@ 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
|
||||
-------
|
||||
@@ -6001,11 +6016,71 @@ 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:
|
||||
|
||||
@@ -755,8 +755,7 @@ def test_delete_table(tmp_db: lancedb.DBConnection):
|
||||
assert tmp_db.table_names() == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_table_async(tmp_db: lancedb.DBConnection):
|
||||
def test_drop_table_async(tmp_db: lancedb.DBConnection):
|
||||
data = pd.DataFrame(
|
||||
{
|
||||
"vector": [[3.1, 4.1], [5.9, 26.5]],
|
||||
@@ -772,7 +771,10 @@ async def test_delete_table_async(tmp_db: lancedb.DBConnection):
|
||||
|
||||
assert tmp_db.table_names() == ["test"]
|
||||
|
||||
tmp_db.drop_table("test")
|
||||
job = tmp_db.drop_table_async("test")
|
||||
assert job.id is None
|
||||
assert job.status() == "finished"
|
||||
job.wait()
|
||||
assert tmp_db.table_names() == []
|
||||
|
||||
tmp_db.create_table("test", data=data)
|
||||
@@ -781,6 +783,17 @@ async def test_delete_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(
|
||||
{
|
||||
|
||||
@@ -632,3 +632,101 @@ 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
|
||||
|
||||
@@ -1,372 +0,0 @@
|
||||
# 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)"
|
||||
@@ -1,181 +0,0 @@
|
||||
# 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=(",", ":")))
|
||||
@@ -1,595 +0,0 @@
|
||||
# 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,
|
||||
},
|
||||
)
|
||||
@@ -1,268 +0,0 @@
|
||||
# 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
|
||||
@@ -1,634 +0,0 @@
|
||||
# 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)
|
||||
@@ -1,398 +0,0 @@
|
||||
# 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
|
||||
@@ -1,719 +0,0 @@
|
||||
# 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
|
||||
@@ -1,579 +0,0 @@
|
||||
# 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
|
||||
@@ -1,728 +0,0 @@
|
||||
# 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
@@ -1,899 +0,0 @@
|
||||
# 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
@@ -1,672 +0,0 @@
|
||||
# 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)
|
||||
@@ -1,291 +0,0 @@
|
||||
# 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)
|
||||
@@ -1,490 +0,0 @@
|
||||
# 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
|
||||
@@ -1,506 +0,0 @@
|
||||
# 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)
|
||||
@@ -1,486 +0,0 @@
|
||||
# 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)
|
||||
@@ -203,6 +203,31 @@ 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.
|
||||
|
||||
@@ -1133,6 +1133,131 @@ 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):
|
||||
@@ -2306,228 +2431,3 @@ 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
|
||||
|
||||
@@ -3854,3 +3854,65 @@ 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]
|
||||
|
||||
+30
-124
@@ -23,7 +23,6 @@ use lancedb::{
|
||||
connection::NamespaceClientPushdownOperation,
|
||||
database::namespace::LanceNamespaceDatabase,
|
||||
database::{CreateTableMode, Database, ReadConsistency},
|
||||
function::{FunctionId, RegisterFunctionJobSpec},
|
||||
};
|
||||
use pyo3::{
|
||||
Bound, FromPyObject, Py, PyAny, PyRef, PyResult, Python,
|
||||
@@ -122,6 +121,8 @@ 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>>,
|
||||
@@ -347,6 +348,23 @@ 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>,
|
||||
@@ -506,14 +524,17 @@ impl Connection {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
let py = self_.py();
|
||||
future_into_py(py, async move {
|
||||
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()?;
|
||||
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()?;
|
||||
Python::attach(|py| -> PyResult<Py<PyDict>> {
|
||||
let dict = PyDict::new(py);
|
||||
dict.set_item("tables", response.tables)?;
|
||||
@@ -590,121 +611,6 @@ 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, ¤t)
|
||||
.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]
|
||||
|
||||
+2
-15
@@ -102,14 +102,11 @@ impl<T> PythonErrorExt<T> for std::result::Result<T, LanceError> {
|
||||
err.setattr(intern!(py, "__cause__"), cause_err)?;
|
||||
Err(PyErr::from_value(err))
|
||||
}),
|
||||
LanceError::JobFailed { failure, .. } => Python::attach(|py| {
|
||||
LanceError::JobFailed { .. } => Python::attach(|py| {
|
||||
let cls = py
|
||||
.import(intern!(py, "lancedb.exceptions"))?
|
||||
.getattr(intern!(py, "JobFailedError"))?;
|
||||
// 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))?))
|
||||
Err(PyErr::from_value(cls.call1((err.to_string(),))?))
|
||||
}),
|
||||
LanceError::JobCancelled { .. } => Python::attach(|py| {
|
||||
let cls = py
|
||||
@@ -117,16 +114,6 @@ 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(),
|
||||
},
|
||||
}
|
||||
|
||||
+21
-29
@@ -10,7 +10,7 @@
|
||||
use std::ops::{Add, Div, Mul, Not, Sub};
|
||||
|
||||
use arrow::{datatypes::DataType, pyarrow::PyArrowType};
|
||||
use datafusion_common::{Column, ScalarValue};
|
||||
use datafusion_common::ScalarValue;
|
||||
use lancedb::expr::{
|
||||
DfExpr, col as ldb_col, contains, expr_cast, is_in, lit as df_lit, lower, upper,
|
||||
};
|
||||
@@ -27,33 +27,6 @@ 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 ──────────────────────────────────────────────────────────
|
||||
@@ -218,8 +191,27 @@ 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 = dt.call_method0("timestamp")?.extract()?;
|
||||
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 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
+4
-31
@@ -3,7 +3,6 @@
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::function::Function;
|
||||
use crate::runtime::future_into_py;
|
||||
use pyo3::{Bound, PyAny, PyRef, PyResult, pyclass, pymethods};
|
||||
|
||||
@@ -22,23 +21,6 @@ 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]
|
||||
@@ -57,8 +39,8 @@ impl Job {
|
||||
pub fn wait(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let result = inner.wait().await.infer_error()?;
|
||||
Ok(project_wait_result(result))
|
||||
inner.wait().await.infer_error()?;
|
||||
Ok(())
|
||||
})
|
||||
}
|
||||
|
||||
@@ -111,16 +93,14 @@ 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={:?}, error_code={:?})",
|
||||
self.phase, self.message, self.retryable, self.error_code
|
||||
"JobFailureInfo(phase={:?}, message={:?}, retryable={:?})",
|
||||
self.phase, self.message, self.retryable
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -135,7 +115,6 @@ pub struct JobDescription {
|
||||
creation_ms: i64,
|
||||
spec_json: Option<String>,
|
||||
failure: Option<JobFailureInfo>,
|
||||
result: Option<Function>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
@@ -160,13 +139,7 @@ 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),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+3
-9
@@ -16,14 +16,14 @@ use query::{FTSQuery, HybridQuery, Query, VectorQuery};
|
||||
use session::Session;
|
||||
use table::{
|
||||
AddColumnsResult, AddResult, AlterColumnsResult, DeleteResult, DropColumnsResult, FtsToken,
|
||||
LsmWriteSpec, MergeResult, PyBlobFile, Table, UpdateFieldMetadataResult, UpdateResult,
|
||||
LsmWriteSpec, MergeResult, PyBlobFile, RefreshColumnResult, 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,9 +46,6 @@ 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>()?;
|
||||
@@ -61,6 +58,7 @@ 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>()?;
|
||||
@@ -92,10 +90,6 @@ 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(())
|
||||
}
|
||||
|
||||
+61
-141
@@ -26,7 +26,7 @@ use lancedb::table::{
|
||||
use lancedb::tokenize as lancedb_tokenize;
|
||||
use pyo3::{
|
||||
Bound, FromPyObject, Py, PyAny, PyRef, PyResult, Python,
|
||||
exceptions::{PyNotImplementedError, PyRuntimeError, PyValueError},
|
||||
exceptions::{PyRuntimeError, PyValueError},
|
||||
pyclass, pyfunction, pymethods,
|
||||
types::{IntoPyDict, PyAnyMethods, PyBytes, PyDict, PyDictMethods, PyList, PyListMethods},
|
||||
};
|
||||
@@ -415,6 +415,32 @@ 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 {
|
||||
@@ -930,146 +956,6 @@ 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 {
|
||||
@@ -1650,6 +1536,40 @@ 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>,
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "lancedb"
|
||||
version = "0.37.1-beta.1"
|
||||
version = "0.38.0-beta.2"
|
||||
edition.workspace = true
|
||||
description = "LanceDB: A serverless, low-latency vector database for AI applications"
|
||||
license.workspace = true
|
||||
@@ -12,7 +12,6 @@ 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 }
|
||||
@@ -189,6 +188,9 @@ required-features = ["bedrock"]
|
||||
[[example]]
|
||||
name = "bench_streaming_dataloader"
|
||||
|
||||
[[example]]
|
||||
name = "bench_open_missing_table"
|
||||
|
||||
[[example]]
|
||||
name = "simple"
|
||||
|
||||
|
||||
@@ -0,0 +1,150 @@
|
||||
// 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(())
|
||||
}
|
||||
@@ -27,7 +27,7 @@ async fn main() -> Result<()> {
|
||||
// --8<-- [end:connect]
|
||||
|
||||
// --8<-- [start:list_names]
|
||||
println!("{:?}", db.table_names().execute().await?);
|
||||
println!("{:?}", db.list_tables().execute().await?.tables);
|
||||
// --8<-- [end:list_names]
|
||||
let tbl = create_table(&db).await?;
|
||||
create_index(&tbl).await?;
|
||||
|
||||
@@ -333,13 +333,11 @@ pub(crate) fn ensure_blob_storage_version(schema: &Schema, params: &mut WritePar
|
||||
.data_storage_version
|
||||
.unwrap_or(LanceFileVersion::Stable)
|
||||
.resolve();
|
||||
// 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 => {}
|
||||
if matches!(
|
||||
resolved,
|
||||
ConcreteFileVersion::V1 | ConcreteFileVersion::V2_0 | ConcreteFileVersion::V2_1
|
||||
) {
|
||||
params.data_storage_version = Some(LanceFileVersion::V2_2);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -504,7 +502,7 @@ mod tests {
|
||||
ensure_blob_storage_version(&blob_schema(), &mut params);
|
||||
assert_eq!(
|
||||
params.data_storage_version.unwrap().resolve(),
|
||||
LanceFileVersion::V2_2.resolve()
|
||||
ConcreteFileVersion::V2_2
|
||||
);
|
||||
}
|
||||
|
||||
@@ -517,7 +515,7 @@ mod tests {
|
||||
ensure_blob_storage_version(&blob_schema(), &mut params);
|
||||
assert_eq!(
|
||||
params.data_storage_version.unwrap().resolve(),
|
||||
LanceFileVersion::V2_2.resolve()
|
||||
ConcreteFileVersion::V2_2
|
||||
);
|
||||
}
|
||||
|
||||
@@ -528,10 +526,7 @@ mod tests {
|
||||
..Default::default()
|
||||
};
|
||||
ensure_blob_storage_version(&blob_schema(), &mut params);
|
||||
assert_eq!(
|
||||
params.data_storage_version.unwrap().resolve(),
|
||||
LanceFileVersion::V2_3.resolve()
|
||||
);
|
||||
assert_eq!(params.data_storage_version.unwrap(), LanceFileVersion::V2_3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
+176
-91
@@ -28,7 +28,6 @@ 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,
|
||||
@@ -74,11 +73,13 @@ 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 {
|
||||
@@ -116,6 +117,57 @@ 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>,
|
||||
@@ -410,7 +462,14 @@ impl Connection {
|
||||
///
|
||||
/// The names will be returned in lexicographical order (ascending)
|
||||
///
|
||||
/// The parameters `page_token` and `limit` can be used to paginate the results
|
||||
/// 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)]
|
||||
pub fn table_names(&self) -> TableNamesBuilder {
|
||||
TableNamesBuilder::new(self.internal.clone())
|
||||
}
|
||||
@@ -457,10 +516,9 @@ impl Connection {
|
||||
///
|
||||
/// # Returns
|
||||
/// Created [`TableRef`], or [`Error::TableNotFound`] if the table does not exist.
|
||||
/// 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.
|
||||
/// 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.
|
||||
pub fn open_table(&self, name: impl Into<String>) -> OpenTableBuilder {
|
||||
OpenTableBuilder::new(
|
||||
self.internal.clone(),
|
||||
@@ -551,88 +609,6 @@ 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
|
||||
@@ -644,6 +620,21 @@ 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
|
||||
@@ -716,9 +707,32 @@ impl Connection {
|
||||
self.internal.namespace_client_config().await
|
||||
}
|
||||
|
||||
/// List tables with pagination support
|
||||
pub async fn list_tables(&self, request: ListTablesRequest) -> Result<ListTablesResponse> {
|
||||
self.internal.list_tables(request).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())
|
||||
}
|
||||
|
||||
/// Get the in-memory embedding registry.
|
||||
@@ -1374,6 +1388,8 @@ 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};
|
||||
@@ -1716,6 +1732,75 @@ 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();
|
||||
|
||||
@@ -439,7 +439,7 @@ mod tests {
|
||||
.unwrap()
|
||||
.data_storage_format
|
||||
.lance_file_format();
|
||||
// Compare concrete stored format to the resolved requested alias.
|
||||
// Compare resolved versions since Stable/Next are aliases that resolve at storage time
|
||||
assert_eq!(storage_format, data_storage_version.resolve());
|
||||
}
|
||||
|
||||
|
||||
@@ -30,16 +30,12 @@ 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>);
|
||||
}
|
||||
@@ -234,12 +230,6 @@ 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>,
|
||||
@@ -321,64 +311,6 @@ 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
|
||||
@@ -391,6 +323,18 @@ 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;
|
||||
|
||||
@@ -1,788 +0,0 @@
|
||||
// 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
Reference in New Issue
Block a user