mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-02 11:38:49 +00:00
Compare commits
60 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 3a3ddfda01 | |||
| c4371eb500 | |||
| 7a08580400 | |||
| 2fccab172f | |||
| e478b80985 | |||
| 72767b17fa | |||
| d7d25cd5ef | |||
| 4843445a7e | |||
| 713510375b | |||
| a9617bf830 | |||
| 208787ae7b | |||
| 9b46e7a448 | |||
| 7a46da2e67 | |||
| 705f7e7760 | |||
| d38a566282 | |||
| 5a27c71ab8 | |||
| 76f6487d92 | |||
| 6f6c3c33e0 | |||
| 1aa3665d67 | |||
| 0194f2317a | |||
| c2a647189d | |||
| df67ee4028 | |||
| d9e41228c8 | |||
| 68597070d2 | |||
| c825737780 | |||
| 0a43795996 | |||
| 0fa2fa05ad | |||
| 93ba442ac2 | |||
| 7a94ab7d6c | |||
| 6ed1a25439 | |||
| ca1d04db25 | |||
| efe3300404 | |||
| ecf87f6371 | |||
| 47213e31f8 | |||
| f65bf89c98 | |||
| d902144605 | |||
| a49dc5c71d | |||
| 98fed41efa | |||
| 1524ee0669 | |||
| 29be3e5509 | |||
| 8cedd50495 | |||
| b71ada0fae | |||
| 206efd98ff | |||
| 65c0968c0f | |||
| 2b10f2a7ce | |||
| f8bb90405f | |||
| 76aac96749 | |||
| 0093bc8179 | |||
| ac35a687f1 | |||
| 203f6536a6 | |||
| 9d3d0d0640 | |||
| a9ed8dba27 | |||
| 04acf1d3b5 | |||
| 3746118374 | |||
| d0b5cbe510 | |||
| 7b195adc3a | |||
| 818d6d1f59 | |||
| 9d589bea44 | |||
| 1798ece362 | |||
| 82b82711ba |
+1
-1
@@ -1,5 +1,5 @@
|
||||
[tool.bumpversion]
|
||||
current_version = "0.38.0-beta.0"
|
||||
current_version = "0.37.1-beta.1"
|
||||
parse = """(?x)
|
||||
(?P<major>0|[1-9]\\d*)\\.
|
||||
(?P<minor>0|[1-9]\\d*)\\.
|
||||
|
||||
@@ -4,14 +4,14 @@ on:
|
||||
workflow_call:
|
||||
inputs:
|
||||
tag:
|
||||
description: "Tag name from Lance (e.g. `v7.2.0-beta.1`). If omitted, the newest release is resolved automatically — stable releases are preferred over pre-releases — and the run is skipped if it is not newer than the version currently pinned in Cargo.toml."
|
||||
description: "Tag name from Lance. If omitted, the skill will use the latest Lance release that needs an update."
|
||||
required: false
|
||||
default: ""
|
||||
type: string
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
tag:
|
||||
description: "Tag name from Lance (e.g. `v7.2.0-beta.1`). Leave empty to resolve the newest release automatically — stable releases are preferred over pre-releases — and skip the run if it is not newer than the version currently pinned in Cargo.toml."
|
||||
description: "Tag name from Lance. Leave empty to use the latest Lance release that needs an update."
|
||||
required: false
|
||||
default: ""
|
||||
type: string
|
||||
|
||||
@@ -69,16 +69,6 @@ jobs:
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: "3.10"
|
||||
- name: Add swap for Arm fat LTO
|
||||
if: matrix.config.platform == 'aarch64'
|
||||
shell: bash
|
||||
run: |
|
||||
swap_file="$RUNNER_TEMP/lancedb-swap"
|
||||
sudo fallocate --length 16G "$swap_file"
|
||||
sudo chmod 600 "$swap_file"
|
||||
sudo mkswap "$swap_file"
|
||||
sudo swapon "$swap_file"
|
||||
free -h
|
||||
- uses: ./.github/workflows/build_linux_wheel
|
||||
with:
|
||||
python-minor-version: 10
|
||||
|
||||
Generated
+62
-48
@@ -3455,8 +3455,8 @@ checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c"
|
||||
|
||||
[[package]]
|
||||
name = "fsst"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"rand 0.9.5",
|
||||
@@ -4815,8 +4815,8 @@ checksum = "e037a2e1d8d5fdbd49b16a4ea09d5d6401c1f29eca5ff29d03d3824dba16256a"
|
||||
|
||||
[[package]]
|
||||
name = "lance"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arc-swap",
|
||||
"arrow",
|
||||
@@ -4832,6 +4832,7 @@ dependencies = [
|
||||
"async-recursion",
|
||||
"async-trait",
|
||||
"async_cell",
|
||||
"aws-credential-types",
|
||||
"aws-sdk-dynamodb",
|
||||
"byteorder",
|
||||
"bytes",
|
||||
@@ -4847,6 +4848,7 @@ dependencies = [
|
||||
"either",
|
||||
"fst",
|
||||
"futures",
|
||||
"half",
|
||||
"humantime",
|
||||
"itertools 0.14.0",
|
||||
"lance-arrow",
|
||||
@@ -4888,8 +4890,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-arrow"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -4911,7 +4913,7 @@ dependencies = [
|
||||
[[package]]
|
||||
name = "lance-arrow-scalar"
|
||||
version = "58.0.0"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -4925,7 +4927,7 @@ dependencies = [
|
||||
[[package]]
|
||||
name = "lance-arrow-stats"
|
||||
version = "58.0.0"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-schema",
|
||||
@@ -4934,8 +4936,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-bitpacking"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrayref",
|
||||
"crunchy",
|
||||
@@ -4945,8 +4947,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-core"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -4954,10 +4956,12 @@ dependencies = [
|
||||
"arrow-schema",
|
||||
"async-trait",
|
||||
"blake3",
|
||||
"byteorder",
|
||||
"bytes",
|
||||
"datafusion-common",
|
||||
"datafusion-sql",
|
||||
"futures",
|
||||
"itertools 0.14.0",
|
||||
"lance-arrow",
|
||||
"lance-derive",
|
||||
"libc",
|
||||
@@ -4975,6 +4979,7 @@ dependencies = [
|
||||
"snafu 0.9.0",
|
||||
"tempfile",
|
||||
"tokio",
|
||||
"tokio-stream",
|
||||
"tokio-util",
|
||||
"tracing",
|
||||
"twox-hash",
|
||||
@@ -4983,8 +4988,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-datafusion"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-array",
|
||||
@@ -5003,6 +5008,7 @@ dependencies = [
|
||||
"jsonb",
|
||||
"lance-arrow",
|
||||
"lance-core",
|
||||
"lance-datagen",
|
||||
"log",
|
||||
"pin-project",
|
||||
"prost",
|
||||
@@ -5013,8 +5019,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-datagen"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-array",
|
||||
@@ -5031,8 +5037,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-derive"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
@@ -5041,8 +5047,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-encoding"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-arith",
|
||||
"arrow-array",
|
||||
@@ -5067,6 +5073,7 @@ dependencies = [
|
||||
"num-traits",
|
||||
"prost",
|
||||
"prost-build",
|
||||
"rand 0.9.5",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"xxhash-rust",
|
||||
@@ -5075,8 +5082,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-file"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-arith",
|
||||
"arrow-array",
|
||||
@@ -5107,8 +5114,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-index"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arc-swap",
|
||||
"arrow",
|
||||
@@ -5123,6 +5130,7 @@ dependencies = [
|
||||
"async-trait",
|
||||
"bitvec",
|
||||
"bytes",
|
||||
"chrono",
|
||||
"crossbeam-queue",
|
||||
"datafusion",
|
||||
"datafusion-common",
|
||||
@@ -5140,6 +5148,7 @@ dependencies = [
|
||||
"lance-bitpacking",
|
||||
"lance-core",
|
||||
"lance-datafusion",
|
||||
"lance-datagen",
|
||||
"lance-encoding",
|
||||
"lance-file",
|
||||
"lance-index-core",
|
||||
@@ -5168,12 +5177,13 @@ dependencies = [
|
||||
"tempfile",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "lance-index-core"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-schema",
|
||||
@@ -5195,8 +5205,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-io"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-array",
|
||||
@@ -5210,6 +5220,7 @@ dependencies = [
|
||||
"futures",
|
||||
"http 1.5.0",
|
||||
"io-uring",
|
||||
"lance-arrow",
|
||||
"lance-core",
|
||||
"lance-namespace",
|
||||
"log",
|
||||
@@ -5227,28 +5238,29 @@ dependencies = [
|
||||
"tokio",
|
||||
"tracing",
|
||||
"url",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "lance-linalg"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
"arrow-schema",
|
||||
"cc",
|
||||
"half",
|
||||
"lance-arrow",
|
||||
"lance-core",
|
||||
"num-traits",
|
||||
"rand 0.9.5",
|
||||
"rayon",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "lance-namespace"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"async-trait",
|
||||
@@ -5260,8 +5272,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-namespace-impls"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-ipc",
|
||||
@@ -5300,9 +5312,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-namespace-reqwest-client"
|
||||
version = "0.11.0"
|
||||
version = "0.8.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0a030196da1c994b63a96a4f0bf5b0cfa459fe6dadc9e962320246ca328da22a"
|
||||
checksum = "ba3f0a235e3ed5f8805205649ccc7d7d0f3df23ce1294242c9265ad488d7f19d"
|
||||
dependencies = [
|
||||
"reqwest 0.12.28",
|
||||
"serde",
|
||||
@@ -5314,13 +5326,14 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-select"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
"arrow-schema",
|
||||
"byteorder",
|
||||
"bytes",
|
||||
"itertools 0.14.0",
|
||||
"lance-core",
|
||||
"roaring",
|
||||
@@ -5329,8 +5342,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-table"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-array",
|
||||
@@ -5370,8 +5383,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-testing"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-schema",
|
||||
@@ -5384,8 +5397,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-tokenizer"
|
||||
version = "11.0.0-beta.13"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.13#ee41152ceb9a78e5df4d2456fdbdb98542eb2059"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"frostem",
|
||||
"icu_segmenter",
|
||||
@@ -5398,7 +5411,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lancedb"
|
||||
version = "0.38.0-beta.0"
|
||||
version = "0.37.1-beta.1"
|
||||
dependencies = [
|
||||
"ahash",
|
||||
"anyhow",
|
||||
@@ -5419,6 +5432,7 @@ dependencies = [
|
||||
"aws-sdk-kms",
|
||||
"aws-sdk-s3",
|
||||
"aws-smithy-runtime",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"candle-core",
|
||||
"candle-nn",
|
||||
@@ -5486,7 +5500,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lancedb-nodejs"
|
||||
version = "0.38.0-beta.0"
|
||||
version = "0.37.1-beta.1"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -5511,7 +5525,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lancedb-python"
|
||||
version = "0.38.0-beta.0"
|
||||
version = "0.37.1-beta.1"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"async-trait",
|
||||
|
||||
+14
-14
@@ -13,20 +13,20 @@ categories = ["database-implementations"]
|
||||
rust-version = "1.91.0"
|
||||
|
||||
[workspace.dependencies]
|
||||
lance = { "version" = "=11.0.0-beta.13", default-features = false, "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-core = { "version" = "=11.0.0-beta.13", "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datagen = { "version" = "=11.0.0-beta.13", "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-file = { "version" = "=11.0.0-beta.13", "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-io = { "version" = "=11.0.0-beta.13", default-features = false, "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-index = { "version" = "=11.0.0-beta.13", "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-linalg = { "version" = "=11.0.0-beta.13", "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace = { "version" = "=11.0.0-beta.13", "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace-impls = { "version" = "=11.0.0-beta.13", default-features = false, "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-table = { "version" = "=11.0.0-beta.13", "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-testing = { "version" = "=11.0.0-beta.13", "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datafusion = { "version" = "=11.0.0-beta.13", "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-encoding = { "version" = "=11.0.0-beta.13", "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-arrow = { "version" = "=11.0.0-beta.13", "tag" = "v11.0.0-beta.13", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance = { version = "=11.0.0-beta.4", default-features = false, rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-core = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-datagen = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-file = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-io = { version = "=11.0.0-beta.4", default-features = false, rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-index = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-linalg = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace-impls = { version = "=11.0.0-beta.4", default-features = false, rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-table = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-testing = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-datafusion = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-encoding = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-arrow = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
ahash = "0.8"
|
||||
# Note that this one does not include pyarrow
|
||||
arrow = { version = "58.0.0", optional = false }
|
||||
|
||||
@@ -101,13 +101,6 @@ ignore = [
|
||||
# https://rustsec.org/advisories/RUSTSEC-2026-0195
|
||||
{ id = "RUSTSEC-2026-0194", reason = "transitive via inferno/lance/opendal; XML from trusted cloud endpoints, not attacker-controlled" },
|
||||
{ id = "RUSTSEC-2026-0195", reason = "transitive via inferno/lance/opendal; XML from trusted cloud endpoints, not attacker-controlled" },
|
||||
# smartstring: unmaintained — the repository was archived by its author on
|
||||
# 2026-05-03. Not a vulnerability. Reached only transitively through polars
|
||||
# (polars-core/-io/-ops/-time/-utils); nothing in LanceDB depends on it directly.
|
||||
# The advisory states no safe upgrade is available: upstream recommends
|
||||
# compact_str/smol_str, so clearing this requires polars to migrate.
|
||||
# https://rustsec.org/advisories/RUSTSEC-2026-0249
|
||||
{ id = "RUSTSEC-2026-0249", reason = "smartstring unmaintained via polars; no fixed upstream release" },
|
||||
]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -14,7 +14,7 @@ Add the following dependency to your `pom.xml`:
|
||||
<dependency>
|
||||
<groupId>com.lancedb</groupId>
|
||||
<artifactId>lancedb-core</artifactId>
|
||||
<version>0.38.0-beta.0</version>
|
||||
<version>0.37.1-beta.1</version>
|
||||
</dependency>
|
||||
```
|
||||
|
||||
|
||||
@@ -386,29 +386,6 @@ Drop an existing table.
|
||||
|
||||
***
|
||||
|
||||
### dropTableAsync()
|
||||
|
||||
```ts
|
||||
abstract dropTableAsync(name, namespacePath?): Promise<Job>
|
||||
```
|
||||
|
||||
Start dropping a table and return its cleanup job.
|
||||
|
||||
The table may become unavailable before its data files are removed. Wait
|
||||
on the returned job to know when cleanup has finished.
|
||||
|
||||
#### Parameters
|
||||
|
||||
* **name**: `string`
|
||||
|
||||
* **namespacePath?**: `string`[]
|
||||
|
||||
#### Returns
|
||||
|
||||
`Promise`<[`Job`](Job.md)>
|
||||
|
||||
***
|
||||
|
||||
### getJob()
|
||||
|
||||
```ts
|
||||
|
||||
@@ -69,34 +69,14 @@ abstract addColumns(newColumnTransforms): Promise<AddColumnsResult>
|
||||
|
||||
Add new columns with defined values.
|
||||
|
||||
The `{ computed }` form stores the expression rather than evaluating it
|
||||
now: the column is committed with no values, and rows get them from
|
||||
[Table#refreshColumn](Table.md#refreshcolumn). Declaring one therefore costs the same on a
|
||||
large table as on an empty one.
|
||||
|
||||
A refresh does not revisit rows it has already filled, so mutating an
|
||||
input leaves the value computed at fill time; recomputing means dropping
|
||||
the column and declaring it again. While a declaration reads a column,
|
||||
that column cannot be renamed, retyped or dropped.
|
||||
|
||||
On LanceDB Cloud and Enterprise the expression is planned by the
|
||||
server, and the refresh runs as a server job -- see
|
||||
[Table#refreshColumnAsync](Table.md#refreshcolumnasync).
|
||||
|
||||
#### Parameters
|
||||
|
||||
* **newColumnTransforms**:
|
||||
\| `Field`<`any`>
|
||||
\| `Field`<`any`>[]
|
||||
\| `Schema`<`any`>
|
||||
\| [`AddColumnsSql`](../interfaces/AddColumnsSql.md)[]
|
||||
\| `object`
|
||||
* **newColumnTransforms**: `Field`<`any`> \| `Field`<`any`>[] \| `Schema`<`any`> \| [`AddColumnsSql`](../interfaces/AddColumnsSql.md)[]
|
||||
Either:
|
||||
- An array of objects with column names and SQL expressions to calculate values
|
||||
- A single Arrow Field defining one column with its data type (column will be initialized with null values)
|
||||
- An array of Arrow Fields defining columns with their data types (columns will be initialized with null values)
|
||||
- An Arrow Schema defining columns with their data types (columns will be initialized with null values)
|
||||
- `{ computed }`, declaring columns defined by a SQL expression whose type and inputs are derived from it
|
||||
|
||||
#### Returns
|
||||
|
||||
@@ -105,13 +85,6 @@ server, and the refresh runs as a server job -- see
|
||||
A promise that resolves to an object
|
||||
containing the new version number of the table after adding the columns.
|
||||
|
||||
#### Example
|
||||
|
||||
```ts
|
||||
await table.addColumns({ computed: [{ name: "doubled", valueSql: "x * 2" }] });
|
||||
const { rowsFilled } = await table.refreshColumn("doubled");
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### alterColumns()
|
||||
@@ -745,67 +718,6 @@ for await (const batch of table.query()) {
|
||||
|
||||
***
|
||||
|
||||
### refreshColumn()
|
||||
|
||||
```ts
|
||||
abstract refreshColumn(column): Promise<RefreshColumnResult>
|
||||
```
|
||||
|
||||
Fill the rows of a computed column that hold no value yet.
|
||||
|
||||
Rows appended since the last refresh are filled by the next one; rows
|
||||
already filled are left as they are, so the call is idempotent and does
|
||||
not observe a mutated input. Local tables only: a remote refresh runs
|
||||
as a server job, through [Table#refreshColumnAsync](Table.md#refreshcolumnasync).
|
||||
|
||||
#### Parameters
|
||||
|
||||
* **column**: `string`
|
||||
The name of the computed column to fill.
|
||||
|
||||
#### Returns
|
||||
|
||||
`Promise`<[`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
|
||||
|
||||
@@ -105,7 +105,6 @@
|
||||
- [OptimizeOptions](interfaces/OptimizeOptions.md)
|
||||
- [OptimizeStats](interfaces/OptimizeStats.md)
|
||||
- [QueryExecutionOptions](interfaces/QueryExecutionOptions.md)
|
||||
- [RefreshColumnResult](interfaces/RefreshColumnResult.md)
|
||||
- [RemovalStats](interfaces/RemovalStats.md)
|
||||
- [RenameTableOptions](interfaces/RenameTableOptions.md)
|
||||
- [RestNamespaceConfig](interfaces/RestNamespaceConfig.md)
|
||||
|
||||
@@ -1,23 +0,0 @@
|
||||
[**@lancedb/lancedb**](../README.md) • **Docs**
|
||||
|
||||
***
|
||||
|
||||
[@lancedb/lancedb](../globals.md) / RefreshColumnResult
|
||||
|
||||
# Interface: RefreshColumnResult
|
||||
|
||||
## Properties
|
||||
|
||||
### rowsFilled
|
||||
|
||||
```ts
|
||||
rowsFilled: number;
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### version
|
||||
|
||||
```ts
|
||||
version: number;
|
||||
```
|
||||
@@ -8,7 +8,7 @@
|
||||
<parent>
|
||||
<groupId>com.lancedb</groupId>
|
||||
<artifactId>lancedb-parent</artifactId>
|
||||
<version>0.38.0-beta.0</version>
|
||||
<version>0.37.1-beta.1</version>
|
||||
<relativePath>../pom.xml</relativePath>
|
||||
</parent>
|
||||
|
||||
|
||||
+2
-2
@@ -6,7 +6,7 @@
|
||||
|
||||
<groupId>com.lancedb</groupId>
|
||||
<artifactId>lancedb-parent</artifactId>
|
||||
<version>0.38.0-beta.0</version>
|
||||
<version>0.37.1-beta.1</version>
|
||||
<packaging>pom</packaging>
|
||||
<name>${project.artifactId}</name>
|
||||
<description>LanceDB Java SDK Parent POM</description>
|
||||
@@ -28,7 +28,7 @@
|
||||
<properties>
|
||||
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
|
||||
<arrow.version>15.0.0</arrow.version>
|
||||
<lance-core.version>11.0.0-beta.13</lance-core.version>
|
||||
<lance-core.version>11.0.0-beta.3</lance-core.version>
|
||||
<spotless.skip>false</spotless.skip>
|
||||
<spotless.version>2.30.0</spotless.version>
|
||||
<spotless.java.googlejavaformat.version>1.7</spotless.java.googlejavaformat.version>
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[package]
|
||||
name = "lancedb-nodejs"
|
||||
edition.workspace = true
|
||||
version = "0.38.0-beta.0"
|
||||
version = "0.37.1-beta.1"
|
||||
publish = false
|
||||
license.workspace = true
|
||||
description.workspace = true
|
||||
|
||||
@@ -89,16 +89,6 @@ describe("given a connection", () => {
|
||||
await db.createTable("test4", [{ id: 1 }, { id: 2 }]);
|
||||
});
|
||||
|
||||
it("should return a completed job when dropping a local table", async () => {
|
||||
await db.createTable("async-drop", [{ id: 1 }]);
|
||||
|
||||
const job = await db.dropTableAsync("async-drop");
|
||||
expect(job.id).toBeNull();
|
||||
await expect(job.status()).resolves.toBe("finished");
|
||||
await job.wait();
|
||||
await expect(db.tableNames()).resolves.toEqual([]);
|
||||
});
|
||||
|
||||
it("should fail if creating table twice, unless overwrite is true", async () => {
|
||||
let tbl = await db.createTable("test", [{ id: 1 }, { id: 2 }]);
|
||||
await expect(tbl.countRows()).resolves.toBe(2);
|
||||
|
||||
@@ -1001,49 +1001,4 @@ describe("remote connection jobs surface", () => {
|
||||
},
|
||||
);
|
||||
});
|
||||
|
||||
it("addBases posts the bases array", async () => {
|
||||
const postedBodies: unknown[] = [];
|
||||
await withMockDatabase(
|
||||
(req, res) => {
|
||||
const path = req.url ?? "";
|
||||
if (path.endsWith("/describe/")) {
|
||||
res.writeHead(200, { "Content-Type": "application/json" }).end(
|
||||
JSON.stringify({
|
||||
name: "photos",
|
||||
version: 1,
|
||||
schema: { fields: [] },
|
||||
}),
|
||||
);
|
||||
return;
|
||||
}
|
||||
if (path.endsWith("/bases/")) {
|
||||
const chunks: Buffer[] = [];
|
||||
req.on("data", (chunk) => chunks.push(chunk));
|
||||
req.on("end", () => {
|
||||
postedBodies.push(JSON.parse(Buffer.concat(chunks).toString()));
|
||||
res
|
||||
.writeHead(200, { "Content-Type": "application/json" })
|
||||
.end(JSON.stringify({ version: 2 }));
|
||||
});
|
||||
return;
|
||||
}
|
||||
res.writeHead(404).end();
|
||||
},
|
||||
async (db) => {
|
||||
const table = await db.openTable("photos");
|
||||
await table.addBases({ path: "s3://bucket/media/" });
|
||||
},
|
||||
);
|
||||
expect(postedBodies).toEqual([
|
||||
{
|
||||
bases: [
|
||||
{
|
||||
path: "s3://bucket/media/",
|
||||
isDatasetRoot: false,
|
||||
},
|
||||
],
|
||||
},
|
||||
]);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -4,7 +4,6 @@
|
||||
import * as fs from "fs";
|
||||
import * as path from "path";
|
||||
import * as tmp from "tmp";
|
||||
import { pathToFileURL } from "url";
|
||||
|
||||
import * as arrow15 from "apache-arrow-15";
|
||||
import * as arrow16 from "apache-arrow-16";
|
||||
@@ -3341,86 +3340,3 @@ describe("LSM merge insert", () => {
|
||||
await expect(table.query().useLsm(true).toArray()).rejects.toThrow();
|
||||
});
|
||||
});
|
||||
|
||||
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]);
|
||||
});
|
||||
});
|
||||
|
||||
describe("table bases", () => {
|
||||
let tmpDir: tmp.DirResult;
|
||||
beforeEach(() => {
|
||||
tmpDir = tmp.dirSync({ unsafeCleanup: true });
|
||||
});
|
||||
afterEach(() => tmpDir.removeCallback());
|
||||
|
||||
it("addBases accepts a file uri", async () => {
|
||||
const conn = await connect(tmpDir.name);
|
||||
const table = await conn.createEmptyTable(
|
||||
"photos",
|
||||
new arrow.Schema([new arrow.Field("id", new arrow.Int64(), false)]),
|
||||
);
|
||||
const media = path.join(tmpDir.name, "media");
|
||||
fs.mkdirSync(media);
|
||||
await table.addBases(pathToFileURL(media).toString());
|
||||
});
|
||||
});
|
||||
|
||||
@@ -327,14 +327,6 @@ export abstract class Connection {
|
||||
*/
|
||||
abstract dropTable(name: string, namespacePath?: string[]): Promise<void>;
|
||||
|
||||
/**
|
||||
* Start dropping a table and return its cleanup job.
|
||||
*
|
||||
* The table may become unavailable before its data files are removed. Wait
|
||||
* on the returned job to know when cleanup has finished.
|
||||
*/
|
||||
abstract dropTableAsync(name: string, namespacePath?: string[]): Promise<Job>;
|
||||
|
||||
/**
|
||||
* Drop all tables in the database.
|
||||
* @param {string[]} namespacePath The namespace path to drop tables from (defaults to root namespace).
|
||||
@@ -713,10 +705,6 @@ export class LocalConnection extends Connection {
|
||||
return this.inner.dropTable(name, namespacePath ?? []);
|
||||
}
|
||||
|
||||
async dropTableAsync(name: string, namespacePath?: string[]): Promise<Job> {
|
||||
return this.inner.dropTableAsync(name, namespacePath ?? []);
|
||||
}
|
||||
|
||||
async dropAllTables(namespacePath?: string[]): Promise<void> {
|
||||
return this.inner.dropAllTables(namespacePath ?? []);
|
||||
}
|
||||
|
||||
@@ -50,7 +50,6 @@ export {
|
||||
MergeResult,
|
||||
AddResult,
|
||||
AddColumnsResult,
|
||||
RefreshColumnResult,
|
||||
AlterColumnsResult,
|
||||
UpdateFieldMetadataResult,
|
||||
DeleteResult,
|
||||
@@ -130,7 +129,6 @@ export {
|
||||
|
||||
export {
|
||||
Table,
|
||||
TableBase,
|
||||
Branches,
|
||||
BranchColumnSummary,
|
||||
BranchColumnChange,
|
||||
|
||||
+2
-131
@@ -33,7 +33,6 @@ import {
|
||||
Job,
|
||||
Branches as NativeBranches,
|
||||
OptimizeStats,
|
||||
RefreshColumnResult,
|
||||
TableStatistics,
|
||||
Tags,
|
||||
UpdateFieldMetadataResult,
|
||||
@@ -78,25 +77,6 @@ export interface WriteProgress {
|
||||
done: boolean;
|
||||
}
|
||||
|
||||
/**
|
||||
* An extra storage prefix registered on a table.
|
||||
*
|
||||
* `path` is an object-store URI. `name` is an optional alias. `isDatasetRoot`
|
||||
* is true when `path` points to a Lance dataset root. When false, `path`
|
||||
* points directly to the directory containing the referenced files.
|
||||
*/
|
||||
export interface TableBase {
|
||||
/** Object store URI such as `s3://bucket/media/`. */
|
||||
path: string;
|
||||
/** Optional alias. */
|
||||
name?: string;
|
||||
/**
|
||||
* True when `path` is a Lance dataset root. When false, `path` is the
|
||||
* directory containing the referenced files.
|
||||
*/
|
||||
isDatasetRoot?: boolean;
|
||||
}
|
||||
|
||||
/**
|
||||
* Options for adding data to a table.
|
||||
*/
|
||||
@@ -545,84 +525,18 @@ export abstract class Table {
|
||||
abstract vectorSearch(vector: IntoVector | MultiVector): VectorQuery;
|
||||
/**
|
||||
* Add new columns with defined values.
|
||||
*
|
||||
* The `{ computed }` form stores the expression rather than evaluating it
|
||||
* now: the column is committed with no values, and rows get them from
|
||||
* {@link Table#refreshColumn}. Declaring one therefore costs the same on a
|
||||
* large table as on an empty one.
|
||||
*
|
||||
* A refresh does not revisit rows it has already filled, so mutating an
|
||||
* input leaves the value computed at fill time; recomputing means dropping
|
||||
* the column and declaring it again. While a declaration reads a column,
|
||||
* that column cannot be renamed, retyped or dropped.
|
||||
*
|
||||
* On LanceDB Cloud and Enterprise the expression is planned by the
|
||||
* server, and the refresh runs as a server job -- see
|
||||
* {@link Table#refreshColumnAsync}.
|
||||
* @param {AddColumnsSql[] | Field | Field[] | Schema} newColumnTransforms Either:
|
||||
* - An array of objects with column names and SQL expressions to calculate values
|
||||
* - A single Arrow Field defining one column with its data type (column will be initialized with null values)
|
||||
* - An array of Arrow Fields defining columns with their data types (columns will be initialized with null values)
|
||||
* - An Arrow Schema defining columns with their data types (columns will be initialized with null values)
|
||||
* - `{ computed }`, declaring columns defined by a SQL expression whose type and inputs are derived from it
|
||||
* @returns {Promise<AddColumnsResult>} A promise that resolves to an object
|
||||
* containing the new version number of the table after adding the columns.
|
||||
* @example
|
||||
* ```ts
|
||||
* await table.addColumns({ computed: [{ name: "doubled", valueSql: "x * 2" }] });
|
||||
* const { rowsFilled } = await table.refreshColumn("doubled");
|
||||
* ```
|
||||
*/
|
||||
abstract addColumns(
|
||||
newColumnTransforms:
|
||||
| AddColumnsSql[]
|
||||
| Field
|
||||
| Field[]
|
||||
| Schema
|
||||
| { computed: AddColumnsSql[] },
|
||||
newColumnTransforms: AddColumnsSql[] | Field | Field[] | Schema,
|
||||
): Promise<AddColumnsResult>;
|
||||
|
||||
/**
|
||||
* Register additional storage bases for this table.
|
||||
*
|
||||
* A URI string is a non-root base with no alias.
|
||||
*/
|
||||
abstract addBases(
|
||||
bases: string | TableBase | Array<string | TableBase>,
|
||||
): Promise<void>;
|
||||
|
||||
/**
|
||||
* 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
|
||||
@@ -1174,22 +1088,8 @@ export class LocalTable extends Table {
|
||||
// TODO: Support BatchUDF
|
||||
|
||||
async addColumns(
|
||||
newColumnTransforms:
|
||||
| AddColumnsSql[]
|
||||
| Field
|
||||
| Field[]
|
||||
| Schema
|
||||
| { computed: AddColumnsSql[] },
|
||||
newColumnTransforms: AddColumnsSql[] | Field | Field[] | Schema,
|
||||
): Promise<AddColumnsResult> {
|
||||
// Columns defined by an expression are declared, not materialized here.
|
||||
if (
|
||||
typeof newColumnTransforms === "object" &&
|
||||
!Array.isArray(newColumnTransforms) &&
|
||||
"computed" in newColumnTransforms
|
||||
) {
|
||||
return await this.inner.addComputedColumns(newColumnTransforms.computed);
|
||||
}
|
||||
|
||||
// Handle single Field -> convert to array of Fields
|
||||
if (newColumnTransforms instanceof Field) {
|
||||
newColumnTransforms = [newColumnTransforms];
|
||||
@@ -1224,20 +1124,6 @@ export class LocalTable extends Table {
|
||||
throw new Error("Invalid input type for addColumns");
|
||||
}
|
||||
|
||||
async addBases(
|
||||
bases: string | TableBase | Array<string | TableBase>,
|
||||
): Promise<void> {
|
||||
await this.inner.addBases(normalizeBases(bases));
|
||||
}
|
||||
|
||||
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> {
|
||||
@@ -1430,21 +1316,6 @@ export class LocalTable extends Table {
|
||||
}
|
||||
}
|
||||
|
||||
function normalizeBases(
|
||||
bases: string | TableBase | Array<string | TableBase>,
|
||||
): TableBase[] {
|
||||
const baseInputs = Array.isArray(bases) ? bases : [bases];
|
||||
return baseInputs.map((base) =>
|
||||
typeof base === "string"
|
||||
? { path: base, isDatasetRoot: false }
|
||||
: {
|
||||
path: base.path,
|
||||
name: base.name,
|
||||
isDatasetRoot: base.isDatasetRoot ?? false,
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* A definition of a column alteration. The alteration changes the column at
|
||||
* `path` to have the new name `name`, to be nullable if `nullable` is true,
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-darwin-arm64",
|
||||
"version": "0.38.0-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": ["darwin"],
|
||||
"cpu": ["arm64"],
|
||||
"main": "lancedb.darwin-arm64.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-linux-arm64-gnu",
|
||||
"version": "0.38.0-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": ["linux"],
|
||||
"cpu": ["arm64"],
|
||||
"main": "lancedb.linux-arm64-gnu.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-linux-arm64-musl",
|
||||
"version": "0.38.0-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": ["linux"],
|
||||
"cpu": ["arm64"],
|
||||
"main": "lancedb.linux-arm64-musl.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-linux-x64-gnu",
|
||||
"version": "0.38.0-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": ["linux"],
|
||||
"cpu": ["x64"],
|
||||
"main": "lancedb.linux-x64-gnu.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-linux-x64-musl",
|
||||
"version": "0.38.0-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": ["linux"],
|
||||
"cpu": ["x64"],
|
||||
"main": "lancedb.linux-x64-musl.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-win32-arm64-msvc",
|
||||
"version": "0.38.0-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": [
|
||||
"win32"
|
||||
],
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-win32-x64-msvc",
|
||||
"version": "0.38.0-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": ["win32"],
|
||||
"cpu": ["x64"],
|
||||
"main": "lancedb.win32-x64-msvc.node",
|
||||
|
||||
Generated
+2
-2
@@ -1,12 +1,12 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb",
|
||||
"version": "0.38.0-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"lockfileVersion": 3,
|
||||
"requires": true,
|
||||
"packages": {
|
||||
"": {
|
||||
"name": "@lancedb/lancedb",
|
||||
"version": "0.38.0-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"cpu": [
|
||||
"x64",
|
||||
"arm64"
|
||||
|
||||
+1
-1
@@ -11,7 +11,7 @@
|
||||
"ann"
|
||||
],
|
||||
"private": false,
|
||||
"version": "0.38.0-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"main": "dist/index.js",
|
||||
"exports": {
|
||||
".": "./dist/index.js",
|
||||
|
||||
@@ -334,22 +334,6 @@ impl Connection {
|
||||
.default_error()
|
||||
}
|
||||
|
||||
/// Start dropping a table and return its cleanup job.
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn drop_table_async(
|
||||
&self,
|
||||
name: String,
|
||||
namespace_path: Option<Vec<String>>,
|
||||
) -> napi::Result<crate::job::Job> {
|
||||
let ns = namespace_path.unwrap_or_default();
|
||||
let job = self
|
||||
.get_inner()?
|
||||
.drop_table_async(&name, &ns)
|
||||
.await
|
||||
.default_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn drop_all_tables(&self, namespace_path: Option<Vec<String>>) -> napi::Result<()> {
|
||||
let ns = namespace_path.unwrap_or_default();
|
||||
|
||||
+11
-1
@@ -42,9 +42,19 @@ impl Job {
|
||||
}
|
||||
|
||||
/// Wait until the operation reaches a terminal state.
|
||||
///
|
||||
/// Jobs that complete without a resource result resolve successfully.
|
||||
/// Resource results are not exposed on this binding yet; unsupported
|
||||
/// success results reject with a generic error.
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn wait(&self) -> napi::Result<()> {
|
||||
self.inner.wait().await.default_error()
|
||||
match self.inner.wait().await.default_error()? {
|
||||
lancedb::JobResult::None => Ok(()),
|
||||
// JobResult is non_exhaustive; Function and future variants fail closed.
|
||||
_ => Err(napi::Error::from_reason(
|
||||
"unsupported job result".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Request cancellation. Cancelling a finished operation is a no-op.
|
||||
|
||||
@@ -10,7 +10,6 @@ use lancedb::table::{
|
||||
AddDataMode, ColumnAlteration as LanceColumnAlteration, Duration,
|
||||
FieldMetadataUpdate as LanceFieldMetadataUpdate, FtsToken as LanceDbFtsToken,
|
||||
NewColumnTransform, OptimizeAction, OptimizeOptions, Ref, Table as LanceDbTable,
|
||||
TableBase as LanceTableBase,
|
||||
};
|
||||
use napi::bindgen_prelude::*;
|
||||
use napi::threadsafe_function::{ThreadsafeFunction, ThreadsafeFunctionCallMode};
|
||||
@@ -348,40 +347,6 @@ impl Table {
|
||||
Ok(res.into())
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn add_computed_columns(
|
||||
&self,
|
||||
columns: Vec<AddColumnsSql>,
|
||||
) -> napi::Result<AddColumnsResult> {
|
||||
let table = self.inner_ref()?;
|
||||
let mut builder = table.add_columns();
|
||||
for column in columns {
|
||||
builder = builder.computed(column.name, column.value_sql);
|
||||
}
|
||||
let res = builder.execute().await.default_error()?;
|
||||
Ok(res.into())
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn refresh_column(&self, column: String) -> napi::Result<RefreshColumnResult> {
|
||||
let res = self
|
||||
.inner_ref()?
|
||||
.refresh_column(column)
|
||||
.await
|
||||
.default_error()?;
|
||||
Ok(res.into())
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn refresh_column_async(&self, column: String) -> napi::Result<crate::job::Job> {
|
||||
let job = self
|
||||
.inner_ref()?
|
||||
.refresh_column_async(column)
|
||||
.await
|
||||
.default_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn add_columns_with_schema(
|
||||
&self,
|
||||
@@ -447,18 +412,6 @@ impl Table {
|
||||
Ok(res.into())
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn add_bases(&self, bases: Vec<TableBase>) -> napi::Result<()> {
|
||||
self.inner_ref()?
|
||||
.add_bases(bases.into_iter().map(|base| LanceTableBase {
|
||||
path: base.path,
|
||||
name: base.name,
|
||||
is_dataset_root: base.is_dataset_root,
|
||||
}))
|
||||
.await
|
||||
.default_error()
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn drop_columns(&self, columns: Vec<String>) -> napi::Result<DropColumnsResult> {
|
||||
let col_refs = columns.iter().map(String::as_str).collect::<Vec<_>>();
|
||||
@@ -713,18 +666,6 @@ impl Table {
|
||||
}
|
||||
}
|
||||
|
||||
#[napi(object)]
|
||||
/// An extra storage prefix registered on a table.
|
||||
pub struct TableBase {
|
||||
/// Object store URI such as `s3://bucket/media/`.
|
||||
pub path: String,
|
||||
/// Optional alias.
|
||||
pub name: Option<String>,
|
||||
/// True when `path` is a Lance dataset root. When false, `path` is the
|
||||
/// directory containing the referenced files.
|
||||
pub is_dataset_root: bool,
|
||||
}
|
||||
|
||||
#[napi(object)]
|
||||
/// A description of an index currently configured on a column
|
||||
pub struct IndexConfig {
|
||||
@@ -1255,21 +1196,6 @@ pub struct AddColumnsResult {
|
||||
pub version: i64,
|
||||
}
|
||||
|
||||
#[napi(object)]
|
||||
pub struct RefreshColumnResult {
|
||||
pub rows_filled: i64,
|
||||
pub version: i64,
|
||||
}
|
||||
|
||||
impl From<lancedb::table::RefreshColumnResult> for RefreshColumnResult {
|
||||
fn from(value: lancedb::table::RefreshColumnResult) -> Self {
|
||||
Self {
|
||||
rows_filled: value.rows_filled as i64,
|
||||
version: value.version as i64,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<lancedb::table::AddColumnsResult> for AddColumnsResult {
|
||||
fn from(value: lancedb::table::AddColumnsResult) -> Self {
|
||||
Self {
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "lancedb-python"
|
||||
version = "0.38.0-beta.0"
|
||||
version = "0.37.1-beta.1"
|
||||
publish = false
|
||||
edition.workspace = true
|
||||
description = "Python bindings for LanceDB"
|
||||
|
||||
@@ -12,6 +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 tokenize as _tokenize
|
||||
from .common import URI, sanitize_uri
|
||||
from urllib.parse import urlparse
|
||||
@@ -21,8 +22,9 @@ from .remote.db import RemoteDBConnection
|
||||
from .expr import Expr, col, lit, func
|
||||
from .schema import blob, vector, BlobType
|
||||
from .job import AsyncJob, Job
|
||||
from .table import AsyncTable, Table, TableBase
|
||||
from .table import AsyncTable, Table
|
||||
from .types import BaseTokenizerType
|
||||
from ._udf import FunctionCapability, udf
|
||||
from ._lancedb import Session
|
||||
from .namespace import (
|
||||
connect_namespace,
|
||||
@@ -507,6 +509,8 @@ __all__ = [
|
||||
"FtsToken",
|
||||
"col",
|
||||
"Expr",
|
||||
"Function",
|
||||
"FunctionCapability",
|
||||
"func",
|
||||
"lit",
|
||||
"URI",
|
||||
@@ -521,6 +525,6 @@ __all__ = [
|
||||
"RemoteDBConnection",
|
||||
"Session",
|
||||
"Table",
|
||||
"TableBase",
|
||||
"udf",
|
||||
"__version__",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Private first-class Function namespace facades for database connections.
|
||||
|
||||
These helpers are internal submission and lookup surfaces. They are not durable
|
||||
resources and are not part of the public top-level export surface.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from . import _udf
|
||||
from ._lancedb import Function
|
||||
from .job import AsyncJob, Job
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .db import AsyncConnection, DBConnection
|
||||
|
||||
|
||||
class _SyncFunctions:
|
||||
"""Synchronous `db.functions` facade."""
|
||||
|
||||
__slots__ = ("_connection",)
|
||||
|
||||
def __init__(self, connection: DBConnection) -> None:
|
||||
self._connection = connection
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return "_SyncFunctions()"
|
||||
|
||||
def register(self, name: str, decorated_udf: Callable[..., object]) -> Job:
|
||||
"""Register a decorated UDF and return a synchronous [Job][lancedb.job.Job]."""
|
||||
definition = _udf._build_function_definition(decorated_udf)
|
||||
native_job = self._connection._submit_register_function(name, definition)
|
||||
return Job(AsyncJob(native_job))
|
||||
|
||||
def replace(
|
||||
self, name: str, current: Function, decorated_udf: Callable[..., object]
|
||||
) -> Job:
|
||||
"""Conditionally replace a Function; return sync [Job][lancedb.job.Job]."""
|
||||
definition = _udf._build_function_definition(decorated_udf)
|
||||
native_job = self._connection._submit_replace_function(
|
||||
name, current, definition
|
||||
)
|
||||
return Job(AsyncJob(native_job))
|
||||
|
||||
def get(self, name: str) -> Function:
|
||||
"""Return the Function currently bound to a database-scoped name."""
|
||||
return self._connection._lookup_function_by_name(name)
|
||||
|
||||
def get_by_id(self, function_id: str) -> Function:
|
||||
"""Return the immutable Function for an exact Function ID."""
|
||||
return self._connection._lookup_function_by_id(function_id)
|
||||
|
||||
def remove(self, name: str, current: Function) -> None:
|
||||
"""Conditionally remove a Function catalog name binding."""
|
||||
return self._connection._remove_function_name(name, current)
|
||||
|
||||
def revoke(self, function: Function) -> None:
|
||||
"""Revoke an exact immutable Function by administrator set-bit."""
|
||||
return self._connection._revoke_function(function)
|
||||
|
||||
|
||||
class _AsyncFunctions:
|
||||
"""Asynchronous `async_db.functions` facade."""
|
||||
|
||||
__slots__ = ("_connection",)
|
||||
|
||||
def __init__(self, connection: AsyncConnection) -> None:
|
||||
self._connection = connection
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return "_AsyncFunctions()"
|
||||
|
||||
async def register(
|
||||
self, name: str, decorated_udf: Callable[..., object]
|
||||
) -> AsyncJob:
|
||||
"""Register a decorated UDF and return an [AsyncJob][lancedb.job.AsyncJob]."""
|
||||
definition = _udf._build_function_definition(decorated_udf)
|
||||
native_job = await self._connection._register_function(name, definition)
|
||||
return AsyncJob(native_job)
|
||||
|
||||
async def replace(
|
||||
self, name: str, current: Function, decorated_udf: Callable[..., object]
|
||||
) -> AsyncJob:
|
||||
"""Conditionally replace a Function; return [AsyncJob][lancedb.job.AsyncJob]."""
|
||||
definition = _udf._build_function_definition(decorated_udf)
|
||||
native_job = await self._connection._replace_function(name, current, definition)
|
||||
return AsyncJob(native_job)
|
||||
|
||||
async def get(self, name: str) -> Function:
|
||||
"""Return the Function currently bound to a database-scoped name."""
|
||||
return await self._connection._lookup_function_by_name(name)
|
||||
|
||||
async def get_by_id(self, function_id: str) -> Function:
|
||||
"""Return the immutable Function for an exact Function ID."""
|
||||
return await self._connection._lookup_function_by_id(function_id)
|
||||
|
||||
async def remove(self, name: str, current: Function) -> None:
|
||||
"""Conditionally remove a Function catalog name binding."""
|
||||
return await self._connection._remove_function_name(name, current)
|
||||
|
||||
async def revoke(self, function: Function) -> None:
|
||||
"""Revoke an exact immutable Function by administrator set-bit."""
|
||||
return await self._connection._revoke_function(function)
|
||||
@@ -153,6 +153,16 @@ class Connection(object):
|
||||
async def job_history(
|
||||
self, job_id: Optional[str] = None
|
||||
) -> List[pa.RecordBatch]: ...
|
||||
async def _register_function(
|
||||
self, name: str, definition: "_FunctionDefinition"
|
||||
) -> Job: ...
|
||||
async def _replace_function(
|
||||
self, name: str, current: Function, definition: "_FunctionDefinition"
|
||||
) -> Job: ...
|
||||
async def _lookup_function_by_name(self, name: str) -> Function: ...
|
||||
async def _lookup_function_by_id(self, function_id: str) -> Function: ...
|
||||
async def _remove_function_name(self, name: str, current: Function) -> None: ...
|
||||
async def _revoke_function(self, function: Function) -> None: ...
|
||||
async def create_table(
|
||||
self,
|
||||
name: str,
|
||||
@@ -198,9 +208,6 @@ class Connection(object):
|
||||
async def drop_table(
|
||||
self, name: str, namespace_path: Optional[List[str]] = None
|
||||
) -> None: ...
|
||||
async def drop_table_async(
|
||||
self, name: str, namespace_path: Optional[List[str]] = None
|
||||
) -> Job: ...
|
||||
async def drop_all_tables(
|
||||
self, namespace_path: Optional[List[str]] = None
|
||||
) -> None: ...
|
||||
@@ -219,11 +226,45 @@ class BlobFile:
|
||||
def read_range(self, offset: int, length: int) -> bytes: ...
|
||||
def read_up_to(self, length: int) -> bytes: ...
|
||||
|
||||
class Function:
|
||||
@property
|
||||
def id(self) -> str: ...
|
||||
@property
|
||||
def parameters(self) -> tuple[tuple[str, pa.DataType], ...]: ...
|
||||
@property
|
||||
def output_type(self) -> pa.DataType: ...
|
||||
@property
|
||||
def output_nullable(self) -> bool: ...
|
||||
def __call__(self, **kwargs: Any) -> "_FunctionCall": ...
|
||||
|
||||
class _FunctionCall:
|
||||
"""Private unresolved Function call authoring value (FF-028)."""
|
||||
|
||||
...
|
||||
|
||||
class _FunctionDefinition:
|
||||
"""Private owner of the Rust FunctionDefinition registration input."""
|
||||
|
||||
def _to_json(self) -> str: ...
|
||||
|
||||
def _new_function_definition(
|
||||
*,
|
||||
parameters: list[tuple[str, pa.DataType]],
|
||||
output_type: pa.DataType,
|
||||
output_nullable: bool,
|
||||
module: str,
|
||||
callable_name: str,
|
||||
source: str,
|
||||
python: str,
|
||||
packages: list[str],
|
||||
capabilities: list[tuple[str, str, Optional[str]]],
|
||||
) -> _FunctionDefinition: ...
|
||||
|
||||
class Job:
|
||||
@property
|
||||
def id(self) -> Optional[str]: ...
|
||||
async def status(self) -> str: ...
|
||||
async def wait(self) -> None: ...
|
||||
async def wait(self) -> Optional[Function]: ...
|
||||
async def cancel(self) -> None: ...
|
||||
|
||||
class JobInfo:
|
||||
@@ -245,6 +286,8 @@ class JobFailureInfo:
|
||||
def message(self) -> Optional[str]: ...
|
||||
@property
|
||||
def retryable(self) -> Optional[bool]: ...
|
||||
@property
|
||||
def error_code(self) -> Optional[str]: ...
|
||||
|
||||
class JobDescription:
|
||||
@property
|
||||
@@ -259,6 +302,8 @@ class JobDescription:
|
||||
def spec_json(self) -> Optional[str]: ...
|
||||
@property
|
||||
def failure(self) -> Optional[JobFailureInfo]: ...
|
||||
@property
|
||||
def result(self) -> Optional[Function]: ...
|
||||
|
||||
class Table:
|
||||
def name(self) -> str: ...
|
||||
@@ -321,6 +366,16 @@ class Table:
|
||||
name: Optional[str],
|
||||
train: Optional[bool],
|
||||
) -> Job: ...
|
||||
async def _add_generated_column(
|
||||
self, column_name: str, call: _FunctionCall
|
||||
) -> Job: ...
|
||||
async def _generated_column_status(
|
||||
self, column_name: str
|
||||
) -> Literal["complete", "incomplete"]: ...
|
||||
async def _refresh_generated_column(self, column_name: str) -> Job: ...
|
||||
async def _alter_generated_column(
|
||||
self, column_name: str, new_call: _FunctionCall
|
||||
) -> Job: ...
|
||||
async def list_versions(self) -> List[Dict[str, Any]]: ...
|
||||
async def version(self) -> int: ...
|
||||
async def checkout(self, version: Union[int, str]): ...
|
||||
@@ -338,11 +393,6 @@ class Table:
|
||||
) -> list[FtsToken]: ...
|
||||
async def delete(self, filter: Union[str, PyExpr]) -> DeleteResult: ...
|
||||
async def add_columns(self, columns: list[tuple[str, str]]) -> AddColumnsResult: ...
|
||||
async def add_computed_columns(
|
||||
self, columns: list[tuple[str, str]]
|
||||
) -> AddColumnsResult: ...
|
||||
async def refresh_column(self, column: str) -> RefreshColumnResult: ...
|
||||
async def refresh_column_async(self, column: str) -> Job: ...
|
||||
async def add_columns_with_schema(self, schema: pa.Schema) -> AddColumnsResult: ...
|
||||
async def alter_columns(
|
||||
self, columns: list[dict[str, Any]]
|
||||
@@ -377,7 +427,6 @@ class Table:
|
||||
def take_offsets(self, offsets: list[int]) -> TakeQuery: ...
|
||||
def take_row_ids(self, row_ids: list[int]) -> TakeQuery: ...
|
||||
async def blob_columns(self) -> list[str]: ...
|
||||
async def add_bases(self, bases: list[Any]) -> None: ...
|
||||
async def fetch_blobs(
|
||||
self, column: str, row_ids: list[int]
|
||||
) -> pa.LargeBinaryArray: ...
|
||||
@@ -689,10 +738,6 @@ class LsmWriteSpec:
|
||||
class AddColumnsResult:
|
||||
version: int
|
||||
|
||||
class RefreshColumnResult:
|
||||
rows_filled: int
|
||||
version: int
|
||||
|
||||
class AlterColumnsResult:
|
||||
version: int
|
||||
|
||||
|
||||
@@ -0,0 +1,538 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Local authoring declaration surface for first-class UDFs.
|
||||
|
||||
This module snapshots declaration metadata onto a Python function and privately
|
||||
validates packagable callables into a source snapshot. It does not mint durable
|
||||
identity or register anything with a database.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
import inspect
|
||||
import stat
|
||||
import symtable
|
||||
import sys
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from types import CodeType, FunctionType
|
||||
from typing import NoReturn, ParamSpec, TypeVar
|
||||
|
||||
import pyarrow as pa
|
||||
|
||||
from . import _lancedb
|
||||
|
||||
__all__ = ["FunctionCapability", "udf"]
|
||||
|
||||
_P = ParamSpec("_P")
|
||||
_R = TypeVar("_R")
|
||||
|
||||
_CONFIG_ATTR = "__lancedb_udf_config__"
|
||||
_SYNTHETIC_SOURCE_FILENAME = "<lancedb-udf>"
|
||||
_PACKAGING_ERROR = "udf is not packagable"
|
||||
_ALLOWED_PARAM_KINDS = (
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
inspect.Parameter.KEYWORD_ONLY,
|
||||
)
|
||||
|
||||
|
||||
class FunctionCapability:
|
||||
"""Local capability declaration for a first-class UDF.
|
||||
|
||||
Construct via :meth:`network` or :meth:`secret`. Direct construction is
|
||||
rejected so callers cannot create an uninitialized capability.
|
||||
"""
|
||||
|
||||
__slots__ = ("_kind", "_origin", "_reference", "_environment_variable")
|
||||
|
||||
def __new__(cls, *args: object, **kwargs: object) -> FunctionCapability:
|
||||
raise TypeError(
|
||||
"FunctionCapability cannot be constructed directly; "
|
||||
"use FunctionCapability.network() or FunctionCapability.secret()"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _create(
|
||||
cls,
|
||||
kind: str,
|
||||
origin: str | None,
|
||||
reference: str | None,
|
||||
environment_variable: str | None,
|
||||
) -> FunctionCapability:
|
||||
obj = object.__new__(cls)
|
||||
object.__setattr__(obj, "_kind", kind)
|
||||
object.__setattr__(obj, "_origin", origin)
|
||||
object.__setattr__(obj, "_reference", reference)
|
||||
object.__setattr__(obj, "_environment_variable", environment_variable)
|
||||
return obj
|
||||
|
||||
@classmethod
|
||||
def network(cls, origin: str) -> FunctionCapability:
|
||||
if not isinstance(origin, str):
|
||||
raise TypeError("origin must be a string")
|
||||
if origin == "":
|
||||
raise ValueError("origin must be non-empty")
|
||||
return cls._create("network", origin, None, None)
|
||||
|
||||
@classmethod
|
||||
def secret(cls, reference: str, *, environment_variable: str) -> FunctionCapability:
|
||||
if not isinstance(reference, str):
|
||||
raise TypeError("reference must be a string")
|
||||
if not isinstance(environment_variable, str):
|
||||
raise TypeError("environment_variable must be a string")
|
||||
if reference == "":
|
||||
raise ValueError("reference must be non-empty")
|
||||
if environment_variable == "":
|
||||
raise ValueError("environment_variable must be non-empty")
|
||||
return cls._create("secret", None, reference, environment_variable)
|
||||
|
||||
@property
|
||||
def kind(self) -> str:
|
||||
return self._kind
|
||||
|
||||
@property
|
||||
def origin(self) -> str | None:
|
||||
return self._origin
|
||||
|
||||
@property
|
||||
def reference(self) -> str | None:
|
||||
return self._reference
|
||||
|
||||
@property
|
||||
def environment_variable(self) -> str | None:
|
||||
return self._environment_variable
|
||||
|
||||
def __setattr__(self, name: str, value: object) -> None:
|
||||
raise AttributeError(
|
||||
f"{type(self).__name__!r} object attribute {name!r} is read-only"
|
||||
)
|
||||
|
||||
def __delattr__(self, name: str) -> None:
|
||||
raise AttributeError(
|
||||
f"{type(self).__name__!r} object attribute {name!r} is read-only"
|
||||
)
|
||||
|
||||
def __eq__(self, other: object) -> bool:
|
||||
if not isinstance(other, FunctionCapability):
|
||||
return NotImplemented
|
||||
return (
|
||||
self._kind == other._kind
|
||||
and self._origin == other._origin
|
||||
and self._reference == other._reference
|
||||
and self._environment_variable == other._environment_variable
|
||||
)
|
||||
|
||||
def __hash__(self) -> int:
|
||||
return hash(
|
||||
(
|
||||
self._kind,
|
||||
self._origin,
|
||||
self._reference,
|
||||
self._environment_variable,
|
||||
)
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
if self._kind == "network":
|
||||
return f"FunctionCapability.network({self._origin!r})"
|
||||
return (
|
||||
"FunctionCapability.secret("
|
||||
f"environment_variable={self._environment_variable!r})"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _UdfConfig:
|
||||
"""Private frozen snapshot of a ``@udf`` declaration."""
|
||||
|
||||
inputs: tuple[tuple[str, pa.DataType], ...]
|
||||
output: pa.DataType
|
||||
output_nullable: bool
|
||||
python: str
|
||||
packages: tuple[str, ...]
|
||||
capabilities: tuple[FunctionCapability, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PackagedUdf:
|
||||
"""Private frozen snapshot of a validated packagable UDF."""
|
||||
|
||||
source: str
|
||||
module: str
|
||||
callable_name: str
|
||||
config: _UdfConfig
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return (
|
||||
f"_PackagedUdf(source=<redacted>, module={self.module!r}, "
|
||||
f"callable_name={self.callable_name!r}, config={self.config!r})"
|
||||
)
|
||||
|
||||
|
||||
def _validate_inputs(
|
||||
inputs: object,
|
||||
) -> tuple[tuple[str, pa.DataType], ...]:
|
||||
if not isinstance(inputs, Mapping):
|
||||
raise TypeError("udf inputs must be a Mapping of name to pyarrow DataType")
|
||||
snapshot: list[tuple[str, pa.DataType]] = []
|
||||
for key, value in inputs.items():
|
||||
if not isinstance(key, str):
|
||||
raise TypeError("udf input names must be strings")
|
||||
if key == "":
|
||||
raise ValueError("udf input names must be non-empty")
|
||||
if not isinstance(value, pa.DataType):
|
||||
raise TypeError("udf input types must be pyarrow DataType values")
|
||||
snapshot.append((key, value))
|
||||
return tuple(snapshot)
|
||||
|
||||
|
||||
def _validate_packages(packages: object) -> tuple[str, ...]:
|
||||
if isinstance(packages, (str, bytes, bytearray)):
|
||||
raise TypeError("udf packages must be a sequence of strings, not a string")
|
||||
if not isinstance(packages, Sequence):
|
||||
raise TypeError("udf packages must be a sequence of strings")
|
||||
snapshot: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for package in packages:
|
||||
if not isinstance(package, str):
|
||||
raise TypeError("udf packages must contain only strings")
|
||||
if package == "":
|
||||
raise ValueError("udf packages must be non-empty strings")
|
||||
if package in seen:
|
||||
raise ValueError(f"duplicate udf package: {package}")
|
||||
seen.add(package)
|
||||
snapshot.append(package)
|
||||
return tuple(snapshot)
|
||||
|
||||
|
||||
def _reject_non_exact_capability() -> NoReturn:
|
||||
# Exact-type only: subclasses are authoring inputs we never accept. Keep the
|
||||
# message fixed so hostile markers never enter exception text.
|
||||
raise TypeError(
|
||||
"udf capabilities must contain only FunctionCapability values"
|
||||
) from None
|
||||
|
||||
|
||||
def _require_exact_capability(capability: object) -> FunctionCapability:
|
||||
if type(capability) is not FunctionCapability:
|
||||
_reject_non_exact_capability()
|
||||
return capability
|
||||
|
||||
|
||||
def _validate_capabilities(capabilities: object) -> tuple[FunctionCapability, ...]:
|
||||
if isinstance(capabilities, (str, bytes, bytearray)):
|
||||
raise TypeError(
|
||||
"udf capabilities must be a sequence of FunctionCapability, not a string"
|
||||
)
|
||||
if not isinstance(capabilities, Sequence):
|
||||
raise TypeError("udf capabilities must be a sequence of FunctionCapability")
|
||||
return tuple(_require_exact_capability(capability) for capability in capabilities)
|
||||
|
||||
|
||||
def udf(
|
||||
*,
|
||||
inputs: Mapping[str, pa.DataType],
|
||||
output: pa.DataType,
|
||||
python: str,
|
||||
packages: Sequence[str] = (),
|
||||
output_nullable: bool = True,
|
||||
capabilities: Sequence[FunctionCapability] = (),
|
||||
) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]:
|
||||
"""Declare a local UDF without packaging or registration.
|
||||
|
||||
Applying the returned decorator attaches a private frozen config snapshot
|
||||
and returns the exact same function object.
|
||||
"""
|
||||
input_snapshot = _validate_inputs(inputs)
|
||||
if not isinstance(output, pa.DataType):
|
||||
raise TypeError("udf output must be a pyarrow DataType")
|
||||
if not isinstance(python, str):
|
||||
raise TypeError("udf python must be a string")
|
||||
if python == "":
|
||||
raise ValueError("udf python must be a non-empty string")
|
||||
package_snapshot = _validate_packages(packages)
|
||||
if not isinstance(output_nullable, bool):
|
||||
raise TypeError("udf output_nullable must be a bool")
|
||||
capability_snapshot = _validate_capabilities(capabilities)
|
||||
|
||||
config = _UdfConfig(
|
||||
inputs=input_snapshot,
|
||||
output=output,
|
||||
output_nullable=output_nullable,
|
||||
python=python,
|
||||
packages=package_snapshot,
|
||||
capabilities=capability_snapshot,
|
||||
)
|
||||
|
||||
def decorator(fn: Callable[_P, _R]) -> Callable[_P, _R]:
|
||||
if not inspect.isfunction(fn):
|
||||
raise TypeError("udf can only decorate a Python function")
|
||||
if hasattr(fn, _CONFIG_ATTR):
|
||||
raise ValueError("function is already decorated with @udf")
|
||||
setattr(fn, _CONFIG_ATTR, config)
|
||||
return fn
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def _get_udf_config(fn: object) -> _UdfConfig:
|
||||
"""Return the private declaration snapshot for a ``@udf``-decorated function."""
|
||||
config = getattr(fn, _CONFIG_ATTR, None)
|
||||
if config is None:
|
||||
raise TypeError("function is not decorated with @udf")
|
||||
if not isinstance(config, _UdfConfig):
|
||||
raise TypeError("function is not decorated with @udf")
|
||||
return config
|
||||
|
||||
|
||||
def _packaging_reject() -> NoReturn:
|
||||
raise ValueError(_PACKAGING_ERROR) from None
|
||||
|
||||
|
||||
def _is_ordinary_function(fn: FunctionType) -> bool:
|
||||
if fn.__name__ == "<lambda>":
|
||||
return False
|
||||
if fn.__qualname__ != fn.__name__:
|
||||
return False
|
||||
if inspect.iscoroutinefunction(fn) or inspect.isasyncgenfunction(fn):
|
||||
return False
|
||||
if inspect.isgeneratorfunction(fn):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _resolve_source_path(fn: FunctionType, module: object) -> Path:
|
||||
try:
|
||||
fn_source: str | None = inspect.getsourcefile(fn)
|
||||
except TypeError:
|
||||
fn_source = None
|
||||
source_lookup_failed = True
|
||||
else:
|
||||
source_lookup_failed = False
|
||||
if source_lookup_failed:
|
||||
_packaging_reject()
|
||||
module_file = vars(module).get("__file__")
|
||||
if not fn_source or not isinstance(module_file, str) or module_file == "":
|
||||
_packaging_reject()
|
||||
try:
|
||||
resolved_paths: tuple[Path, Path] | None = (
|
||||
Path(fn_source).resolve(),
|
||||
Path(module_file).resolve(),
|
||||
)
|
||||
except (OSError, RuntimeError):
|
||||
resolved_paths = None
|
||||
if resolved_paths is None:
|
||||
_packaging_reject()
|
||||
fn_path, module_path = resolved_paths
|
||||
if fn_path != module_path:
|
||||
_packaging_reject()
|
||||
if fn_path.suffix != ".py":
|
||||
_packaging_reject()
|
||||
try:
|
||||
mode: int | None = fn_path.stat().st_mode
|
||||
except OSError:
|
||||
mode = None
|
||||
if mode is None:
|
||||
_packaging_reject()
|
||||
if not stat.S_ISREG(mode):
|
||||
_packaging_reject()
|
||||
return fn_path
|
||||
|
||||
|
||||
def _validate_source(
|
||||
source: str, callable_name: str
|
||||
) -> tuple[CodeType, symtable.SymbolTable]:
|
||||
try:
|
||||
module_code = compile(
|
||||
source,
|
||||
_SYNTHETIC_SOURCE_FILENAME,
|
||||
"exec",
|
||||
optimize=sys.flags.optimize,
|
||||
)
|
||||
ast.parse(source, filename=_SYNTHETIC_SOURCE_FILENAME, mode="exec")
|
||||
table = symtable.symtable(source, _SYNTHETIC_SOURCE_FILENAME, "exec")
|
||||
parsed: tuple[CodeType, symtable.SymbolTable] | None = (module_code, table)
|
||||
except Exception:
|
||||
parsed = None
|
||||
if parsed is None:
|
||||
_packaging_reject()
|
||||
module_code, table = parsed
|
||||
|
||||
for child in table.get_children():
|
||||
if child.get_name() == callable_name and child.get_type() == "function":
|
||||
return module_code, table
|
||||
_packaging_reject()
|
||||
|
||||
|
||||
def _source_bound_names(table: symtable.SymbolTable) -> set[str]:
|
||||
names: set[str] = set()
|
||||
for symbol in table.get_symbols():
|
||||
if symbol.is_imported() or symbol.is_assigned() or symbol.is_namespace():
|
||||
names.add(symbol.get_name())
|
||||
return names
|
||||
|
||||
|
||||
def _code_fingerprint(code: CodeType) -> tuple[object, ...]:
|
||||
"""Structural fingerprint ignoring only location/debug fields."""
|
||||
constants = tuple(
|
||||
_code_fingerprint(constant) if isinstance(constant, CodeType) else constant
|
||||
for constant in code.co_consts
|
||||
)
|
||||
return (
|
||||
code.co_name,
|
||||
getattr(code, "co_qualname", code.co_name),
|
||||
code.co_argcount,
|
||||
code.co_posonlyargcount,
|
||||
code.co_kwonlyargcount,
|
||||
code.co_flags,
|
||||
code.co_code,
|
||||
code.co_names,
|
||||
code.co_varnames,
|
||||
code.co_freevars,
|
||||
code.co_cellvars,
|
||||
getattr(code, "co_exceptiontable", b""),
|
||||
constants,
|
||||
)
|
||||
|
||||
|
||||
def _toplevel_code_candidates(
|
||||
module_code: CodeType, callable_name: str
|
||||
) -> list[CodeType]:
|
||||
candidates: list[CodeType] = []
|
||||
for constant in module_code.co_consts:
|
||||
if not isinstance(constant, CodeType):
|
||||
continue
|
||||
if constant.co_name != callable_name:
|
||||
continue
|
||||
if getattr(constant, "co_qualname", callable_name) != callable_name:
|
||||
continue
|
||||
candidates.append(constant)
|
||||
return candidates
|
||||
|
||||
|
||||
def _validate_loaded_code_matches_source(
|
||||
fn: FunctionType, module_code: CodeType
|
||||
) -> None:
|
||||
candidates = _toplevel_code_candidates(module_code, fn.__name__)
|
||||
if not candidates:
|
||||
_packaging_reject()
|
||||
target = _code_fingerprint(fn.__code__)
|
||||
if not any(_code_fingerprint(candidate) == target for candidate in candidates):
|
||||
_packaging_reject()
|
||||
|
||||
|
||||
def _validate_signature(fn: FunctionType, config: _UdfConfig) -> None:
|
||||
try:
|
||||
signature: inspect.Signature | None = inspect.signature(fn)
|
||||
except (TypeError, ValueError):
|
||||
signature = None
|
||||
if signature is None:
|
||||
_packaging_reject()
|
||||
parameters = list(signature.parameters.values())
|
||||
expected = [name for name, _ in config.inputs]
|
||||
actual = [parameter.name for parameter in parameters]
|
||||
if actual != expected:
|
||||
_packaging_reject()
|
||||
for parameter in parameters:
|
||||
if parameter.kind not in _ALLOWED_PARAM_KINDS:
|
||||
_packaging_reject()
|
||||
|
||||
|
||||
def _validate_ambient_globals(fn: FunctionType, table: symtable.SymbolTable) -> None:
|
||||
try:
|
||||
closure_vars: inspect.ClosureVars | None = inspect.getclosurevars(fn)
|
||||
except (TypeError, ValueError):
|
||||
closure_vars = None
|
||||
if closure_vars is None:
|
||||
_packaging_reject()
|
||||
if closure_vars.nonlocals:
|
||||
_packaging_reject()
|
||||
bound_names = _source_bound_names(table)
|
||||
for name in closure_vars.globals:
|
||||
if name not in bound_names:
|
||||
_packaging_reject()
|
||||
|
||||
|
||||
def _package_udf(fn: object) -> _PackagedUdf:
|
||||
"""Validate and snapshot a packagable ``@udf``-decorated function."""
|
||||
config = _get_udf_config(fn)
|
||||
if not isinstance(fn, FunctionType) or not _is_ordinary_function(fn):
|
||||
_packaging_reject()
|
||||
|
||||
module_name = fn.__module__
|
||||
if (
|
||||
not isinstance(module_name, str)
|
||||
or module_name == ""
|
||||
or module_name == "__main__"
|
||||
):
|
||||
_packaging_reject()
|
||||
module = sys.modules.get(module_name)
|
||||
if module is None:
|
||||
_packaging_reject()
|
||||
callable_name = fn.__name__
|
||||
if vars(module).get(callable_name) is not fn:
|
||||
_packaging_reject()
|
||||
|
||||
source_path = _resolve_source_path(fn, module)
|
||||
try:
|
||||
source: str | None = source_path.read_text(encoding="utf-8")
|
||||
except (OSError, UnicodeError):
|
||||
source = None
|
||||
if source is None:
|
||||
_packaging_reject()
|
||||
|
||||
module_code, table = _validate_source(source, callable_name)
|
||||
_validate_signature(fn, config)
|
||||
_validate_ambient_globals(fn, table)
|
||||
_validate_loaded_code_matches_source(fn, module_code)
|
||||
|
||||
return _PackagedUdf(
|
||||
source=source,
|
||||
module=module_name,
|
||||
callable_name=callable_name,
|
||||
config=config,
|
||||
)
|
||||
|
||||
|
||||
def _normalize_capability_triple(
|
||||
capability: FunctionCapability,
|
||||
) -> tuple[str, str, str | None]:
|
||||
"""Normalize a local capability declaration to the native triple shape."""
|
||||
# Private config is untrusted; re-check exact type before any property access.
|
||||
capability = _require_exact_capability(capability)
|
||||
if capability.kind == "network":
|
||||
origin = capability.origin
|
||||
if origin is None:
|
||||
raise ValueError("invalid network capability") from None
|
||||
return ("network", origin, None)
|
||||
if capability.kind == "secret":
|
||||
reference = capability.reference
|
||||
environment_variable = capability.environment_variable
|
||||
if reference is None or environment_variable is None:
|
||||
raise ValueError("invalid secret capability") from None
|
||||
return ("secret", reference, environment_variable)
|
||||
# Fail closed without echoing the unknown kind.
|
||||
raise ValueError("unsupported capability kind") from None
|
||||
|
||||
|
||||
def _build_function_definition(fn: object) -> _lancedb._FunctionDefinition:
|
||||
"""Package a ``@udf`` and bridge it to the private native definition."""
|
||||
packaged = _package_udf(fn)
|
||||
config = packaged.config
|
||||
capabilities = [
|
||||
_normalize_capability_triple(capability) for capability in config.capabilities
|
||||
]
|
||||
return _lancedb._new_function_definition(
|
||||
parameters=list(config.inputs),
|
||||
output_type=config.output,
|
||||
output_nullable=config.output_nullable,
|
||||
module=packaged.module,
|
||||
callable_name=packaged.callable_name,
|
||||
source=packaged.source,
|
||||
python=config.python,
|
||||
packages=list(config.packages),
|
||||
capabilities=capabilities,
|
||||
)
|
||||
+126
-37
@@ -63,8 +63,12 @@ if TYPE_CHECKING:
|
||||
import pyarrow as pa
|
||||
from .pydantic import LanceModel
|
||||
|
||||
from ._functions import _AsyncFunctions, _SyncFunctions
|
||||
from ._lancedb import Connection as LanceDbConnection
|
||||
from ._lancedb import Function
|
||||
from ._lancedb import Job as NativeJob
|
||||
from ._lancedb import JobDescription, JobInfo
|
||||
from ._lancedb import _FunctionDefinition
|
||||
from .common import DATA, URI
|
||||
from .embeddings import EmbeddingFunctionConfig
|
||||
from ._lancedb import Session
|
||||
@@ -524,12 +528,6 @@ class DBConnection(EnforceOverrides):
|
||||
namespace_path = []
|
||||
raise NotImplementedError
|
||||
|
||||
def drop_table_async(
|
||||
self, name: str, namespace_path: Optional[List[str]] = None
|
||||
) -> Job:
|
||||
"""Start dropping a table and return its cleanup job."""
|
||||
raise NotImplementedError
|
||||
|
||||
def rename_table(
|
||||
self,
|
||||
cur_name: str,
|
||||
@@ -656,6 +654,71 @@ class DBConnection(EnforceOverrides):
|
||||
"job_history is not supported for this connection type"
|
||||
)
|
||||
|
||||
@property
|
||||
def functions(self) -> "_SyncFunctions":
|
||||
"""First-class Function operations for this connection."""
|
||||
from ._functions import _SyncFunctions
|
||||
|
||||
return _SyncFunctions(self)
|
||||
|
||||
def _submit_register_function(
|
||||
self, name: str, definition: "_FunctionDefinition"
|
||||
) -> NativeJob:
|
||||
"""Submit a Function registration job via the native connection.
|
||||
|
||||
Connection subclasses that support registration override this hook.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"function registration is not supported for this connection type"
|
||||
)
|
||||
|
||||
def _submit_replace_function(
|
||||
self, name: str, current: "Function", definition: "_FunctionDefinition"
|
||||
) -> NativeJob:
|
||||
"""Submit a Function conditional replace job via the native connection.
|
||||
|
||||
Connection subclasses that support registration override this hook.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"function replace is not supported for this connection type"
|
||||
)
|
||||
|
||||
def _lookup_function_by_name(self, name: str) -> "Function":
|
||||
"""Look up a Function by database-scoped name via the native connection.
|
||||
|
||||
Connection subclasses that share the native Connection override this hook.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"function lookup is not supported for this connection type"
|
||||
)
|
||||
|
||||
def _lookup_function_by_id(self, function_id: str) -> "Function":
|
||||
"""Look up a Function by exact Function ID via the native connection.
|
||||
|
||||
Connection subclasses that share the native Connection override this hook.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"function lookup is not supported for this connection type"
|
||||
)
|
||||
|
||||
def _remove_function_name(self, name: str, current: "Function") -> None:
|
||||
"""Conditionally remove a Function catalog name via the native connection.
|
||||
|
||||
Connection subclasses that share the native Connection override this hook.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"function name removal is not supported for this connection type"
|
||||
)
|
||||
|
||||
def _revoke_function(self, function: "Function") -> None:
|
||||
"""Revoke an exact immutable Function via the native connection.
|
||||
|
||||
Connection subclasses that share the native Connection override this hook.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"function revocation is not supported for this connection type"
|
||||
)
|
||||
|
||||
|
||||
class LanceDBConnection(DBConnection):
|
||||
"""
|
||||
@@ -1192,20 +1255,6 @@ class LanceDBConnection(DBConnection):
|
||||
)
|
||||
)
|
||||
|
||||
@override
|
||||
def drop_table_async(
|
||||
self, name: str, namespace_path: Optional[List[str]] = None
|
||||
) -> Job:
|
||||
"""Start dropping a table and return its cleanup job.
|
||||
|
||||
The table may become unavailable before its data files are removed.
|
||||
Call :meth:`Job.wait` to wait for cleanup to finish.
|
||||
"""
|
||||
if namespace_path is None:
|
||||
namespace_path = []
|
||||
job = LOOP.run(self._conn.drop_table_async(name, namespace_path=namespace_path))
|
||||
return Job(job if isinstance(job, AsyncJob) else AsyncJob(job))
|
||||
|
||||
@override
|
||||
def drop_all_tables(self, namespace_path: Optional[List[str]] = None):
|
||||
if namespace_path is None:
|
||||
@@ -1287,6 +1336,34 @@ class LanceDBConnection(DBConnection):
|
||||
"""
|
||||
return LOOP.run(self._conn.job_history(job_id))
|
||||
|
||||
@override
|
||||
def _submit_register_function(
|
||||
self, name: str, definition: "_FunctionDefinition"
|
||||
) -> NativeJob:
|
||||
return LOOP.run(self._conn._register_function(name, definition))
|
||||
|
||||
@override
|
||||
def _submit_replace_function(
|
||||
self, name: str, current: "Function", definition: "_FunctionDefinition"
|
||||
) -> NativeJob:
|
||||
return LOOP.run(self._conn._replace_function(name, current, definition))
|
||||
|
||||
@override
|
||||
def _lookup_function_by_name(self, name: str) -> "Function":
|
||||
return LOOP.run(self._conn._lookup_function_by_name(name))
|
||||
|
||||
@override
|
||||
def _lookup_function_by_id(self, function_id: str) -> "Function":
|
||||
return LOOP.run(self._conn._lookup_function_by_id(function_id))
|
||||
|
||||
@override
|
||||
def _remove_function_name(self, name: str, current: "Function") -> None:
|
||||
return LOOP.run(self._conn._remove_function_name(name, current))
|
||||
|
||||
@override
|
||||
def _revoke_function(self, function: "Function") -> None:
|
||||
return LOOP.run(self._conn._revoke_function(function))
|
||||
|
||||
@override
|
||||
def namespace_client(self) -> LanceNamespace:
|
||||
"""Get the equivalent namespace client for this connection.
|
||||
@@ -1983,23 +2060,6 @@ class AsyncConnection(object):
|
||||
if f"Table '{name}' was not found" not in str(e):
|
||||
raise e
|
||||
|
||||
async def drop_table_async(
|
||||
self,
|
||||
name: str,
|
||||
*,
|
||||
namespace_path: Optional[List[str]] = None,
|
||||
) -> AsyncJob:
|
||||
"""Start dropping a table and return its cleanup job.
|
||||
|
||||
The table may become unavailable before its data files are removed.
|
||||
Await :meth:`AsyncJob.wait` to wait for cleanup to finish.
|
||||
"""
|
||||
if namespace_path is None:
|
||||
namespace_path = []
|
||||
return AsyncJob(
|
||||
await self._inner.drop_table_async(name, namespace_path=namespace_path)
|
||||
)
|
||||
|
||||
async def drop_all_tables(self, namespace_path: Optional[List[str]] = None):
|
||||
"""Drop all tables from the database.
|
||||
|
||||
@@ -2050,6 +2110,35 @@ class AsyncConnection(object):
|
||||
"""
|
||||
return await self._inner.job_history(job_id)
|
||||
|
||||
@property
|
||||
def functions(self) -> "_AsyncFunctions":
|
||||
"""First-class Function operations for this connection."""
|
||||
from ._functions import _AsyncFunctions
|
||||
|
||||
return _AsyncFunctions(self)
|
||||
|
||||
async def _register_function(
|
||||
self, name: str, definition: "_FunctionDefinition"
|
||||
) -> NativeJob:
|
||||
return await self._inner._register_function(name, definition)
|
||||
|
||||
async def _replace_function(
|
||||
self, name: str, current: "Function", definition: "_FunctionDefinition"
|
||||
) -> NativeJob:
|
||||
return await self._inner._replace_function(name, current, definition)
|
||||
|
||||
async def _lookup_function_by_name(self, name: str) -> "Function":
|
||||
return await self._inner._lookup_function_by_name(name)
|
||||
|
||||
async def _lookup_function_by_id(self, function_id: str) -> "Function":
|
||||
return await self._inner._lookup_function_by_id(function_id)
|
||||
|
||||
async def _remove_function_name(self, name: str, current: "Function") -> None:
|
||||
return await self._inner._remove_function_name(name, current)
|
||||
|
||||
async def _revoke_function(self, function: "Function") -> None:
|
||||
return await self._inner._revoke_function(function)
|
||||
|
||||
async def namespace_client(self) -> LanceNamespace:
|
||||
"""Get the equivalent namespace client for this connection.
|
||||
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
|
||||
"""Custom exception handling"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class MissingValueError(ValueError):
|
||||
"""Exception raised when a required value is missing."""
|
||||
@@ -26,12 +28,47 @@ class MissingColumnError(KeyError):
|
||||
|
||||
|
||||
class JobFailedError(RuntimeError):
|
||||
"""Exception raised when an asynchronous job reaches the failed state."""
|
||||
"""Exception raised when an asynchronous job reaches the failed state.
|
||||
|
||||
pass
|
||||
``error_code`` is the optional exact category string projected from the
|
||||
native job failure when the backend supplied one. The RuntimeError
|
||||
message remains the existing diagnostic text and must not be used to
|
||||
recover or override the code.
|
||||
"""
|
||||
|
||||
__slots__ = ("_error_code",)
|
||||
|
||||
def __init__(self, message: str, error_code: Optional[str] = None) -> None:
|
||||
super().__init__(message)
|
||||
self._error_code = error_code
|
||||
|
||||
@property
|
||||
def error_code(self) -> Optional[str]:
|
||||
"""Exact job failure error category string, when supplied."""
|
||||
return self._error_code
|
||||
|
||||
|
||||
class JobCancelledError(RuntimeError):
|
||||
"""Exception raised when an asynchronous job was cancelled."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class FunctionError(RuntimeError):
|
||||
"""Exception raised when a first-class Function operation fails.
|
||||
|
||||
``code`` is the stable semantic category from the native error. The
|
||||
message is a sanitized client diagnostic and must not be used to recover
|
||||
or override the code.
|
||||
"""
|
||||
|
||||
__slots__ = ("_code",)
|
||||
|
||||
def __init__(self, message: str, code: str) -> None:
|
||||
super().__init__(message)
|
||||
self._code = code
|
||||
|
||||
@property
|
||||
def code(self) -> str:
|
||||
"""Stable Function error category string."""
|
||||
return self._code
|
||||
|
||||
@@ -10,6 +10,7 @@ from typing import Optional
|
||||
from lancedb.background_loop import LOOP
|
||||
|
||||
from . import _lancedb
|
||||
from ._lancedb import Function
|
||||
|
||||
|
||||
class AsyncJob:
|
||||
@@ -44,18 +45,22 @@ class AsyncJob:
|
||||
return "finished"
|
||||
return await self._inner.status()
|
||||
|
||||
async def wait(self, timeout: Optional[timedelta] = None):
|
||||
async def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
|
||||
"""Wait until the operation reaches a terminal state.
|
||||
|
||||
Returns the success result when present (currently a
|
||||
:class:`~lancedb.Function`), or `None` when the job finished without
|
||||
one.
|
||||
|
||||
Raises `JobFailedError` if the operation failed, `JobCancelledError`
|
||||
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
|
||||
"""
|
||||
if self._inner is None:
|
||||
return
|
||||
return None
|
||||
if timeout is None:
|
||||
await self._inner.wait()
|
||||
return await self._inner.wait()
|
||||
else:
|
||||
await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
|
||||
return await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
|
||||
|
||||
async def cancel(self):
|
||||
"""Request cancellation. Cancelling a finished operation is a no-op."""
|
||||
@@ -88,15 +93,19 @@ class Job:
|
||||
return "finished"
|
||||
return LOOP.run(self._inner.status())
|
||||
|
||||
def wait(self, timeout: Optional[timedelta] = None):
|
||||
def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
|
||||
"""Block until the operation reaches a terminal state.
|
||||
|
||||
Returns the success result when present (currently a
|
||||
:class:`~lancedb.Function`), or `None` when the job finished without
|
||||
one.
|
||||
|
||||
Raises `JobFailedError` if the operation failed, `JobCancelledError`
|
||||
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
|
||||
"""
|
||||
if self._inner is None:
|
||||
return
|
||||
LOOP.run(self._inner.wait(timeout))
|
||||
return None
|
||||
return LOOP.run(self._inner.wait(timeout))
|
||||
|
||||
def cancel(self):
|
||||
"""Request cancellation. Cancelling a finished operation is a no-op."""
|
||||
|
||||
@@ -49,7 +49,6 @@ from lancedb._lancedb import (
|
||||
)
|
||||
from lancedb.background_loop import LOOP
|
||||
from lancedb.db import AsyncConnection, DBConnection
|
||||
from lancedb.job import AsyncJob, Job
|
||||
from lance_namespace import (
|
||||
LanceNamespace,
|
||||
connect as namespace_connect,
|
||||
@@ -625,18 +624,6 @@ class LanceNamespaceDBConnection(DBConnection):
|
||||
namespace_path = []
|
||||
LOOP.run(self._inner.drop_table(name, namespace_path=namespace_path))
|
||||
|
||||
@override
|
||||
def drop_table_async(
|
||||
self, name: str, namespace_path: Optional[List[str]] = None
|
||||
) -> Job:
|
||||
"""Start dropping a table and return its cleanup job."""
|
||||
if namespace_path is None:
|
||||
namespace_path = []
|
||||
job = LOOP.run(
|
||||
self._inner.drop_table_async(name, namespace_path=namespace_path)
|
||||
)
|
||||
return Job(job if isinstance(job, AsyncJob) else AsyncJob(job))
|
||||
|
||||
@override
|
||||
def rename_table(
|
||||
self,
|
||||
@@ -1147,14 +1134,6 @@ class AsyncLanceNamespaceDBConnection:
|
||||
namespace_path = []
|
||||
await self._inner.drop_table(name, namespace_path=namespace_path)
|
||||
|
||||
async def drop_table_async(
|
||||
self, name: str, namespace_path: Optional[List[str]] = None
|
||||
) -> AsyncJob:
|
||||
"""Start dropping a table and return its cleanup job."""
|
||||
if namespace_path is None:
|
||||
namespace_path = []
|
||||
return await self._inner.drop_table_async(name, namespace_path=namespace_path)
|
||||
|
||||
async def rename_table(
|
||||
self,
|
||||
cur_name: str,
|
||||
|
||||
@@ -23,10 +23,12 @@ import pyarrow as pa
|
||||
|
||||
from ..common import DATA
|
||||
from ..db import DBConnection, LOOP
|
||||
from ..job import AsyncJob, Job
|
||||
from ..job import Job
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .._lancedb import JobDescription, JobInfo
|
||||
from .._lancedb import Function
|
||||
from .._lancedb import Job as NativeJob
|
||||
from .._lancedb import JobDescription, JobInfo, _FunctionDefinition
|
||||
from ..embeddings import EmbeddingFunctionConfig
|
||||
from lance_namespace import (
|
||||
LanceNamespace,
|
||||
@@ -663,16 +665,6 @@ class RemoteDBConnection(DBConnection):
|
||||
namespace_path = []
|
||||
LOOP.run(self._conn.drop_table(name, namespace_path=namespace_path))
|
||||
|
||||
@override
|
||||
def drop_table_async(
|
||||
self, name: str, namespace_path: Optional[List[str]] = None
|
||||
) -> Job:
|
||||
"""Start dropping a table and return its cleanup job."""
|
||||
if namespace_path is None:
|
||||
namespace_path = []
|
||||
job = LOOP.run(self._conn.drop_table_async(name, namespace_path=namespace_path))
|
||||
return Job(job if isinstance(job, AsyncJob) else AsyncJob(job))
|
||||
|
||||
@override
|
||||
def rename_table(
|
||||
self,
|
||||
@@ -744,6 +736,34 @@ class RemoteDBConnection(DBConnection):
|
||||
"""
|
||||
return LOOP.run(self._conn.job_history(job_id))
|
||||
|
||||
@override
|
||||
def _submit_register_function(
|
||||
self, name: str, definition: "_FunctionDefinition"
|
||||
) -> "NativeJob":
|
||||
return LOOP.run(self._conn._register_function(name, definition))
|
||||
|
||||
@override
|
||||
def _submit_replace_function(
|
||||
self, name: str, current: "Function", definition: "_FunctionDefinition"
|
||||
) -> "NativeJob":
|
||||
return LOOP.run(self._conn._replace_function(name, current, definition))
|
||||
|
||||
@override
|
||||
def _lookup_function_by_name(self, name: str) -> "Function":
|
||||
return LOOP.run(self._conn._lookup_function_by_name(name))
|
||||
|
||||
@override
|
||||
def _lookup_function_by_id(self, function_id: str) -> "Function":
|
||||
return LOOP.run(self._conn._lookup_function_by_id(function_id))
|
||||
|
||||
@override
|
||||
def _remove_function_name(self, name: str, current: "Function") -> None:
|
||||
return LOOP.run(self._conn._remove_function_name(name, current))
|
||||
|
||||
@override
|
||||
def _revoke_function(self, function: "Function") -> None:
|
||||
return LOOP.run(self._conn._revoke_function(function))
|
||||
|
||||
@override
|
||||
def namespace_client(self) -> LanceNamespace:
|
||||
"""Get the equivalent namespace client for this connection.
|
||||
|
||||
@@ -7,6 +7,7 @@ import logging
|
||||
from functools import cached_property
|
||||
import os
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
@@ -50,7 +51,7 @@ from lancedb.index import (
|
||||
)
|
||||
from lancedb.job import Job
|
||||
from lancedb.remote.db import LOOP
|
||||
from lancedb.table import IndexConfigType, KNOWN_METRICS, TableBase
|
||||
from lancedb.table import IndexConfigType, KNOWN_METRICS
|
||||
import pyarrow as pa
|
||||
|
||||
from lancedb.common import DATA, VEC, VECTOR_COLUMN_NAME
|
||||
@@ -67,6 +68,9 @@ from ..query import (
|
||||
from ..table import AsyncTable, BlobMode, Branches, IndexStatistics, Query, Table, Tags
|
||||
from ..types import BaseTokenizerType
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from lancedb._lancedb import _FunctionCall
|
||||
|
||||
|
||||
class RemoteTable(Table):
|
||||
def __init__(
|
||||
@@ -570,6 +574,45 @@ class RemoteTable(Table):
|
||||
)
|
||||
)
|
||||
|
||||
def add_generated_column(self, column_name: str, call: "_FunctionCall") -> Job:
|
||||
"""Add a generated column from an authored Function call.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
|
||||
of the Job does not publish the column; callers must wait and re-read
|
||||
the table to observe the new definition and values.
|
||||
"""
|
||||
return Job(LOOP.run(self._table.add_generated_column(column_name, call)))
|
||||
|
||||
def generated_column_status(
|
||||
self, column_name: str
|
||||
) -> Literal["complete", "incomplete"]:
|
||||
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
|
||||
|
||||
Projection-only: reads the named column's stored definition status.
|
||||
Does not refresh values, submit a Job, or mutate table state.
|
||||
"""
|
||||
return LOOP.run(self._table.generated_column_status(column_name))
|
||||
|
||||
def refresh_generated_column(self, column_name: str) -> Job:
|
||||
"""Refresh values for an existing generated column.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
|
||||
of the Job does not publish new values; callers must wait and re-read
|
||||
the table to observe refreshed results.
|
||||
"""
|
||||
return Job(LOOP.run(self._table.refresh_generated_column(column_name)))
|
||||
|
||||
def alter_generated_column(
|
||||
self, column_name: str, new_call: "_FunctionCall"
|
||||
) -> Job:
|
||||
"""Alter the Function call for an existing generated column.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
|
||||
of the Job does not publish the new definition; callers must wait and
|
||||
re-read the table to observe the updated column.
|
||||
"""
|
||||
return Job(LOOP.run(self._table.alter_generated_column(column_name, new_call)))
|
||||
|
||||
def _is_legacy_create_index_call(
|
||||
self,
|
||||
first_arg: str,
|
||||
@@ -958,19 +1001,8 @@ class RemoteTable(Table):
|
||||
def count_rows(self, filter: Optional[str] = None) -> int:
|
||||
return LOOP.run(self._table.count_rows(filter))
|
||||
|
||||
def add_columns(
|
||||
self,
|
||||
transforms: Dict[str, str] | None = None,
|
||||
*,
|
||||
computed: Dict[str, str] | None = None,
|
||||
) -> AddColumnsResult:
|
||||
return LOOP.run(self._table.add_columns(transforms, computed=computed))
|
||||
|
||||
def refresh_column(self, column: str):
|
||||
return LOOP.run(self._table.refresh_column(column))
|
||||
|
||||
def refresh_column_async(self, column: str) -> Job:
|
||||
return Job(LOOP.run(self._table.refresh_column_async(column)))
|
||||
def add_columns(self, transforms: Dict[str, str]) -> AddColumnsResult:
|
||||
return LOOP.run(self._table.add_columns(transforms))
|
||||
|
||||
def alter_columns(
|
||||
self, *alterations: Iterable[Dict[str, str]]
|
||||
@@ -1082,13 +1114,6 @@ class RemoteTable(Table):
|
||||
def blob_columns(self) -> list[str]:
|
||||
return LOOP.run(self._table.blob_columns())
|
||||
|
||||
def add_bases(
|
||||
self,
|
||||
bases: Union[str, TableBase, Iterable[Union[str, TableBase]]],
|
||||
) -> None:
|
||||
"""Register additional storage bases for this table."""
|
||||
LOOP.run(self._table.add_bases(bases))
|
||||
|
||||
def fetch_blobs(
|
||||
self, column: str, row_ids: Union[list[int], pa.Table]
|
||||
) -> pa.LargeBinaryArray:
|
||||
|
||||
+123
-268
@@ -19,7 +19,6 @@ from typing import (
|
||||
Iterable,
|
||||
List,
|
||||
Literal,
|
||||
Mapping,
|
||||
Optional,
|
||||
Sequence,
|
||||
Tuple,
|
||||
@@ -177,7 +176,6 @@ if TYPE_CHECKING:
|
||||
CompactionStats,
|
||||
Tag,
|
||||
AddColumnsResult,
|
||||
RefreshColumnResult,
|
||||
AddResult,
|
||||
AlterColumnsResult,
|
||||
UpdateFieldMetadataResult,
|
||||
@@ -187,6 +185,7 @@ if TYPE_CHECKING:
|
||||
LsmWriteSpec,
|
||||
MergeResult,
|
||||
UpdateResult,
|
||||
_FunctionCall,
|
||||
)
|
||||
from .index import IndexConfig
|
||||
import pandas
|
||||
@@ -711,21 +710,6 @@ def _normalize_progress(progress):
|
||||
return progress, False
|
||||
|
||||
|
||||
@dataclass
|
||||
class TableBase:
|
||||
"""An extra storage prefix registered on a table.
|
||||
|
||||
``path`` is an object-store URI. ``name`` is an optional alias.
|
||||
``is_dataset_root`` is true when ``path`` points to a Lance dataset
|
||||
root. When false, ``path`` points directly to the directory containing
|
||||
the referenced files.
|
||||
"""
|
||||
|
||||
path: str
|
||||
name: Optional[str] = None
|
||||
is_dataset_root: bool = False
|
||||
|
||||
|
||||
class Table(ABC):
|
||||
"""
|
||||
A Table is a collection of Records in a LanceDB Database.
|
||||
@@ -1024,6 +1008,43 @@ class Table(ABC):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def add_generated_column(self, column_name: str, call: _FunctionCall) -> Job:
|
||||
"""Add a generated column from an authored Function call.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
|
||||
of the Job does not publish the column; callers must wait and re-read
|
||||
the table to observe the new definition and values.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def generated_column_status(
|
||||
self, column_name: str
|
||||
) -> Literal["complete", "incomplete"]:
|
||||
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
|
||||
|
||||
Projection-only: reads the named column's stored definition status.
|
||||
Does not refresh values, submit a Job, or mutate table state.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def refresh_generated_column(self, column_name: str) -> Job:
|
||||
"""Refresh values for an existing generated column.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
|
||||
of the Job does not publish new values; callers must wait and re-read
|
||||
the table to observe refreshed results.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def alter_generated_column(self, column_name: str, new_call: _FunctionCall) -> Job:
|
||||
"""Alter the Function call for an existing generated column.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
|
||||
of the Job does not publish the new definition; callers must wait and
|
||||
re-read the table to observe the updated column.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def drop_index(self, name: str) -> None:
|
||||
"""
|
||||
Drop an index from the table.
|
||||
@@ -1584,18 +1605,6 @@ class Table(ABC):
|
||||
def blob_columns(self) -> list[str]:
|
||||
"""Names of the blob v2 columns declared on this table."""
|
||||
|
||||
def add_bases(
|
||||
self,
|
||||
bases: Union[str, TableBase, Iterable[Union[str, TableBase]]],
|
||||
) -> None:
|
||||
"""Register additional storage bases for this table.
|
||||
|
||||
A URI string is a non-root base with no alias::
|
||||
|
||||
table.add_bases("s3://bucket/media/")
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def fetch_blobs(
|
||||
self, column: str, row_ids: Union[list[int], pa.Table]
|
||||
@@ -1945,14 +1954,7 @@ class Table(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def add_columns(
|
||||
self,
|
||||
transforms: Dict[str, str]
|
||||
| pa.Field
|
||||
| List[pa.Field]
|
||||
| pa.Schema
|
||||
| None = None,
|
||||
*,
|
||||
computed: Dict[str, str] | None = None,
|
||||
self, transforms: Dict[str, str] | pa.Field | List[pa.Field] | pa.Schema
|
||||
):
|
||||
"""
|
||||
Add new columns with defined values.
|
||||
@@ -1966,95 +1968,11 @@ class Table(ABC):
|
||||
Alternatively, a pyarrow Field or Schema can be provided to add
|
||||
new columns with the specified data types. The new columns will
|
||||
be initialized with null values.
|
||||
computed: Dict[str, str], optional
|
||||
A map of column name to a SQL expression defining the column. The
|
||||
column's type and inputs are derived from the expression, so no
|
||||
data type is supplied.
|
||||
|
||||
Unlike ``transforms``, the expression is stored rather than
|
||||
evaluated now: the column is committed with no values, and rows get
|
||||
them from [`refresh_column`][lancedb.table.Table.refresh_column].
|
||||
Declaring one therefore costs the same on a large table as on an
|
||||
empty one.
|
||||
|
||||
A refresh does not revisit rows it has already filled, so mutating
|
||||
an input leaves the value computed at fill time; recomputing means
|
||||
dropping the column and declaring it again. While a declaration
|
||||
reads a column, that column cannot be renamed, retyped or dropped.
|
||||
|
||||
On LanceDB Cloud and Enterprise the expression is planned by the
|
||||
server, and the refresh runs as a server job -- see
|
||||
[`refresh_column_async`][lancedb.table.Table.refresh_column_async].
|
||||
Cannot be combined with ``transforms``.
|
||||
|
||||
Returns
|
||||
-------
|
||||
AddColumnsResult
|
||||
version: the new version number of the table after adding columns.
|
||||
|
||||
Examples
|
||||
--------
|
||||
>>> import lancedb
|
||||
>>> db = lancedb.connect("./.lancedb")
|
||||
>>> table = db.create_table("computed_demo", [{"x": 1}, {"x": 2}])
|
||||
>>> table.add_columns(computed={"doubled": "x * 2"})
|
||||
AddColumnsResult(version=2)
|
||||
>>> table.refresh_column("doubled")
|
||||
RefreshColumnResult(rows_filled=2, version=3)
|
||||
>>> table.to_arrow().sort_by("x").to_pandas()
|
||||
x doubled
|
||||
0 1 2
|
||||
1 2 4
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def refresh_column(self, column: str) -> "RefreshColumnResult":
|
||||
"""
|
||||
Fill the rows of a computed column that hold no value yet.
|
||||
|
||||
Declared with ``add_columns(computed=...)``, a column starts empty and
|
||||
gets its values here. Rows appended since the last refresh are filled
|
||||
by the next one; rows already filled are left as they are, so the call
|
||||
is idempotent and does not observe a mutated input.
|
||||
|
||||
Local tables only: a remote refresh runs as a server job, through
|
||||
[`refresh_column_async`][lancedb.table.Table.refresh_column_async].
|
||||
|
||||
Parameters
|
||||
----------
|
||||
column: str
|
||||
The name of the computed column to fill.
|
||||
|
||||
Returns
|
||||
-------
|
||||
RefreshColumnResult
|
||||
rows_filled: the number of rows given a value.
|
||||
version: the new version number of the table.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def refresh_column_async(self, column: str) -> Job:
|
||||
"""
|
||||
Like :meth:`refresh_column`, but returns a handle to the refresh job
|
||||
instead of blocking until it completes.
|
||||
|
||||
The job may already be complete when returned; callers must not assume
|
||||
the column is filled until :meth:`Job.wait` returns. Invalid input --
|
||||
an unknown column, or one that is not computed -- raises here rather
|
||||
than failing the job. On local tables the job runs in-process; on
|
||||
LanceDB Cloud and Enterprise it is the server's backfill job.
|
||||
|
||||
Examples
|
||||
--------
|
||||
>>> import lancedb
|
||||
>>> db = lancedb.connect("./.lancedb")
|
||||
>>> table = db.create_table("computed_job_demo", [{"x": 1}, {"x": 2}])
|
||||
>>> table.add_columns(computed={"doubled": "x * 2"})
|
||||
AddColumnsResult(version=2)
|
||||
>>> job = table.refresh_column_async("doubled")
|
||||
>>> job.wait()
|
||||
>>> job.status()
|
||||
'finished'
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
@@ -2442,12 +2360,6 @@ class LanceTable(Table):
|
||||
def blob_columns(self) -> list[str]:
|
||||
return LOOP.run(self._table.blob_columns())
|
||||
|
||||
def add_bases(
|
||||
self,
|
||||
bases: Union[str, TableBase, Iterable[Union[str, TableBase]]],
|
||||
) -> None:
|
||||
LOOP.run(self._table.add_bases(bases))
|
||||
|
||||
def fetch_blobs(
|
||||
self, column: str, row_ids: Union[list[int], pa.Table]
|
||||
) -> pa.LargeBinaryArray:
|
||||
@@ -2975,6 +2887,43 @@ class LanceTable(Table):
|
||||
)
|
||||
)
|
||||
|
||||
def add_generated_column(self, column_name: str, call: _FunctionCall) -> Job:
|
||||
"""Add a generated column from an authored Function call.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
|
||||
of the Job does not publish the column; callers must wait and re-read
|
||||
the table to observe the new definition and values.
|
||||
"""
|
||||
return Job(LOOP.run(self._table.add_generated_column(column_name, call)))
|
||||
|
||||
def generated_column_status(
|
||||
self, column_name: str
|
||||
) -> Literal["complete", "incomplete"]:
|
||||
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
|
||||
|
||||
Projection-only: reads the named column's stored definition status.
|
||||
Does not refresh values, submit a Job, or mutate table state.
|
||||
"""
|
||||
return LOOP.run(self._table.generated_column_status(column_name))
|
||||
|
||||
def refresh_generated_column(self, column_name: str) -> Job:
|
||||
"""Refresh values for an existing generated column.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
|
||||
of the Job does not publish new values; callers must wait and re-read
|
||||
the table to observe refreshed results.
|
||||
"""
|
||||
return Job(LOOP.run(self._table.refresh_generated_column(column_name)))
|
||||
|
||||
def alter_generated_column(self, column_name: str, new_call: _FunctionCall) -> Job:
|
||||
"""Alter the Function call for an existing generated column.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
|
||||
of the Job does not publish the new definition; callers must wait and
|
||||
re-read the table to observe the updated column.
|
||||
"""
|
||||
return Job(LOOP.run(self._table.alter_generated_column(column_name, new_call)))
|
||||
|
||||
def _is_legacy_create_index_call(
|
||||
self,
|
||||
first_arg: str,
|
||||
@@ -4065,28 +4014,9 @@ class LanceTable(Table):
|
||||
return LOOP.run(self._table.index_stats(index_name))
|
||||
|
||||
def add_columns(
|
||||
self,
|
||||
transforms: Dict[str, str]
|
||||
| pa.field
|
||||
| List[pa.field]
|
||||
| pa.Schema
|
||||
| None = None,
|
||||
*,
|
||||
computed: Dict[str, str] | None = None,
|
||||
self, transforms: Dict[str, str] | pa.field | List[pa.field] | pa.Schema
|
||||
) -> AddColumnsResult:
|
||||
return LOOP.run(self._table.add_columns(transforms, computed=computed))
|
||||
|
||||
def refresh_column(self, column: str) -> "RefreshColumnResult":
|
||||
"""Fill a computed column's unfilled rows. See
|
||||
[`AsyncTable.refresh_column`][lancedb.AsyncTable.refresh_column]."""
|
||||
return LOOP.run(self._table.refresh_column(column))
|
||||
|
||||
def refresh_column_async(self, column: str) -> Job:
|
||||
"""Fill a computed column's unfilled rows, returning a handle to the
|
||||
refresh job. See
|
||||
[`Table.refresh_column_async`][lancedb.table.Table.refresh_column_async].
|
||||
"""
|
||||
return Job(LOOP.run(self._table.refresh_column_async(column)))
|
||||
return LOOP.run(self._table.add_columns(transforms))
|
||||
|
||||
def alter_columns(
|
||||
self, *alterations: Iterable[Dict[str, str]]
|
||||
@@ -5211,6 +5141,50 @@ class AsyncTable:
|
||||
)
|
||||
return AsyncJob(job)
|
||||
|
||||
async def add_generated_column(
|
||||
self, column_name: str, call: _FunctionCall
|
||||
) -> AsyncJob:
|
||||
"""Add a generated column from an authored Function call.
|
||||
|
||||
Returns an :class:`~lancedb.job.AsyncJob` for the create operation.
|
||||
Acceptance of the Job does not publish the column; callers must wait
|
||||
and re-read the table to observe the new definition and values.
|
||||
"""
|
||||
job = await self._inner._add_generated_column(column_name, call)
|
||||
return AsyncJob(job)
|
||||
|
||||
async def generated_column_status(
|
||||
self, column_name: str
|
||||
) -> Literal["complete", "incomplete"]:
|
||||
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
|
||||
|
||||
Projection-only: reads the named column's stored definition status.
|
||||
Does not refresh values, submit a Job, or mutate table state.
|
||||
"""
|
||||
return await self._inner._generated_column_status(column_name)
|
||||
|
||||
async def refresh_generated_column(self, column_name: str) -> AsyncJob:
|
||||
"""Refresh values for an existing generated column.
|
||||
|
||||
Returns an :class:`~lancedb.job.AsyncJob` for the refresh operation.
|
||||
Acceptance of the Job does not publish new values; callers must wait
|
||||
and re-read the table to observe refreshed results.
|
||||
"""
|
||||
job = await self._inner._refresh_generated_column(column_name)
|
||||
return AsyncJob(job)
|
||||
|
||||
async def alter_generated_column(
|
||||
self, column_name: str, new_call: _FunctionCall
|
||||
) -> AsyncJob:
|
||||
"""Alter the Function call for an existing generated column.
|
||||
|
||||
Returns an :class:`~lancedb.job.AsyncJob` for the change operation.
|
||||
Acceptance of the Job does not publish the new definition; callers must
|
||||
wait and re-read the table to observe the updated column.
|
||||
"""
|
||||
job = await self._inner._alter_generated_column(column_name, new_call)
|
||||
return AsyncJob(job)
|
||||
|
||||
async def drop_index(self, name: str) -> None:
|
||||
"""
|
||||
Drop an index from the table.
|
||||
@@ -6001,14 +5975,7 @@ class AsyncTable:
|
||||
return await self._inner.update(updates_sql, where)
|
||||
|
||||
async def add_columns(
|
||||
self,
|
||||
transforms: dict[str, str]
|
||||
| pa.field
|
||||
| List[pa.field]
|
||||
| pa.Schema
|
||||
| None = None,
|
||||
*,
|
||||
computed: dict[str, str] | None = None,
|
||||
self, transforms: dict[str, str] | pa.field | List[pa.field] | pa.Schema
|
||||
) -> AddColumnsResult:
|
||||
"""
|
||||
Add new columns with defined values.
|
||||
@@ -6021,22 +5988,6 @@ class AsyncTable:
|
||||
each row in the table, and can reference existing columns.
|
||||
Alternatively, you can pass a pyarrow field or schema to add
|
||||
new columns with NULLs.
|
||||
computed: Dict[str, str], optional
|
||||
A map of column name to a SQL expression defining the column. The
|
||||
column's type and inputs are derived from the expression.
|
||||
|
||||
Unlike ``transforms``, the expression is stored rather than
|
||||
evaluated now: the column is committed with no values, and rows get
|
||||
them from
|
||||
[`refresh_column`][lancedb.table.AsyncTable.refresh_column].
|
||||
|
||||
A refresh does not revisit rows it has already filled, so mutating
|
||||
an input leaves the value computed at fill time. While a
|
||||
declaration reads a column, that column cannot be renamed, retyped
|
||||
or dropped.
|
||||
|
||||
On LanceDB Cloud and Enterprise the expression is planned by
|
||||
the server. Cannot be combined with ``transforms``.
|
||||
|
||||
Returns
|
||||
-------
|
||||
@@ -6050,71 +6001,11 @@ class AsyncTable:
|
||||
{isinstance(f, pa.Field) for f in transforms}
|
||||
):
|
||||
transforms = pa.schema(transforms)
|
||||
if computed:
|
||||
if transforms:
|
||||
raise ValueError(
|
||||
"add_columns cannot take both transforms and computed columns"
|
||||
)
|
||||
return await self._inner.add_computed_columns(list(computed.items()))
|
||||
if transforms is None:
|
||||
raise ValueError("add_columns requires transforms or computed columns")
|
||||
if isinstance(transforms, pa.Schema):
|
||||
return await self._inner.add_columns_with_schema(transforms)
|
||||
else:
|
||||
return await self._inner.add_columns(list(transforms.items()))
|
||||
|
||||
async def refresh_column(self, column: str) -> RefreshColumnResult:
|
||||
"""
|
||||
Fill the rows of a computed column that hold no value yet.
|
||||
|
||||
Declared with ``add_columns(computed=...)``, a column starts empty and
|
||||
gets its values here. Rows appended since the last refresh are filled
|
||||
by the next one; rows already filled are left as they are, so the call
|
||||
is idempotent and does not observe a mutated input.
|
||||
|
||||
Local tables only: a remote refresh runs as a server job, through
|
||||
[`refresh_column_async`][lancedb.table.Table.refresh_column_async].
|
||||
|
||||
Parameters
|
||||
----------
|
||||
column: str
|
||||
The name of the computed column to fill.
|
||||
|
||||
Returns
|
||||
-------
|
||||
RefreshColumnResult
|
||||
The number of rows filled and the new version of the table.
|
||||
"""
|
||||
return await self._inner.refresh_column(column)
|
||||
|
||||
async def refresh_column_async(self, column: str) -> AsyncJob:
|
||||
"""
|
||||
Like :meth:`refresh_column`, but returns a handle to the refresh job
|
||||
instead of blocking until it completes.
|
||||
|
||||
The job may already be complete when returned; callers must not assume
|
||||
the column is filled until :meth:`AsyncJob.wait` resolves. Invalid
|
||||
input -- an unknown column, or one that is not computed -- raises here
|
||||
rather than failing the job. On local tables the job runs
|
||||
in-process; on LanceDB Cloud and Enterprise it is the server's
|
||||
backfill job.
|
||||
|
||||
Examples
|
||||
--------
|
||||
>>> import asyncio
|
||||
>>> import lancedb
|
||||
>>> async def refresh_in_background():
|
||||
... db = await lancedb.connect_async("./.lancedb")
|
||||
... table = await db.create_table("computed_job_async_demo", [{"x": 1}])
|
||||
... await table.add_columns(computed={"doubled": "x * 2"})
|
||||
... job = await table.refresh_column_async("doubled")
|
||||
... await job.wait()
|
||||
... return await job.status()
|
||||
>>> asyncio.run(refresh_in_background())
|
||||
'finished'
|
||||
"""
|
||||
return AsyncJob(await self._inner.refresh_column_async(column))
|
||||
|
||||
async def alter_columns(
|
||||
self, *alterations: Iterable[dict[str, Any]]
|
||||
) -> AlterColumnsResult:
|
||||
@@ -6300,18 +6191,6 @@ class AsyncTable:
|
||||
async def blob_columns(self) -> list[str]:
|
||||
return await self._inner.blob_columns()
|
||||
|
||||
async def add_bases(
|
||||
self,
|
||||
bases: Union[str, TableBase, Iterable[Union[str, TableBase]]],
|
||||
) -> None:
|
||||
"""Register additional storage bases for this table.
|
||||
|
||||
A URI string is a non-root base with no alias::
|
||||
|
||||
await table.add_bases("s3://bucket/media/")
|
||||
"""
|
||||
await self._inner.add_bases(_normalize_bases(bases))
|
||||
|
||||
async def fetch_blobs(
|
||||
self, column: str, row_ids: Union[list[int], pa.Table]
|
||||
) -> pa.LargeBinaryArray:
|
||||
@@ -6530,30 +6409,6 @@ class AsyncTable:
|
||||
await self._inner.replace_field_metadata(field_name, new_metadata)
|
||||
|
||||
|
||||
def _normalize_bases(
|
||||
base_inputs: Union[str, TableBase, Iterable[Union[str, TableBase]]],
|
||||
) -> list[TableBase]:
|
||||
if isinstance(base_inputs, (str, TableBase)):
|
||||
items: Iterable[Union[str, TableBase]] = [base_inputs]
|
||||
elif isinstance(base_inputs, Mapping):
|
||||
raise TypeError(
|
||||
"Expected a URI string, TableBase, or an iterable of those values"
|
||||
)
|
||||
else:
|
||||
items = base_inputs
|
||||
normalized_bases: list[TableBase] = []
|
||||
for base in items:
|
||||
if isinstance(base, str):
|
||||
normalized_bases.append(TableBase(path=base))
|
||||
elif isinstance(base, TableBase):
|
||||
normalized_bases.append(base)
|
||||
else:
|
||||
raise TypeError(
|
||||
f"Expected a URI string or TableBase, got {type(base).__name__}"
|
||||
)
|
||||
return normalized_bases
|
||||
|
||||
|
||||
@dataclass
|
||||
class IndexStatistics:
|
||||
"""
|
||||
|
||||
@@ -1,72 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
|
||||
|
||||
def test_add_bases_accepts_named_and_dataset_root(tmp_path):
|
||||
media = tmp_path / "media"
|
||||
parent = tmp_path / "parent"
|
||||
media.mkdir()
|
||||
parent.mkdir()
|
||||
db = lancedb.connect(tmp_path / "db")
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = db.create_table("photos", schema=schema)
|
||||
table.add_bases(
|
||||
[
|
||||
lancedb.TableBase(path=media.as_uri(), name="media", is_dataset_root=False),
|
||||
lancedb.TableBase(
|
||||
path=parent.as_uri(), name="parent", is_dataset_root=True
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def test_add_bases_accepts_two_unnamed_paths(tmp_path):
|
||||
media = tmp_path / "media"
|
||||
other = tmp_path / "other"
|
||||
media.mkdir()
|
||||
other.mkdir()
|
||||
db = lancedb.connect(tmp_path / "db")
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = db.create_table("photos", schema=schema)
|
||||
table.add_bases([media.as_uri(), other.as_uri()])
|
||||
|
||||
|
||||
def test_add_bases_rejects_dict_input(tmp_path):
|
||||
db = lancedb.connect(tmp_path / "db")
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = db.create_table("photos", schema=schema)
|
||||
with pytest.raises(TypeError, match="TableBase"):
|
||||
table.add_bases({"path": "s3://bucket/media/"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_add_bases_accepts_file_uri(tmp_path):
|
||||
media = tmp_path / "media"
|
||||
media.mkdir()
|
||||
db = await lancedb.connect_async(tmp_path / "db")
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = await db.create_table("photos", schema=schema)
|
||||
await table.add_bases(media.as_uri())
|
||||
|
||||
|
||||
def test_memory_add_bases_accepts_file_uri(tmp_path):
|
||||
media = tmp_path / "media"
|
||||
media.mkdir()
|
||||
db = lancedb.connect("memory:///")
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = db.create_table("photos", schema=schema)
|
||||
table.add_bases(media.as_uri())
|
||||
|
||||
|
||||
def test_namespace_add_bases_accepts_file_uri(tmp_path):
|
||||
media = tmp_path / "media"
|
||||
media.mkdir()
|
||||
db = lancedb.connect_namespace("dir", {"root": str(tmp_path / "ns")})
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = db.create_table("photos", schema=schema)
|
||||
table.add_bases(media.as_uri())
|
||||
@@ -755,7 +755,8 @@ def test_delete_table(tmp_db: lancedb.DBConnection):
|
||||
assert tmp_db.table_names() == []
|
||||
|
||||
|
||||
def test_drop_table_async(tmp_db: lancedb.DBConnection):
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_table_async(tmp_db: lancedb.DBConnection):
|
||||
data = pd.DataFrame(
|
||||
{
|
||||
"vector": [[3.1, 4.1], [5.9, 26.5]],
|
||||
@@ -771,10 +772,7 @@ def test_drop_table_async(tmp_db: lancedb.DBConnection):
|
||||
|
||||
assert tmp_db.table_names() == ["test"]
|
||||
|
||||
job = tmp_db.drop_table_async("test")
|
||||
assert job.id is None
|
||||
assert job.status() == "finished"
|
||||
job.wait()
|
||||
tmp_db.drop_table("test")
|
||||
assert tmp_db.table_names() == []
|
||||
|
||||
tmp_db.create_table("test", data=data)
|
||||
@@ -783,17 +781,6 @@ def test_drop_table_async(tmp_db: lancedb.DBConnection):
|
||||
tmp_db.drop_table("does_not_exist", ignore_missing=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_drop_table_async_connection(tmp_db_async: lancedb.AsyncConnection):
|
||||
await tmp_db_async.create_table("test", data=pa.table({"id": [1, 2]}))
|
||||
|
||||
job = await tmp_db_async.drop_table_async("test")
|
||||
assert job.id is None
|
||||
assert await job.status() == "finished"
|
||||
await job.wait()
|
||||
assert await tmp_db_async.table_names() == []
|
||||
|
||||
|
||||
def test_drop_database(tmp_db: lancedb.DBConnection):
|
||||
data = pd.DataFrame(
|
||||
{
|
||||
|
||||
@@ -0,0 +1,372 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python exact Function handle call authoring (FF-028)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.expr import Expr, col, func, lit
|
||||
|
||||
_CALL_PATH = "/v1/functions/lookup"
|
||||
_CALL_CATALOG_NAME = "text.normalize.call-name"
|
||||
_CALL_FUNCTION_ID = "fn.exact.call-handle"
|
||||
_LITERAL_PAYLOAD_SENTINEL = "LITERAL_PAYLOAD_SENTINEL_call_xyz_42"
|
||||
_INT_PAYLOAD_SENTINEL = 2_147_000_123
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64).
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
_LIST_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////+4AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAABAAAANz///8c"
|
||||
"AAAADAAAAAAAAQxcAAAAAQAAABwAAAAEAAQABAAAABAAFAAQAA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECH"
|
||||
"AAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAQAAABpdGVtAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////8AAAAAFAAAAAAAAAAMABQAEgAMAAgABAAMAAAAnAAAAKAAAAAQAAAAAAAEAAgACAAAAAQACAAAAAQAAAA"
|
||||
"BAAAABAAAANz///8cAAAADAAAAAAAAQxcAAAAAQAAABwAAAAEAAQABAAAABAAFAAQAA4ADwAEAAAACAAQAAAA"
|
||||
"GAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAQAAABpdGVtAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAAAwAAAAEFSUk9XMQ=="
|
||||
)
|
||||
|
||||
_OVERDESIGN_ATTRS = (
|
||||
"id",
|
||||
"function_id",
|
||||
"name",
|
||||
"connection",
|
||||
"table",
|
||||
"snapshot",
|
||||
"field_id",
|
||||
"field_ids",
|
||||
"job",
|
||||
"job_id",
|
||||
"artifact",
|
||||
"digest",
|
||||
"retry_key",
|
||||
"idempotency_key",
|
||||
"user_version",
|
||||
"execute",
|
||||
"status",
|
||||
"wait",
|
||||
"cancel",
|
||||
"to_json",
|
||||
"_to_json",
|
||||
"serialize",
|
||||
"geneva",
|
||||
)
|
||||
|
||||
|
||||
def _sample_function_wire(
|
||||
*,
|
||||
function_id: str = _CALL_FUNCTION_ID,
|
||||
parameters: list[dict[str, str]] | None = None,
|
||||
output_type_ipc: str = _UTF8_TYPE_IPC_B64,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": function_id,
|
||||
"signature": {
|
||||
"parameters": parameters
|
||||
or [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": output_type_ipc,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _lookup_success_body(function: dict[str, Any] | None = None) -> bytes:
|
||||
return json.dumps({"function": function or _sample_function_wire()}).encode("utf-8")
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _lookup_function(function: dict[str, Any] | None = None):
|
||||
body = _lookup_success_body(function)
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _CALL_PATH
|
||||
_read_body(request)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(body)
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
return db.functions.get(_CALL_CATALOG_NAME)
|
||||
|
||||
|
||||
def _authored_call_type():
|
||||
cls = getattr(_native, "_FunctionCall", None)
|
||||
if cls is None:
|
||||
pytest.fail("lancedb._lancedb._FunctionCall is missing")
|
||||
return cls
|
||||
|
||||
|
||||
def _exception_text(exc: BaseException) -> str:
|
||||
parts = [str(exc), repr(exc)]
|
||||
current: BaseException | None = exc
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(f"{type(current).__name__}: {current}")
|
||||
current = current.__cause__ or current.__context__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def test_function_keyword_call_returns_private_frozen_authored_value():
|
||||
function = _lookup_function()
|
||||
assert callable(function)
|
||||
|
||||
authored = function(text=col("text"), limit=8)
|
||||
authored_type = _authored_call_type()
|
||||
assert type(authored) is authored_type
|
||||
assert authored_type.__module__ == "lancedb._lancedb"
|
||||
assert authored_type.__name__ == "_FunctionCall"
|
||||
|
||||
# Keyword order must not matter; bindings store/render in signature order.
|
||||
authored_reversed = function(limit=8, text=col("text"))
|
||||
assert type(authored_reversed) is authored_type
|
||||
rendered = repr(authored_reversed)
|
||||
assert rendered.index("text=") < rendered.index("limit=")
|
||||
assert 'text=field("text")' in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
|
||||
|
||||
def test_function_call_rejects_positional_missing_and_unknown_args():
|
||||
function = _lookup_function()
|
||||
|
||||
with pytest.raises(TypeError, match="keyword"):
|
||||
function(col("text"), 8)
|
||||
|
||||
with pytest.raises((TypeError, ValueError), match="limit"):
|
||||
function(text=col("text"))
|
||||
|
||||
with pytest.raises((TypeError, ValueError), match="text"):
|
||||
function(limit=8)
|
||||
|
||||
with pytest.raises((TypeError, ValueError), match="unknown|extra"):
|
||||
function(text=col("text"), limit=8, extra=1)
|
||||
|
||||
|
||||
def test_function_call_accepts_direct_case_sensitive_column_and_rejects_complex_exprs():
|
||||
function = _lookup_function()
|
||||
|
||||
authored = function(text=col("firstName"), limit=1)
|
||||
assert type(authored) is _authored_call_type()
|
||||
rendered = repr(authored)
|
||||
assert 'text=field("firstName")' in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
|
||||
complex_exprs = (
|
||||
col("text") + lit("x"),
|
||||
col("text").cast(pa.string()),
|
||||
func("lower", col("text")),
|
||||
col("text") == lit("x"),
|
||||
col("text").lower(),
|
||||
)
|
||||
for expr in complex_exprs:
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
function(text=expr, limit=1)
|
||||
|
||||
# Raw native PyExpr is not the public col() wrapper.
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
function(text=col("text")._inner, limit=1)
|
||||
|
||||
# Non-expression / non-literal objects are rejected for field-shaped misuse
|
||||
# when a column binding is required; plain strings are literals for utf8.
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
function(text=object(), limit=1)
|
||||
|
||||
|
||||
def test_function_call_plain_literal_declared_type_null_and_nested():
|
||||
function = _lookup_function()
|
||||
|
||||
authored = function(text="hello", limit=7)
|
||||
assert type(authored) is _authored_call_type()
|
||||
rendered = repr(authored)
|
||||
assert "text=literal(Utf8, null=false)" in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
|
||||
# Plain Python int normalizes to declared Int32 and non-null.
|
||||
authored_int32 = function(text="hello", limit=2_147_483_647)
|
||||
assert type(authored_int32) is _authored_call_type()
|
||||
rendered_int32 = repr(authored_int32)
|
||||
assert "limit=literal(Int32, null=false)" in rendered_int32
|
||||
assert "Int64" not in rendered_int32
|
||||
assert "2147483647" not in rendered_int32
|
||||
|
||||
# Plain None keeps each declared parameter type with null=true.
|
||||
authored_null = function(text=None, limit=None)
|
||||
assert type(authored_null) is _authored_call_type()
|
||||
rendered_null = repr(authored_null)
|
||||
assert "text=literal(Utf8, null=true)" in rendered_null
|
||||
assert "limit=literal(Int32, null=true)" in rendered_null
|
||||
|
||||
list_function = _lookup_function(
|
||||
_sample_function_wire(
|
||||
parameters=[
|
||||
{"name": "values", "data_type_ipc": _LIST_INT32_TYPE_IPC_B64},
|
||||
]
|
||||
)
|
||||
)
|
||||
authored_list = list_function(values=[1, 2, 3])
|
||||
assert type(authored_list) is _authored_call_type()
|
||||
rendered_list = repr(authored_list)
|
||||
assert "values=literal(List(Int32), null=false)" in rendered_list
|
||||
assert "[1, 2, 3]" not in rendered_list
|
||||
|
||||
authored_list_null = list_function(values=None)
|
||||
assert type(authored_list_null) is _authored_call_type()
|
||||
rendered_list_null = repr(authored_list_null)
|
||||
assert "values=literal(List(Int32), null=true)" in rendered_list_null
|
||||
|
||||
|
||||
def test_function_call_direct_literal_expr_exact_type_only():
|
||||
function = _lookup_function()
|
||||
|
||||
# lit(int) is Int64 in the expression builder; int32 parameter must reject it.
|
||||
with pytest.raises((TypeError, ValueError), match="limit|int32|type") as raised:
|
||||
function(text="hello", limit=lit(8))
|
||||
reject_text = _exception_text(raised.value)
|
||||
assert "Int64" in reject_text or "int64" in reject_text.lower()
|
||||
assert "Int32" in reject_text or "int32" in reject_text.lower()
|
||||
|
||||
# Exact utf8 literal expression is accepted and stored as Utf8/non-null.
|
||||
authored = function(text=lit("hello"), limit=8)
|
||||
assert type(authored) is _authored_call_type()
|
||||
rendered = repr(authored)
|
||||
assert "text=literal(Utf8, null=false)" in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
assert "hello" not in rendered
|
||||
|
||||
# Cast / arithmetic around a literal is not a direct Literal node.
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
function(text=lit("hello").cast(pa.string()), limit=8)
|
||||
|
||||
|
||||
def test_function_call_conversion_error_and_repr_are_payload_free():
|
||||
function = _lookup_function()
|
||||
|
||||
with pytest.raises((TypeError, ValueError)) as raised:
|
||||
function(text="ok", limit=_LITERAL_PAYLOAD_SENTINEL)
|
||||
text = _exception_text(raised.value)
|
||||
assert _LITERAL_PAYLOAD_SENTINEL not in text
|
||||
assert "limit" in text
|
||||
assert "int32" in text.lower() or "Int32" in text
|
||||
|
||||
authored = function(text=_LITERAL_PAYLOAD_SENTINEL, limit=_INT_PAYLOAD_SENTINEL)
|
||||
rendered = f"{authored!r}\n{authored!s}"
|
||||
assert _LITERAL_PAYLOAD_SENTINEL not in rendered
|
||||
assert str(_INT_PAYLOAD_SENTINEL) not in rendered
|
||||
assert "text=literal(Utf8, null=false)" in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
assert type(authored).__name__ == "_FunctionCall"
|
||||
assert "_FunctionCall" in rendered
|
||||
|
||||
|
||||
def test_function_call_private_type_nonconstructible_immutable_and_not_exported():
|
||||
function = _lookup_function()
|
||||
authored = function(text=col("text"), limit=1)
|
||||
authored_type = _authored_call_type()
|
||||
|
||||
assert "_FunctionCall" not in getattr(lancedb, "__all__", [])
|
||||
assert not hasattr(lancedb, "_FunctionCall")
|
||||
assert getattr(_native, "_FunctionCall", None) is authored_type
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
authored_type()
|
||||
|
||||
for attr in _OVERDESIGN_ATTRS:
|
||||
assert not hasattr(authored, attr)
|
||||
|
||||
for attr in ("function", "bindings", "arguments", "parameters", "text", "limit"):
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(authored, attr, None)
|
||||
|
||||
# Existing Function handle stays frozen / connection-free / name-free.
|
||||
assert not hasattr(function, "name")
|
||||
assert not hasattr(function, "connection")
|
||||
with pytest.raises(AttributeError):
|
||||
function.id = "mutated"
|
||||
|
||||
|
||||
def test_function_call_does_not_change_col_query_expression_behavior():
|
||||
# Regression guard: authoring must not alter public col()/Expr query behavior.
|
||||
expr = col("firstName") > lit(1)
|
||||
assert isinstance(expr, Expr)
|
||||
assert expr.to_sql() == "(`firstName` > 1)"
|
||||
@@ -0,0 +1,181 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pyarrow as pa
|
||||
|
||||
from lancedb import udf
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"value": pa.int64()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pyarrow==24.0.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
def double_nullable(value):
|
||||
if value is None:
|
||||
return None
|
||||
return value * 2
|
||||
|
||||
|
||||
def test_first_class_function_enterprise_lifecycle():
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
from datetime import timedelta
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import FunctionError
|
||||
from lancedb.expr import col
|
||||
|
||||
host = os.environ.get("LANCEDB_FCF_E2E_HOST")
|
||||
if not host:
|
||||
pytest.skip("LANCEDB_FCF_E2E_HOST is required for the live enterprise test")
|
||||
|
||||
database_uri = os.environ.get("LANCEDB_FCF_E2E_DB_URI", "db://fcf-e2e-local")
|
||||
api_key = os.environ.get("LANCEDB_FCF_E2E_API_KEY", "fake")
|
||||
run_suffix = uuid.uuid4().hex[:12]
|
||||
table_name = f"fcf_e2e_{run_suffix}"
|
||||
function_name = f"fcf_e2e.double_{run_suffix}"
|
||||
job_timeout = timedelta(minutes=5)
|
||||
query_timeout = timedelta(seconds=30)
|
||||
|
||||
def connect():
|
||||
return lancedb.connect(
|
||||
database_uri,
|
||||
api_key=api_key,
|
||||
host_override=host,
|
||||
)
|
||||
|
||||
setup_db = connect()
|
||||
setup_db.create_table(
|
||||
table_name,
|
||||
data=pa.Table.from_pylist(
|
||||
[
|
||||
{"row_id": 1, "value": 2},
|
||||
{"row_id": 2, "value": 5},
|
||||
{"row_id": 3, "value": None},
|
||||
],
|
||||
schema=pa.schema(
|
||||
[
|
||||
pa.field("row_id", pa.int64(), nullable=False),
|
||||
pa.field("value", pa.int64(), nullable=True),
|
||||
]
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
registration_job = setup_db.functions.register(function_name, double_nullable)
|
||||
registration_job_id = registration_job.id
|
||||
assert isinstance(registration_job_id, str) and registration_job_id
|
||||
registered_function = registration_job.wait(timeout=job_timeout)
|
||||
assert type(registered_function) is lancedb.Function
|
||||
assert isinstance(registered_function.id, str) and registered_function.id
|
||||
with pytest.raises(AttributeError):
|
||||
registered_function.id = "mutated"
|
||||
|
||||
catalog_reader = connect()
|
||||
function_by_name = catalog_reader.functions.get(function_name)
|
||||
function_by_id = catalog_reader.functions.get_by_id(registered_function.id)
|
||||
expected_signature = ((("value", pa.int64()),), pa.int64(), True)
|
||||
expected_identity = (
|
||||
registered_function.id,
|
||||
*expected_signature,
|
||||
)
|
||||
for function in (registered_function, function_by_name, function_by_id):
|
||||
assert type(function) is lancedb.Function
|
||||
assert (
|
||||
function.id,
|
||||
function.parameters,
|
||||
function.output_type,
|
||||
function.output_nullable,
|
||||
) == expected_identity
|
||||
|
||||
generated_column_table = catalog_reader.open_table(table_name)
|
||||
generated_column_job = generated_column_table.add_generated_column(
|
||||
"derived",
|
||||
registered_function(value=col("value")),
|
||||
)
|
||||
generated_column_job_id = generated_column_job.id
|
||||
assert isinstance(generated_column_job_id, str) and generated_column_job_id
|
||||
assert generated_column_job.wait(timeout=job_timeout) is None
|
||||
|
||||
complete_reader = connect().open_table(table_name)
|
||||
complete_status = complete_reader.generated_column_status("derived")
|
||||
assert complete_status == "complete"
|
||||
initial_rows = sorted(
|
||||
complete_reader.search()
|
||||
.select(["row_id", "value", "derived"])
|
||||
.limit(3)
|
||||
.to_list(timeout=query_timeout),
|
||||
key=lambda row: row["row_id"],
|
||||
)
|
||||
assert initial_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 5, "derived": 10},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
update_result = complete_reader.update(
|
||||
where="row_id = 2",
|
||||
values={"value": 7},
|
||||
)
|
||||
assert update_result.rows_updated == 1
|
||||
|
||||
incomplete_reader = connect().open_table(table_name)
|
||||
incomplete_status = incomplete_reader.generated_column_status("derived")
|
||||
assert incomplete_status == "incomplete"
|
||||
with pytest.raises(FunctionError) as raised:
|
||||
(
|
||||
incomplete_reader.search()
|
||||
.select(["row_id", "derived"])
|
||||
.limit(3)
|
||||
.to_list(timeout=query_timeout)
|
||||
)
|
||||
assert raised.value.code == "generated_column_incomplete"
|
||||
|
||||
refresh_job = incomplete_reader.refresh_generated_column("derived")
|
||||
refresh_job_id = refresh_job.id
|
||||
assert isinstance(refresh_job_id, str) and refresh_job_id
|
||||
assert refresh_job.wait(timeout=job_timeout) is None
|
||||
|
||||
refreshed_reader = connect().open_table(table_name)
|
||||
refreshed_status = refreshed_reader.generated_column_status("derived")
|
||||
assert refreshed_status == "complete"
|
||||
final_rows = sorted(
|
||||
refreshed_reader.search()
|
||||
.select(["row_id", "value", "derived"])
|
||||
.limit(3)
|
||||
.to_list(timeout=query_timeout),
|
||||
key=lambda row: row["row_id"],
|
||||
)
|
||||
assert final_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 7, "derived": 14},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
evidence = {
|
||||
"run_suffix": run_suffix,
|
||||
"database": database_uri.removeprefix("db://"),
|
||||
"table": table_name,
|
||||
"function": function_name,
|
||||
"function_id": registered_function.id,
|
||||
"job_ids": {
|
||||
"register": registration_job_id,
|
||||
"add_generated_column": generated_column_job_id,
|
||||
"refresh_generated_column": refresh_job_id,
|
||||
},
|
||||
"status_transitions": [
|
||||
complete_status,
|
||||
incomplete_status,
|
||||
refreshed_status,
|
||||
],
|
||||
"final_rows": final_rows,
|
||||
}
|
||||
print(json.dumps(evidence, sort_keys=True, separators=(",", ":")))
|
||||
@@ -0,0 +1,595 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pyarrow as pa
|
||||
|
||||
from lancedb import udf
|
||||
|
||||
|
||||
_RUNNING_DEADLINE_SECONDS = 30
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"value": pa.int64()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pyarrow==24.0.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
def reliable_double(value):
|
||||
if value is None:
|
||||
return None
|
||||
return value * 2
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"value": pa.int64()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pyarrow==24.0.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
def terminate_worker_on_input(value):
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
if len(value) == 0:
|
||||
return value
|
||||
except TypeError:
|
||||
pass
|
||||
|
||||
import os
|
||||
|
||||
os._exit(73)
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"value": pa.int64()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pyarrow==24.0.0"],
|
||||
output_nullable=False,
|
||||
)
|
||||
def slow_triple(value):
|
||||
import time
|
||||
|
||||
time.sleep(0.02)
|
||||
return value * 3
|
||||
|
||||
|
||||
def _require_live() -> str:
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
host = os.environ.get("LANCEDB_FCF_E2E_HOST")
|
||||
if not host:
|
||||
pytest.skip(
|
||||
"LANCEDB_FCF_E2E_HOST is required for live enterprise reliability tests"
|
||||
)
|
||||
return host
|
||||
|
||||
|
||||
def _job_timeout():
|
||||
from datetime import timedelta
|
||||
|
||||
return timedelta(minutes=5)
|
||||
|
||||
|
||||
def _query_timeout():
|
||||
from datetime import timedelta
|
||||
|
||||
return timedelta(seconds=30)
|
||||
|
||||
|
||||
def _connect():
|
||||
import os
|
||||
|
||||
import lancedb
|
||||
|
||||
return lancedb.connect(
|
||||
os.environ.get("LANCEDB_FCF_E2E_DB_URI", "db://fcf-e2e-local"),
|
||||
api_key=os.environ.get("LANCEDB_FCF_E2E_API_KEY", "fake"),
|
||||
host_override=_require_live(),
|
||||
)
|
||||
|
||||
|
||||
def _run_names(case: str) -> tuple[str, str]:
|
||||
import uuid
|
||||
|
||||
suffix = uuid.uuid4().hex[:12]
|
||||
return f"fcf_rel_{case}_{suffix}", f"fcf_rel.{case}_{suffix}"
|
||||
|
||||
|
||||
def _read_rows(table, columns: list[str], row_count: int) -> list[dict]:
|
||||
return sorted(
|
||||
table.search()
|
||||
.select(columns)
|
||||
.limit(row_count)
|
||||
.to_list(timeout=_query_timeout()),
|
||||
key=lambda row: row["row_id"],
|
||||
)
|
||||
|
||||
|
||||
def _emit_evidence(case: str, evidence: dict) -> None:
|
||||
import json
|
||||
|
||||
print(
|
||||
json.dumps(
|
||||
{"case": case, **evidence},
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_enterprise_reliability_core_lifecycle():
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import FunctionError
|
||||
from lancedb.expr import col
|
||||
|
||||
_require_live()
|
||||
table_name, function_name = _run_names("lifecycle")
|
||||
setup_db = _connect()
|
||||
setup_db.create_table(
|
||||
table_name,
|
||||
data=pa.Table.from_pylist(
|
||||
[
|
||||
{"row_id": 1, "value": 2},
|
||||
{"row_id": 2, "value": 5},
|
||||
{"row_id": 3, "value": None},
|
||||
],
|
||||
schema=pa.schema(
|
||||
[
|
||||
pa.field("row_id", pa.int64(), nullable=False),
|
||||
pa.field("value", pa.int64(), nullable=True),
|
||||
]
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
registration_job = setup_db.functions.register(function_name, reliable_double)
|
||||
registration_job_id = registration_job.id
|
||||
assert isinstance(registration_job_id, str) and registration_job_id
|
||||
registered = registration_job.wait(timeout=_job_timeout())
|
||||
assert type(registered) is lancedb.Function
|
||||
assert isinstance(registered.id, str) and registered.id
|
||||
with pytest.raises(AttributeError):
|
||||
registered.id = "mutated"
|
||||
|
||||
catalog_reader = _connect()
|
||||
by_name = catalog_reader.functions.get(function_name)
|
||||
by_id = catalog_reader.functions.get_by_id(registered.id)
|
||||
expected_identity = (
|
||||
registered.id,
|
||||
(("value", pa.int64()),),
|
||||
pa.int64(),
|
||||
True,
|
||||
)
|
||||
for function in (registered, by_name, by_id):
|
||||
assert type(function) is lancedb.Function
|
||||
assert (
|
||||
function.id,
|
||||
function.parameters,
|
||||
function.output_type,
|
||||
function.output_nullable,
|
||||
) == expected_identity
|
||||
|
||||
table = catalog_reader.open_table(table_name)
|
||||
create_job = table.add_generated_column(
|
||||
"derived",
|
||||
registered(value=col("value")),
|
||||
)
|
||||
create_job_id = create_job.id
|
||||
assert isinstance(create_job_id, str) and create_job_id
|
||||
assert create_job.wait(timeout=_job_timeout()) is None
|
||||
|
||||
complete_reader = _connect().open_table(table_name)
|
||||
complete_status = complete_reader.generated_column_status("derived")
|
||||
assert complete_status == "complete"
|
||||
initial_rows = _read_rows(
|
||||
complete_reader,
|
||||
["row_id", "value", "derived"],
|
||||
3,
|
||||
)
|
||||
assert initial_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 5, "derived": 10},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
complete_reader.update(where="row_id = 2", values={"value": 7})
|
||||
|
||||
incomplete_reader = _connect().open_table(table_name)
|
||||
changed_rows = _read_rows(incomplete_reader, ["row_id", "value"], 3)
|
||||
assert changed_rows == [
|
||||
{"row_id": 1, "value": 2},
|
||||
{"row_id": 2, "value": 7},
|
||||
{"row_id": 3, "value": None},
|
||||
]
|
||||
incomplete_status = incomplete_reader.generated_column_status("derived")
|
||||
assert incomplete_status == "incomplete"
|
||||
with pytest.raises(FunctionError) as raised:
|
||||
(
|
||||
incomplete_reader.search()
|
||||
.select(["row_id", "derived"])
|
||||
.limit(3)
|
||||
.to_list(timeout=_query_timeout())
|
||||
)
|
||||
assert raised.value.code == "generated_column_incomplete"
|
||||
|
||||
refresh_job = incomplete_reader.refresh_generated_column("derived")
|
||||
refresh_job_id = refresh_job.id
|
||||
assert isinstance(refresh_job_id, str) and refresh_job_id
|
||||
assert refresh_job.wait(timeout=_job_timeout()) is None
|
||||
|
||||
refreshed_reader = _connect().open_table(table_name)
|
||||
refreshed_status = refreshed_reader.generated_column_status("derived")
|
||||
assert refreshed_status == "complete"
|
||||
final_rows = _read_rows(
|
||||
refreshed_reader,
|
||||
["row_id", "value", "derived"],
|
||||
3,
|
||||
)
|
||||
assert final_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 7, "derived": 14},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
_emit_evidence(
|
||||
"core_lifecycle",
|
||||
{
|
||||
"final_rows": final_rows,
|
||||
"function_id": registered.id,
|
||||
"job_ids": {
|
||||
"create": create_job_id,
|
||||
"refresh": refresh_job_id,
|
||||
"register": registration_job_id,
|
||||
},
|
||||
"status": [
|
||||
complete_status,
|
||||
incomplete_status,
|
||||
refreshed_status,
|
||||
],
|
||||
"table": table_name,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_enterprise_reliability_restart_retention():
|
||||
import json
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
|
||||
_require_live()
|
||||
raw_evidence = os.environ.get("LANCEDB_FCF_E2E_RESTART_EVIDENCE")
|
||||
if not raw_evidence:
|
||||
pytest.skip(
|
||||
"LANCEDB_FCF_E2E_RESTART_EVIDENCE is required for restart retention"
|
||||
)
|
||||
|
||||
try:
|
||||
evidence = json.loads(raw_evidence)
|
||||
except json.JSONDecodeError as error:
|
||||
pytest.fail(f"LANCEDB_FCF_E2E_RESTART_EVIDENCE must be valid JSON: {error.msg}")
|
||||
|
||||
assert isinstance(evidence, dict), (
|
||||
"LANCEDB_FCF_E2E_RESTART_EVIDENCE must be a JSON object"
|
||||
)
|
||||
table_name = evidence.get("table")
|
||||
function_id = evidence.get("function_id")
|
||||
raw_job_ids = evidence.get("job_ids")
|
||||
assert isinstance(table_name, str) and table_name, (
|
||||
"restart evidence must contain a non-empty table"
|
||||
)
|
||||
assert isinstance(function_id, str) and function_id, (
|
||||
"restart evidence must contain a non-empty function_id"
|
||||
)
|
||||
assert isinstance(raw_job_ids, dict), (
|
||||
"restart evidence must contain a job_ids object"
|
||||
)
|
||||
job_ids = {}
|
||||
for job_kind in ("register", "create", "refresh"):
|
||||
job_id = raw_job_ids.get(job_kind)
|
||||
assert isinstance(job_id, str) and job_id, (
|
||||
f"restart evidence must contain a non-empty job_ids.{job_kind}"
|
||||
)
|
||||
job_ids[job_kind] = job_id
|
||||
|
||||
db = _connect()
|
||||
function = db.functions.get_by_id(function_id)
|
||||
expected_identity = (
|
||||
function_id,
|
||||
(("value", pa.int64()),),
|
||||
pa.int64(),
|
||||
True,
|
||||
)
|
||||
assert type(function) is lancedb.Function
|
||||
assert (
|
||||
function.id,
|
||||
function.parameters,
|
||||
function.output_type,
|
||||
function.output_nullable,
|
||||
) == expected_identity
|
||||
|
||||
jobs = {}
|
||||
for job_kind in ("register", "create", "refresh"):
|
||||
job = db.get_job(job_ids[job_kind])
|
||||
assert job is not None
|
||||
assert job.job_id == job_ids[job_kind]
|
||||
assert job.state == "finished"
|
||||
assert job.failure is None
|
||||
jobs[job_kind] = job
|
||||
|
||||
registered_result = jobs["register"].result
|
||||
assert type(registered_result) is lancedb.Function
|
||||
assert (
|
||||
registered_result.id,
|
||||
registered_result.parameters,
|
||||
registered_result.output_type,
|
||||
registered_result.output_nullable,
|
||||
) == expected_identity
|
||||
assert jobs["create"].result is None
|
||||
assert jobs["refresh"].result is None
|
||||
|
||||
table = db.open_table(table_name)
|
||||
status = table.generated_column_status("derived")
|
||||
assert status == "complete"
|
||||
assert table.count_rows() == 3
|
||||
final_rows = _read_rows(table, ["row_id", "value", "derived"], 3)
|
||||
assert final_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 7, "derived": 14},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
_emit_evidence(
|
||||
"restart_retention",
|
||||
{
|
||||
"final_rows": final_rows,
|
||||
"function_id": function_id,
|
||||
"generated_column_status": status,
|
||||
"job_ids": job_ids,
|
||||
"job_states": {
|
||||
job_kind: jobs[job_kind].state
|
||||
for job_kind in ("register", "create", "refresh")
|
||||
},
|
||||
"table": table_name,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_enterprise_reliability_failure_atomicity_and_worker_recovery():
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import JobFailedError
|
||||
from lancedb.expr import col
|
||||
|
||||
_require_live()
|
||||
table_name, failing_function_name = _run_names("worker_failure")
|
||||
_, healthy_function_name = _run_names("worker_recovery")
|
||||
row_count = 4
|
||||
setup_db = _connect()
|
||||
setup_db.create_table(
|
||||
table_name,
|
||||
data=pa.table(
|
||||
{
|
||||
"row_id": list(range(row_count)),
|
||||
"value": [1, 2, 3, 4],
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
registration_job = setup_db.functions.register(
|
||||
failing_function_name,
|
||||
terminate_worker_on_input,
|
||||
)
|
||||
failing_function = registration_job.wait(timeout=_job_timeout())
|
||||
assert type(failing_function) is lancedb.Function
|
||||
|
||||
table = setup_db.open_table(table_name)
|
||||
failed_create_job = table.add_generated_column(
|
||||
"must_not_publish",
|
||||
failing_function(value=col("value")),
|
||||
)
|
||||
failed_job_id = failed_create_job.id
|
||||
assert isinstance(failed_job_id, str) and failed_job_id
|
||||
with pytest.raises(JobFailedError) as raised:
|
||||
failed_create_job.wait(timeout=_job_timeout())
|
||||
assert raised.value.error_code == "udf_execution_failure"
|
||||
|
||||
first_description = _connect().get_job(failed_job_id)
|
||||
second_description = _connect().get_job(failed_job_id)
|
||||
for description in (first_description, second_description):
|
||||
assert description is not None
|
||||
assert description.job_id == failed_job_id
|
||||
assert description.state == "failed"
|
||||
assert description.failure is not None
|
||||
assert description.failure.error_code == "udf_execution_failure"
|
||||
|
||||
atomic_reader = _connect().open_table(table_name)
|
||||
assert "must_not_publish" not in atomic_reader.schema.names
|
||||
assert _read_rows(atomic_reader, ["row_id", "value"], row_count) == [
|
||||
{"row_id": 0, "value": 1},
|
||||
{"row_id": 1, "value": 2},
|
||||
{"row_id": 2, "value": 3},
|
||||
{"row_id": 3, "value": 4},
|
||||
]
|
||||
|
||||
healthy_registration_job = setup_db.functions.register(
|
||||
healthy_function_name,
|
||||
reliable_double,
|
||||
)
|
||||
healthy_function = healthy_registration_job.wait(timeout=_job_timeout())
|
||||
assert type(healthy_function) is lancedb.Function
|
||||
recovery_job = atomic_reader.add_generated_column(
|
||||
"recovered",
|
||||
healthy_function(value=col("value")),
|
||||
)
|
||||
recovery_job_id = recovery_job.id
|
||||
assert isinstance(recovery_job_id, str) and recovery_job_id
|
||||
assert recovery_job.wait(timeout=_job_timeout()) is None
|
||||
|
||||
recovered_reader = _connect().open_table(table_name)
|
||||
assert "must_not_publish" not in recovered_reader.schema.names
|
||||
assert recovered_reader.generated_column_status("recovered") == "complete"
|
||||
recovered_rows = _read_rows(
|
||||
recovered_reader,
|
||||
["row_id", "value", "recovered"],
|
||||
row_count,
|
||||
)
|
||||
assert recovered_rows == [
|
||||
{"row_id": 0, "value": 1, "recovered": 2},
|
||||
{"row_id": 1, "value": 2, "recovered": 4},
|
||||
{"row_id": 2, "value": 3, "recovered": 6},
|
||||
{"row_id": 3, "value": 4, "recovered": 8},
|
||||
]
|
||||
|
||||
_emit_evidence(
|
||||
"failure_atomicity_and_worker_recovery",
|
||||
{
|
||||
"failure_code": first_description.failure.error_code,
|
||||
"failed_job_id": failed_job_id,
|
||||
"recovered_rows": recovered_rows,
|
||||
"recovery_job_id": recovery_job_id,
|
||||
"table": table_name,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_enterprise_reliability_concurrent_refresh_fencing():
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import FunctionError, JobFailedError
|
||||
from lancedb.expr import col
|
||||
|
||||
_require_live()
|
||||
table_name, function_name = _run_names("refresh_fencing")
|
||||
row_count = 1024
|
||||
setup_db = _connect()
|
||||
setup_db.create_table(
|
||||
table_name,
|
||||
data=pa.table(
|
||||
{
|
||||
"row_id": list(range(row_count)),
|
||||
"value": list(range(row_count)),
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
registration_job = setup_db.functions.register(function_name, slow_triple)
|
||||
function = registration_job.wait(timeout=_job_timeout())
|
||||
assert type(function) is lancedb.Function
|
||||
|
||||
table = setup_db.open_table(table_name)
|
||||
create_job = table.add_generated_column(
|
||||
"derived",
|
||||
function(value=col("value")),
|
||||
)
|
||||
assert create_job.wait(timeout=_job_timeout()) is None
|
||||
initial_reader = _connect().open_table(table_name)
|
||||
assert initial_reader.generated_column_status("derived") == "complete"
|
||||
|
||||
initial_reader.update(where="row_id = 0", values={"value": 10_000})
|
||||
incomplete_reader = _connect().open_table(table_name)
|
||||
assert incomplete_reader.generated_column_status("derived") == "incomplete"
|
||||
|
||||
refresh_job = incomplete_reader.refresh_generated_column("derived")
|
||||
refresh_job_id = refresh_job.id
|
||||
assert isinstance(refresh_job_id, str) and refresh_job_id
|
||||
deadline = time.monotonic() + _RUNNING_DEADLINE_SECONDS
|
||||
observed_states = []
|
||||
running_observations = 0
|
||||
while running_observations < 2:
|
||||
state = refresh_job.status()
|
||||
if not observed_states or observed_states[-1] != state:
|
||||
observed_states.append(state)
|
||||
if state == "running":
|
||||
running_observations += 1
|
||||
else:
|
||||
running_observations = 0
|
||||
assert state not in {"finished", "failed", "cancelled"}
|
||||
assert time.monotonic() < deadline
|
||||
if running_observations < 2:
|
||||
time.sleep(0.05)
|
||||
|
||||
concurrent_writer = _connect().open_table(table_name)
|
||||
concurrent_writer.update(where="row_id = 1", values={"value": 20_000})
|
||||
with pytest.raises(JobFailedError) as raised:
|
||||
refresh_job.wait(timeout=_job_timeout())
|
||||
assert raised.value.error_code == "stale_or_conflicting_input"
|
||||
|
||||
stale_job = _connect().get_job(refresh_job_id)
|
||||
assert stale_job is not None
|
||||
assert stale_job.job_id == refresh_job_id
|
||||
assert stale_job.state == "failed"
|
||||
assert stale_job.failure is not None
|
||||
assert stale_job.failure.error_code == raised.value.error_code
|
||||
if observed_states[-1] != stale_job.state:
|
||||
observed_states.append(stale_job.state)
|
||||
|
||||
stale_reader = _connect().open_table(table_name)
|
||||
stale_rows = _read_rows(stale_reader, ["row_id", "value"], row_count)
|
||||
assert len(stale_rows) == row_count
|
||||
for row_id, row in enumerate(stale_rows):
|
||||
expected_value = 10_000 if row_id == 0 else 20_000 if row_id == 1 else row_id
|
||||
assert (row["row_id"], row["value"]) == (row_id, expected_value)
|
||||
assert stale_reader.generated_column_status("derived") == "incomplete"
|
||||
with pytest.raises(FunctionError) as incomplete:
|
||||
(
|
||||
stale_reader.search()
|
||||
.select(["row_id", "derived"])
|
||||
.limit(row_count)
|
||||
.to_list(timeout=_query_timeout())
|
||||
)
|
||||
assert incomplete.value.code == "generated_column_incomplete"
|
||||
|
||||
resubmitted_job = stale_reader.refresh_generated_column("derived")
|
||||
resubmitted_job_id = resubmitted_job.id
|
||||
assert isinstance(resubmitted_job_id, str) and resubmitted_job_id
|
||||
assert resubmitted_job.wait(timeout=_job_timeout()) is None
|
||||
|
||||
final_reader = _connect().open_table(table_name)
|
||||
final_status = final_reader.generated_column_status("derived")
|
||||
assert final_status == "complete"
|
||||
final_rows = _read_rows(
|
||||
final_reader,
|
||||
["row_id", "value", "derived"],
|
||||
row_count,
|
||||
)
|
||||
assert len(final_rows) == row_count
|
||||
final_checksum = 0
|
||||
for row_id, row in enumerate(final_rows):
|
||||
expected_value = 10_000 if row_id == 0 else 20_000 if row_id == 1 else row_id
|
||||
assert (row["row_id"], row["value"], row["derived"]) == (
|
||||
row_id,
|
||||
expected_value,
|
||||
expected_value * 3,
|
||||
)
|
||||
final_checksum += row["derived"]
|
||||
|
||||
_emit_evidence(
|
||||
"concurrent_refresh_fencing",
|
||||
{
|
||||
"failure_code": stale_job.failure.error_code,
|
||||
"final_checksum": final_checksum,
|
||||
"final_status": final_status,
|
||||
"observed_states": observed_states,
|
||||
"resubmitted_job_id": resubmitted_job_id,
|
||||
"row_count": row_count,
|
||||
"stale_job_id": refresh_job_id,
|
||||
"table": table_name,
|
||||
},
|
||||
)
|
||||
@@ -0,0 +1,268 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract: Python projection of JobFailure.error_code / JobFailedError.error_code.
|
||||
|
||||
Public Function failures expose eight stable string categories. Asynchronous
|
||||
errors remain the unified JobFailedError and JobFailureInfo. Python must
|
||||
project the optional exact error_code string already supplied structurally by
|
||||
Rust: preserve a known code, preserve an unknown nonempty future code
|
||||
byte-for-byte, and return None for legacy failure payloads without error_code.
|
||||
Never infer or override a code from message, phase, retryable, HTTP status,
|
||||
job type, or state.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from datetime import timedelta
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import JobFailedError
|
||||
|
||||
_DESCRIBE_PATH = "/v1/jobs/describe"
|
||||
_KNOWN_CODE = "name_or_function_not_found"
|
||||
_CONFLICTING_STABLE_IN_MESSAGE = "definition_validation_failure"
|
||||
_UNKNOWN_CODE = "enterprise_future_category_xyz"
|
||||
_WAIT_KNOWN_CODE = "unsupported_runtime_or_capability"
|
||||
_WAIT_CONFLICTING_IN_MESSAGE = "revoked_function"
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _failed_describe_body(
|
||||
*,
|
||||
job_id: str,
|
||||
error_code: Optional[str] = None,
|
||||
include_error_code: bool = True,
|
||||
phase: str = "execute",
|
||||
message: str = "worker died",
|
||||
retryable: bool = False,
|
||||
job_type: str = "create_index",
|
||||
) -> dict[str, Any]:
|
||||
failure: dict[str, Any] = {
|
||||
"phase": phase,
|
||||
"message": message,
|
||||
"retryable": retryable,
|
||||
}
|
||||
if include_error_code:
|
||||
failure["error_code"] = error_code
|
||||
return {
|
||||
"job_id": job_id,
|
||||
"job_type": job_type,
|
||||
"job_state": "FAILED",
|
||||
"creation_ms": 1000,
|
||||
"spec": {},
|
||||
"failure": failure,
|
||||
}
|
||||
|
||||
|
||||
def _describe_handler(bodies_by_job_id: dict[str, dict[str, Any]]):
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _DESCRIBE_PATH
|
||||
payload = json.loads(_read_body(request).decode("utf-8") or "{}")
|
||||
job_id = payload["job_id"]
|
||||
body = bodies_by_job_id.get(job_id)
|
||||
if body is None:
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
return
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def test_get_job_failure_error_code_known_not_inferred_from_message():
|
||||
"""Structural error_code wins; conflicting message text must not override."""
|
||||
body = _failed_describe_body(
|
||||
job_id="job-known",
|
||||
error_code=_KNOWN_CODE,
|
||||
phase="validate",
|
||||
message=f"looks like {_CONFLICTING_STABLE_IN_MESSAGE}",
|
||||
retryable=False,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-known": body})) as db:
|
||||
description = db.get_job("job-known")
|
||||
assert description is not None
|
||||
failure = description.failure
|
||||
assert failure is not None
|
||||
assert failure.error_code == _KNOWN_CODE
|
||||
assert failure.error_code != _CONFLICTING_STABLE_IN_MESSAGE
|
||||
assert failure.phase == "validate"
|
||||
assert failure.message == f"looks like {_CONFLICTING_STABLE_IN_MESSAGE}"
|
||||
assert failure.retryable is False
|
||||
|
||||
|
||||
def test_get_job_failure_error_code_unknown_preserved_byte_for_byte():
|
||||
body = _failed_describe_body(
|
||||
job_id="job-unknown",
|
||||
error_code=_UNKNOWN_CODE,
|
||||
phase="execute",
|
||||
message=f"new category mentioning {_KNOWN_CODE}",
|
||||
retryable=True,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-unknown": body})) as db:
|
||||
failure = db.get_job("job-unknown").failure
|
||||
assert failure.error_code == _UNKNOWN_CODE
|
||||
assert failure.error_code != _KNOWN_CODE
|
||||
|
||||
|
||||
def test_get_job_failure_error_code_absent_is_none():
|
||||
"""Legacy describe payloads without error_code must not invent a category."""
|
||||
body = _failed_describe_body(
|
||||
job_id="job-legacy",
|
||||
include_error_code=False,
|
||||
phase="execute",
|
||||
message=f"{_KNOWN_CODE} in logs",
|
||||
retryable=True,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-legacy": body})) as db:
|
||||
failure = db.get_job("job-legacy").failure
|
||||
assert failure.error_code is None
|
||||
assert failure.phase == "execute"
|
||||
assert failure.retryable is True
|
||||
|
||||
|
||||
def test_sync_job_wait_job_failed_error_code_known_not_inferred():
|
||||
body = _failed_describe_body(
|
||||
job_id="job-wait-known",
|
||||
error_code=_WAIT_KNOWN_CODE,
|
||||
phase="dispatch",
|
||||
message=f"{_WAIT_CONFLICTING_IN_MESSAGE} in transport logs",
|
||||
retryable=False,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-wait-known": body})) as db:
|
||||
with pytest.raises(JobFailedError) as exc_info:
|
||||
db.job("job-wait-known").wait(timeout=timedelta(seconds=5))
|
||||
|
||||
err = exc_info.value
|
||||
assert isinstance(err, JobFailedError)
|
||||
assert err.error_code == _WAIT_KNOWN_CODE
|
||||
assert err.error_code != _WAIT_CONFLICTING_IN_MESSAGE
|
||||
|
||||
|
||||
def test_sync_job_wait_job_failed_error_code_absent_is_none():
|
||||
body = _failed_describe_body(
|
||||
job_id="job-wait-legacy",
|
||||
include_error_code=False,
|
||||
phase="execute",
|
||||
message=f"{_WAIT_KNOWN_CODE} mentioned only in message",
|
||||
retryable=True,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-wait-legacy": body})) as db:
|
||||
with pytest.raises(JobFailedError) as exc_info:
|
||||
db.job("job-wait-legacy").wait(timeout=timedelta(seconds=5))
|
||||
|
||||
assert exc_info.value.error_code is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_job_wait_job_failed_error_code_unknown_preserved():
|
||||
body = _failed_describe_body(
|
||||
job_id="job-wait-unknown",
|
||||
error_code=_UNKNOWN_CODE,
|
||||
phase="execute",
|
||||
message=f"future code with {_WAIT_KNOWN_CODE} in text",
|
||||
retryable=False,
|
||||
)
|
||||
async with _mock_remote_db_async(
|
||||
_describe_handler({"job-wait-unknown": body})
|
||||
) as db:
|
||||
with pytest.raises(JobFailedError) as exc_info:
|
||||
await db.job("job-wait-unknown").wait(timeout=timedelta(seconds=5))
|
||||
|
||||
err = exc_info.value
|
||||
assert err.error_code == _UNKNOWN_CODE
|
||||
assert err.error_code != _WAIT_KNOWN_CODE
|
||||
|
||||
|
||||
def test_job_failed_error_legacy_message_construction_error_code_is_none():
|
||||
err = JobFailedError("legacy construction with only a message")
|
||||
assert err.error_code is None
|
||||
|
||||
|
||||
def test_job_failed_error_error_code_is_read_only():
|
||||
err = JobFailedError("message")
|
||||
with pytest.raises(AttributeError):
|
||||
err.error_code = _KNOWN_CODE
|
||||
@@ -0,0 +1,634 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python first-class Function catalog lookup."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.remote.errors import HttpError
|
||||
|
||||
|
||||
def _function_error_cls(*, required: bool = True):
|
||||
"""Resolve FunctionError from the live module (records RED when absent)."""
|
||||
from lancedb import exceptions as exc_mod
|
||||
|
||||
cls = getattr(exc_mod, "FunctionError", None)
|
||||
if cls is None:
|
||||
if required:
|
||||
pytest.fail("lancedb.exceptions.FunctionError is missing")
|
||||
return type("MissingFunctionError", (), {})
|
||||
return cls
|
||||
|
||||
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_LOOKUP_CATALOG_NAME = "text.normalize.lookup-name"
|
||||
_LOOKUP_FUNCTION_ID = "fn.exact.lookup-handle"
|
||||
_LOOKUP_SERVER_MESSAGE_MARKER = (
|
||||
"SERVER_LOOKUP_DIAGNOSTIC_MARKER name=text.normalize.lookup-name "
|
||||
"id=fn.exact.lookup-handle"
|
||||
)
|
||||
_SENSITIVE_BODY_MARKER = "SENSITIVE_LOOKUP_BODY_MARKER"
|
||||
_UNKNOWN_CODE = "enterprise_future_lookup_category_xyz"
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as job-result
|
||||
# tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust FileWriter.
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_DELETED_LOOKUP_KEYWORDS = (
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"replace",
|
||||
"expected_current_function_id",
|
||||
"list",
|
||||
"alias",
|
||||
"lineage",
|
||||
"FunctionVersion",
|
||||
)
|
||||
|
||||
|
||||
def _sample_function_wire() -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": _LOOKUP_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _lookup_success_body(
|
||||
*,
|
||||
function: dict[str, Any] | None = None,
|
||||
extra_outer: dict[str, Any] | None = None,
|
||||
) -> bytes:
|
||||
body: dict[str, Any] = {"function": function or _sample_function_wire()}
|
||||
if extra_outer:
|
||||
body.update(extra_outer)
|
||||
return json.dumps(body).encode("utf-8")
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_chain_text(exc: BaseException) -> str:
|
||||
parts: list[str] = []
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(str(current))
|
||||
parts.append(repr(current))
|
||||
current = current.__cause__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _assert_payload_free(exc: BaseException) -> None:
|
||||
text = _exception_chain_text(exc)
|
||||
assert _LOOKUP_SERVER_MESSAGE_MARKER not in text
|
||||
assert _LOOKUP_CATALOG_NAME not in text
|
||||
assert _LOOKUP_FUNCTION_ID not in text
|
||||
assert _SENSITIVE_BODY_MARKER not in text
|
||||
|
||||
|
||||
def _assert_exact_lookup_function(function: object) -> None:
|
||||
assert type(function) is lancedb.Function
|
||||
assert function.id == _LOOKUP_FUNCTION_ID
|
||||
assert not hasattr(function, "name")
|
||||
assert function.parameters == (
|
||||
("text", pa.string()),
|
||||
("limit", pa.int32()),
|
||||
)
|
||||
assert function.output_type == pa.string()
|
||||
assert function.output_nullable is True
|
||||
assert _LOOKUP_CATALOG_NAME not in repr(function)
|
||||
assert _LOOKUP_CATALOG_NAME not in str(function)
|
||||
|
||||
|
||||
def _assert_name_request(raw: bytes, body: dict[str, Any]) -> None:
|
||||
assert raw
|
||||
assert body == {"name": _LOOKUP_CATALOG_NAME}
|
||||
assert "function_id" not in body
|
||||
|
||||
|
||||
def _assert_id_request(raw: bytes, body: dict[str, Any]) -> None:
|
||||
assert raw
|
||||
assert body == {"function_id": _LOOKUP_FUNCTION_ID}
|
||||
assert "name" not in body
|
||||
|
||||
|
||||
def _assert_native_lookup_methods_present() -> None:
|
||||
assert hasattr(_native.Connection, "_lookup_function_by_name")
|
||||
assert hasattr(_native.Connection, "_lookup_function_by_id")
|
||||
assert callable(getattr(_native.Connection, "_lookup_function_by_name"))
|
||||
assert callable(getattr(_native.Connection, "_lookup_function_by_id"))
|
||||
|
||||
|
||||
def test_native_connection_exposes_private_lookup_methods():
|
||||
_assert_native_lookup_methods_present()
|
||||
|
||||
|
||||
def test_sync_remote_get_by_name_exact_request_and_function_shape():
|
||||
_assert_native_lookup_methods_present()
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _LOOKUP_PATH
|
||||
raw = _read_body(request)
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
assert not hasattr(db, "get_function")
|
||||
function = db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
_assert_name_request(seen["raw"], seen["body"])
|
||||
_assert_exact_lookup_function(function)
|
||||
|
||||
|
||||
def test_sync_remote_get_by_id_exact_request_and_function_shape():
|
||||
_assert_native_lookup_methods_present()
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _LOOKUP_PATH
|
||||
raw = _read_body(request)
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
function = db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
_assert_id_request(seen["raw"], seen["body"])
|
||||
_assert_exact_lookup_function(function)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_get_by_name_and_id():
|
||||
_assert_native_lookup_methods_present()
|
||||
name_seen: dict[str, Any] = {}
|
||||
id_seen: dict[str, Any] = {}
|
||||
stage = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _LOOKUP_PATH
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
stage["n"] += 1
|
||||
if stage["n"] == 1:
|
||||
name_seen["raw"] = raw
|
||||
name_seen["body"] = body
|
||||
else:
|
||||
id_seen["raw"] = raw
|
||||
id_seen["body"] = body
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
|
||||
async with _mock_remote_db_async(handler) as db:
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
by_name = await db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
by_id = await db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
_assert_name_request(name_seen["raw"], name_seen["body"])
|
||||
_assert_id_request(id_seen["raw"], id_seen["body"])
|
||||
_assert_exact_lookup_function(by_name)
|
||||
_assert_exact_lookup_function(by_id)
|
||||
|
||||
|
||||
def test_sync_remote_get_accepts_additive_outer_success_fields():
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
_lookup_success_body(
|
||||
extra_outer={
|
||||
"server_extra": {"ok": True},
|
||||
"request_echo_name": _LOOKUP_CATALOG_NAME,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
function = db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
_assert_exact_lookup_function(function)
|
||||
|
||||
|
||||
def test_empty_name_and_id_reject_before_transport():
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
_read_body(request)
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"should not be reached")
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.get("")
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.get_by_id("")
|
||||
|
||||
assert received["n"] == 0
|
||||
|
||||
|
||||
def test_local_sync_lookup_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_lookup_methods_present()
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_lookup_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_lookup_methods_present()
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
|
||||
|
||||
def test_explicit_known_code_is_function_error_with_exact_code():
|
||||
body = {
|
||||
"error_code": "name_or_function_not_found",
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
"looks_like": "definition_validation_failure",
|
||||
}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
request.send_response(404)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
err = exc_info.value
|
||||
assert isinstance(err, function_error)
|
||||
assert err.code == "name_or_function_not_found"
|
||||
assert err.code != "definition_validation_failure"
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
def test_explicit_unknown_code_preserved_despite_status_and_message():
|
||||
body = {
|
||||
"error_code": _UNKNOWN_CODE,
|
||||
"message": (
|
||||
f"{_LOOKUP_SERVER_MESSAGE_MARKER} revoked_function "
|
||||
"name_or_function_not_found"
|
||||
),
|
||||
}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
raw = _read_body(request)
|
||||
assert json.loads(raw.decode("utf-8")) == {"function_id": _LOOKUP_FUNCTION_ID}
|
||||
request.send_response(409)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
err = exc_info.value
|
||||
assert err.code == _UNKNOWN_CODE
|
||||
assert err.code != "revoked_function"
|
||||
assert err.code != "name_or_function_not_found"
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,status,response_body",
|
||||
[
|
||||
(
|
||||
"missing_code_404",
|
||||
404,
|
||||
{
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"empty_code",
|
||||
400,
|
||||
{
|
||||
"error_code": "",
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"wrong_type_code",
|
||||
400,
|
||||
{
|
||||
"error_code": 123,
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"null_code",
|
||||
404,
|
||||
{
|
||||
"error_code": None,
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"non_json",
|
||||
404,
|
||||
f"not-json {_LOOKUP_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_invalid_or_missing_error_code_is_payload_free_http(
|
||||
label: str, status: int, response_body: object
|
||||
):
|
||||
del label # parametrize label for failure diagnosis only
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
request.send_response(status)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
if isinstance(response_body, str):
|
||||
request.wfile.write(response_body.encode("utf-8"))
|
||||
else:
|
||||
request.wfile.write(json.dumps(response_body).encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert not isinstance(err, _function_error_cls(required=False))
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,response_body",
|
||||
[
|
||||
(
|
||||
"missing_function",
|
||||
{
|
||||
"server_extra": True,
|
||||
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
},
|
||||
),
|
||||
(
|
||||
"null_function",
|
||||
{
|
||||
"function": None,
|
||||
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
},
|
||||
),
|
||||
(
|
||||
"wrong_type_function",
|
||||
{
|
||||
"function": "not-an-object",
|
||||
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
},
|
||||
),
|
||||
(
|
||||
"invalid_function_shape",
|
||||
{
|
||||
"function": {
|
||||
"format_version": 1,
|
||||
"id": _LOOKUP_FUNCTION_ID,
|
||||
# missing signature
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
}
|
||||
},
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_malformed_success_is_payload_free_http(label: str, response_body: dict):
|
||||
del label
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(response_body).encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert not isinstance(err, _function_error_cls(required=False))
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
def test_function_error_surface_omits_server_marker_name_and_id():
|
||||
body = {
|
||||
"error_code": "name_or_function_not_found",
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
"function_id": _LOOKUP_FUNCTION_ID,
|
||||
"name": _LOOKUP_CATALOG_NAME,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
_read_body(request)
|
||||
request.send_response(404)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
err = exc_info.value
|
||||
_assert_payload_free(err)
|
||||
assert getattr(err, "code", None) == "name_or_function_not_found"
|
||||
|
||||
|
||||
def test_no_direct_db_lookup_methods_and_no_deleted_keywords():
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
_read_body(request)
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"should not be reached")
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
assert not hasattr(db, "lookup_function")
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
assert not hasattr(db, "get_function")
|
||||
assert not hasattr(db.functions, "get_by_name")
|
||||
assert not hasattr(db.functions, "list")
|
||||
|
||||
for keyword in _DELETED_LOOKUP_KEYWORDS:
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME, **{keyword: True})
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.get_by_id(_LOOKUP_FUNCTION_ID, **{keyword: True})
|
||||
|
||||
assert received["n"] == 0
|
||||
|
||||
|
||||
def test_function_error_is_not_top_level_export():
|
||||
assert not hasattr(lancedb, "FunctionError")
|
||||
function_error = _function_error_cls()
|
||||
assert issubclass(function_error, RuntimeError)
|
||||
@@ -0,0 +1,398 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for Python first-class Function registration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from datetime import timedelta
|
||||
from typing import Any, Callable
|
||||
from unittest import mock
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
import lancedb._udf as _udf_mod
|
||||
import lancedb.job
|
||||
from lancedb import FunctionCapability, udf
|
||||
from lancedb.remote.errors import HttpError
|
||||
|
||||
_SOURCE_MARKER = "registration-source-marker-unique-xyz"
|
||||
_SECRET_REFERENCE = "secret://team/registration-redact-token-xyz"
|
||||
_SECRET_ENV = "REGISTER_API_TOKEN"
|
||||
_NETWORK_ORIGIN = "https://api.registration-example.com"
|
||||
_FUNCTION_NAME = "text.normalize"
|
||||
_FUNCTION_ID_RETRY = "fn.register-retry-1"
|
||||
_JOB_ID_RETRY = "job-register-retry-1"
|
||||
_JOB_ID_ASYNC = "job-register-async-1"
|
||||
_REGISTER_PATH = "/v1/functions/register"
|
||||
_DESCRIBE_PATH = "/v1/jobs/describe"
|
||||
|
||||
_DELETED_REGISTER_KEYWORDS = (
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"replace",
|
||||
"expected_current_function_id",
|
||||
)
|
||||
|
||||
_SPEC_KEYS = {
|
||||
"format_version",
|
||||
"name",
|
||||
"definition",
|
||||
"expected_current_function_id",
|
||||
}
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"text": pa.string(), "limit": pa.int32()},
|
||||
output=pa.string(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=True,
|
||||
capabilities=[
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
],
|
||||
)
|
||||
def packable_register_normalize(text, limit):
|
||||
"""registration-source-marker-unique-xyz."""
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def _definition_json(fn: object) -> dict[str, Any]:
|
||||
payload = _udf_mod._build_function_definition(fn)._to_json()
|
||||
if isinstance(payload, bytes):
|
||||
return json.loads(payload.decode("utf-8"))
|
||||
assert isinstance(payload, str)
|
||||
return json.loads(payload)
|
||||
|
||||
|
||||
def _expected_register_spec(name: str, fn: object) -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"name": name,
|
||||
"definition": _definition_json(fn),
|
||||
"expected_current_function_id": None,
|
||||
}
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_chain_text(exc: BaseException) -> str:
|
||||
parts: list[str] = []
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(str(current))
|
||||
parts.append(repr(current))
|
||||
current = current.__cause__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _assert_markers_absent_from_exception(exc: BaseException) -> None:
|
||||
text = _exception_chain_text(exc)
|
||||
assert _SOURCE_MARKER not in text
|
||||
assert _SECRET_REFERENCE not in text
|
||||
|
||||
|
||||
def _assert_exact_register_spec(body: dict[str, Any], expected: dict[str, Any]) -> None:
|
||||
assert set(body) == _SPEC_KEYS
|
||||
assert body == expected
|
||||
assert body["format_version"] == 1
|
||||
assert body["expected_current_function_id"] is None
|
||||
assert _SOURCE_MARKER in json.dumps(body["definition"])
|
||||
assert any(
|
||||
capability.get("reference") == _SECRET_REFERENCE
|
||||
for capability in body["definition"]["capabilities"]
|
||||
)
|
||||
|
||||
|
||||
def test_sync_remote_register_retries_exact_wire_and_returns_job():
|
||||
expected_spec = _expected_register_spec(_FUNCTION_NAME, packable_register_normalize)
|
||||
attempts: list[dict[str, Any]] = []
|
||||
describe_calls: list[dict[str, Any]] = []
|
||||
function_result_wire = {
|
||||
"kind": "function",
|
||||
"format_version": 1,
|
||||
"function": {
|
||||
"format_version": 1,
|
||||
"id": _FUNCTION_ID_RETRY,
|
||||
"signature": expected_spec["definition"]["signature"],
|
||||
},
|
||||
}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _REGISTER_PATH:
|
||||
request_id = request.headers.get("x-request-id")
|
||||
attempts.append(
|
||||
{
|
||||
"request_id": request_id,
|
||||
"raw": raw,
|
||||
"body": json.loads(raw.decode("utf-8")),
|
||||
}
|
||||
)
|
||||
if len(attempts) == 1:
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"transient register failure")
|
||||
return
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps({"job_id": _JOB_ID_RETRY}).encode("utf-8"))
|
||||
return
|
||||
|
||||
assert request.path == _DESCRIBE_PATH
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body["job_id"] == _JOB_ID_RETRY
|
||||
describe_calls.append(body)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
"job_id": _JOB_ID_RETRY,
|
||||
"job_state": "DONE",
|
||||
"job_type": "register_function",
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
"result": function_result_wire,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
package_calls = {"n": 0}
|
||||
original_package = _udf_mod._package_udf
|
||||
|
||||
def counting_package(fn: object):
|
||||
package_calls["n"] += 1
|
||||
return original_package(fn)
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
assert not hasattr(db, "register_function")
|
||||
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
|
||||
job = db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_RETRY
|
||||
waited = job.wait(timeout=timedelta(seconds=5))
|
||||
|
||||
assert package_calls["n"] == 1
|
||||
assert len(attempts) == 2
|
||||
first, second = attempts
|
||||
assert isinstance(first["request_id"], str) and first["request_id"]
|
||||
assert first["request_id"] == second["request_id"]
|
||||
assert first["raw"] == second["raw"]
|
||||
assert first["raw"]
|
||||
_assert_exact_register_spec(first["body"], expected_spec)
|
||||
_assert_exact_register_spec(second["body"], expected_spec)
|
||||
|
||||
assert len(describe_calls) == 1
|
||||
assert describe_calls[0]["job_id"] == _JOB_ID_RETRY
|
||||
assert type(waited) is lancedb.Function
|
||||
assert waited.id == _FUNCTION_ID_RETRY
|
||||
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
|
||||
assert waited.output_type == pa.string()
|
||||
assert waited.output_nullable is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_register_returns_async_job_with_exact_spec():
|
||||
expected_spec = _expected_register_spec(_FUNCTION_NAME, packable_register_normalize)
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _REGISTER_PATH
|
||||
raw = _read_body(request)
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
seen["request_id"] = request.headers.get("x-request-id")
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
|
||||
|
||||
async with _mock_remote_db_async(handler) as db:
|
||||
assert not hasattr(db, "register_function")
|
||||
job = await db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert seen.get("raw")
|
||||
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
|
||||
_assert_exact_register_spec(seen["body"], expected_spec)
|
||||
assert type(job) is lancedb.job.AsyncJob
|
||||
assert job.id == _JOB_ID_ASYNC
|
||||
|
||||
|
||||
def test_sync_remote_register_http_error_omits_source_and_secret_markers():
|
||||
echoed = f"register failed with {_SOURCE_MARKER} and {_SECRET_REFERENCE}"
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
assert request.path == _REGISTER_PATH
|
||||
_read_body(request)
|
||||
request.send_response(400)
|
||||
request.end_headers()
|
||||
request.wfile.write(echoed.encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert received["n"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert err.status_code == 400
|
||||
_assert_markers_absent_from_exception(err)
|
||||
|
||||
|
||||
def test_empty_name_rejects_before_http():
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
_read_body(request)
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"should not be reached")
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.register("", packable_register_normalize)
|
||||
|
||||
assert received["n"] == 0
|
||||
|
||||
|
||||
def test_local_sync_register_not_implemented_without_table_mutation(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "register_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_register_not_implemented_without_table_mutation(tmp_path):
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "register_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
|
||||
|
||||
@pytest.mark.parametrize("keyword", _DELETED_REGISTER_KEYWORDS)
|
||||
def test_register_rejects_deleted_overdesign_keywords_before_submission(keyword):
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
_read_body(request)
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"should not be reached")
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.register(
|
||||
_FUNCTION_NAME,
|
||||
packable_register_normalize,
|
||||
**{keyword: True},
|
||||
)
|
||||
|
||||
assert received["n"] == 0
|
||||
@@ -0,0 +1,719 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python conditional first-class Function name removal."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.remote.errors import HttpError
|
||||
|
||||
_REMOVE_PATH = "/v1/functions/remove"
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_REMOVE_CATALOG_NAME = "text.normalize.remove-name"
|
||||
_REMOVE_FUNCTION_ID = "fn.exact.remove-handle"
|
||||
_REMOVE_SERVER_MESSAGE_MARKER = (
|
||||
"SERVER_REMOVE_DIAGNOSTIC_MARKER name=text.normalize.remove-name "
|
||||
"id=fn.exact.remove-handle"
|
||||
)
|
||||
_SENSITIVE_BODY_MARKER = "SENSITIVE_REMOVE_BODY_MARKER"
|
||||
_CONFLICTING_MESSAGE_CODE = "revoked_function"
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
|
||||
# replace tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_DELETED_REMOVE_KEYWORDS = (
|
||||
"expected_current_function_id",
|
||||
"function_id",
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"version",
|
||||
"force",
|
||||
"if_exists",
|
||||
"revoke",
|
||||
"delete",
|
||||
)
|
||||
|
||||
|
||||
def _function_error_cls(*, required: bool = True):
|
||||
"""Resolve FunctionError from the live module (records RED when absent)."""
|
||||
from lancedb import exceptions as exc_mod
|
||||
|
||||
cls = getattr(exc_mod, "FunctionError", None)
|
||||
if cls is None:
|
||||
if required:
|
||||
pytest.fail("lancedb.exceptions.FunctionError is missing")
|
||||
return type("MissingFunctionError", (), {})
|
||||
return cls
|
||||
|
||||
|
||||
def _sample_function_wire() -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": _REMOVE_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _lookup_success_body() -> bytes:
|
||||
return json.dumps({"function": _sample_function_wire()}).encode("utf-8")
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
def _close_db(db: Any) -> None:
|
||||
with contextlib.suppress(Exception):
|
||||
inner = getattr(db, "_conn", None)
|
||||
if inner is not None:
|
||||
inner.close()
|
||||
return
|
||||
close = getattr(db, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
db = None
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
if db is not None:
|
||||
_close_db(db)
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
db = None
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
if db is not None:
|
||||
_close_db(db)
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_chain_text(exc: BaseException) -> str:
|
||||
parts: list[str] = []
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(str(current))
|
||||
parts.append(repr(current))
|
||||
current = current.__cause__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _assert_payload_free(exc: BaseException) -> None:
|
||||
text = _exception_chain_text(exc)
|
||||
assert _REMOVE_SERVER_MESSAGE_MARKER not in text
|
||||
assert _REMOVE_CATALOG_NAME not in text
|
||||
assert _REMOVE_FUNCTION_ID not in text
|
||||
assert _SENSITIVE_BODY_MARKER not in text
|
||||
|
||||
|
||||
def _assert_exact_remove_function(function: object) -> None:
|
||||
assert type(function) is lancedb.Function
|
||||
assert function.id == _REMOVE_FUNCTION_ID
|
||||
assert not hasattr(function, "name")
|
||||
assert function.parameters == (
|
||||
("text", pa.string()),
|
||||
("limit", pa.int32()),
|
||||
)
|
||||
assert function.output_type == pa.string()
|
||||
assert function.output_nullable is True
|
||||
assert _REMOVE_CATALOG_NAME not in repr(function)
|
||||
assert _REMOVE_CATALOG_NAME not in str(function)
|
||||
|
||||
|
||||
def _assert_exact_remove_request(
|
||||
request: http.server.BaseHTTPRequestHandler,
|
||||
raw: bytes,
|
||||
body: dict[str, Any],
|
||||
*,
|
||||
expected_id: str,
|
||||
) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _REMOVE_PATH
|
||||
assert "?" not in request.path
|
||||
assert raw
|
||||
assert body == {
|
||||
"name": _REMOVE_CATALOG_NAME,
|
||||
"expected_current_function_id": expected_id,
|
||||
}
|
||||
assert set(body) == {"name", "expected_current_function_id"}
|
||||
assert "format_version" not in body
|
||||
assert "function_id" not in body
|
||||
assert "function" not in body
|
||||
assert "signature" not in body
|
||||
assert "job_id" not in body
|
||||
assert "idempotency_key" not in body
|
||||
assert "user_version" not in body
|
||||
assert "force" not in body
|
||||
assert "if_exists" not in body
|
||||
request_id = request.headers.get("x-request-id")
|
||||
assert isinstance(request_id, str) and request_id
|
||||
|
||||
|
||||
def _assert_native_remove_method_present() -> None:
|
||||
assert hasattr(_native.Connection, "_remove_function_name")
|
||||
assert callable(getattr(_native.Connection, "_remove_function_name"))
|
||||
|
||||
|
||||
def _lookup_success_handler(
|
||||
counters: dict[str, int],
|
||||
*,
|
||||
after_lookup: (
|
||||
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
|
||||
) = None,
|
||||
):
|
||||
"""Serve exact name lookup; optionally continue for remove."""
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
counters["lookup"] = counters.get("lookup", 0) + 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body == {"name": _REMOVE_CATALOG_NAME}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
if after_lookup is not None:
|
||||
after_lookup(request, raw)
|
||||
return
|
||||
counters["remove"] = counters.get("remove", 0) + 1
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected remove")
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _observe_current(db) -> lancedb.Function:
|
||||
current = db.functions.get(_REMOVE_CATALOG_NAME)
|
||||
_assert_exact_remove_function(current)
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
return current
|
||||
|
||||
|
||||
async def _observe_current_async(db) -> lancedb.Function:
|
||||
current = await db.functions.get(_REMOVE_CATALOG_NAME)
|
||||
_assert_exact_remove_function(current)
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
return current
|
||||
|
||||
|
||||
def test_native_connection_exposes_private_remove_function_name():
|
||||
_assert_native_remove_method_present()
|
||||
|
||||
|
||||
def test_sync_remote_remove_exact_body_path_request_id_returns_none():
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
remove_attempts: list[dict[str, Any]] = []
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
body = json.loads(payload.decode("utf-8"))
|
||||
remove_attempts.append(
|
||||
{
|
||||
"request": request,
|
||||
"raw": payload,
|
||||
"body": body,
|
||||
"request_id": request.headers.get("x-request-id"),
|
||||
}
|
||||
)
|
||||
# Illegal body on 204 must be ignored; success is status-driven only.
|
||||
request.send_response(204)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
"message": _REMOVE_SERVER_MESSAGE_MARKER,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
|
||||
result = db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert result is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 1
|
||||
assert len(remove_attempts) == 1
|
||||
attempt = remove_attempts[0]
|
||||
_assert_exact_remove_request(
|
||||
attempt["request"],
|
||||
attempt["raw"],
|
||||
attempt["body"],
|
||||
expected_id=current.id,
|
||||
)
|
||||
assert attempt["body"]["expected_current_function_id"] == current.id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_remove_exact_body_returns_none():
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
seen["request"] = request
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
seen["request_id"] = request.headers.get("x-request-id")
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
async with _mock_remote_db_async(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
current = await _observe_current_async(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
result = await db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert result is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 1
|
||||
assert seen.get("raw")
|
||||
_assert_exact_remove_request(
|
||||
seen["request"],
|
||||
seen["raw"],
|
||||
seen["body"],
|
||||
expected_id=current.id,
|
||||
)
|
||||
|
||||
|
||||
def test_after_remove_name_lookup_not_found_id_lookup_same_function():
|
||||
"""Catalog-pointer SDK sequence via a stateful fixture; not server atomicity."""
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {
|
||||
"lookup_name": 0,
|
||||
"lookup_id": 0,
|
||||
"remove": 0,
|
||||
}
|
||||
removed = {"yes": False}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
if "name" in body:
|
||||
counters["lookup_name"] += 1
|
||||
assert body == {"name": _REMOVE_CATALOG_NAME}
|
||||
if removed["yes"]:
|
||||
request.send_response(404)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
"error_code": "name_or_function_not_found",
|
||||
"message": _REMOVE_SERVER_MESSAGE_MARKER,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
return
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
|
||||
counters["lookup_id"] += 1
|
||||
assert body == {"function_id": _REMOVE_FUNCTION_ID}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
_assert_exact_remove_request(
|
||||
request, raw, body, expected_id=_REMOVE_FUNCTION_ID
|
||||
)
|
||||
removed["yes"] = True
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(handler) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup_name"] == 1
|
||||
assert counters["lookup_id"] == 0
|
||||
assert counters["remove"] == 0
|
||||
|
||||
result = db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
assert result is None
|
||||
assert counters["lookup_name"] == 1
|
||||
assert counters["remove"] == 1
|
||||
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.get(_REMOVE_CATALOG_NAME)
|
||||
err = exc_info.value
|
||||
assert err.code == "name_or_function_not_found"
|
||||
_assert_payload_free(err)
|
||||
|
||||
by_id = db.functions.get_by_id(_REMOVE_FUNCTION_ID)
|
||||
|
||||
assert counters["lookup_name"] == 2
|
||||
assert counters["lookup_id"] == 1
|
||||
assert counters["remove"] == 1
|
||||
_assert_exact_remove_function(by_id)
|
||||
assert by_id.id == current.id
|
||||
assert by_id.parameters == current.parameters
|
||||
assert by_id.output_type == current.output_type
|
||||
assert by_id.output_nullable is current.output_nullable
|
||||
|
||||
|
||||
def test_explicit_name_conflict_is_function_error_payload_free():
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
body = {
|
||||
"error_code": "name_conflict",
|
||||
"message": (
|
||||
f"{_REMOVE_SERVER_MESSAGE_MARKER} looks_like {_CONFLICTING_MESSAGE_CODE}"
|
||||
),
|
||||
"name": _REMOVE_CATALOG_NAME,
|
||||
"function_id": _REMOVE_FUNCTION_ID,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
parsed = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_remove_request(
|
||||
request, payload, parsed, expected_id=_REMOVE_FUNCTION_ID
|
||||
)
|
||||
request.send_response(409)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, function_error)
|
||||
assert err.code == "name_conflict"
|
||||
assert err.code != _CONFLICTING_MESSAGE_CODE
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,status,response_body",
|
||||
[
|
||||
(
|
||||
"200_with_body",
|
||||
200,
|
||||
{
|
||||
"ok": True,
|
||||
"message": _REMOVE_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
"job_id": "must-not-infer-job",
|
||||
},
|
||||
),
|
||||
(
|
||||
"202_empty",
|
||||
202,
|
||||
f"{_REMOVE_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
|
||||
),
|
||||
("200_empty", 200, ""),
|
||||
],
|
||||
)
|
||||
def test_http_200_202_cannot_return_success(
|
||||
label: str, status: int, response_body: object
|
||||
):
|
||||
del label
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
parsed = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_remove_request(
|
||||
request, payload, parsed, expected_id=_REMOVE_FUNCTION_ID
|
||||
)
|
||||
request.send_response(status)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
if isinstance(response_body, str):
|
||||
request.wfile.write(response_body.encode("utf-8"))
|
||||
else:
|
||||
request.wfile.write(json.dumps(response_body).encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert not isinstance(err, _function_error_cls(required=False))
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
def test_empty_name_rejects_before_remove_transport():
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.remove("", current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_current",
|
||||
[
|
||||
_REMOVE_FUNCTION_ID,
|
||||
{"id": _REMOVE_FUNCTION_ID},
|
||||
object(),
|
||||
123,
|
||||
],
|
||||
)
|
||||
def test_raw_id_or_arbitrary_current_rejected_without_remove(bad_current):
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
# Observe a real handle separately so the bad-current path is isolated.
|
||||
_ = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.remove(_REMOVE_CATALOG_NAME, bad_current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
|
||||
|
||||
def test_local_sync_remove_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
_close_db(db)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_remove_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("keyword", _DELETED_REMOVE_KEYWORDS)
|
||||
def test_remove_rejects_deleted_cas_retry_version_kwargs(keyword):
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.remove(
|
||||
_REMOVE_CATALOG_NAME,
|
||||
current,
|
||||
**{keyword: True},
|
||||
)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
|
||||
|
||||
def test_no_direct_remove_methods_and_function_has_no_remove_facade_private():
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
assert callable(getattr(db.functions, "remove", None))
|
||||
assert not hasattr(lancedb, "_SyncFunctions")
|
||||
assert not hasattr(lancedb, "_AsyncFunctions")
|
||||
assert type(db.functions).__name__.startswith("_")
|
||||
assert "_SyncFunctions" not in getattr(lancedb, "__all__", [])
|
||||
assert "_AsyncFunctions" not in getattr(lancedb, "__all__", [])
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
@@ -0,0 +1,579 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python conditional first-class Function replacement."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from datetime import timedelta
|
||||
from typing import Any, Callable
|
||||
from unittest import mock
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
import lancedb._udf as _udf_mod
|
||||
import lancedb.job
|
||||
from lancedb import FunctionCapability, udf
|
||||
from lancedb.exceptions import JobFailedError
|
||||
|
||||
_SOURCE_MARKER = "replace-source-marker-unique-xyz"
|
||||
_SECRET_REFERENCE = "secret://team/replace-redact-token-xyz"
|
||||
_SECRET_ENV = "REPLACE_API_TOKEN"
|
||||
_NETWORK_ORIGIN = "https://api.replace-example.com"
|
||||
_FUNCTION_NAME = "text.normalize"
|
||||
_CURRENT_FUNCTION_ID = "fn.replace-current-1"
|
||||
_REPLACED_FUNCTION_ID = "fn.replace-result-1"
|
||||
_JOB_ID_SYNC = "job-replace-sync-1"
|
||||
_JOB_ID_ASYNC = "job-replace-async-1"
|
||||
_JOB_ID_CONFLICT = "job-replace-conflict-1"
|
||||
_REGISTER_PATH = "/v1/functions/register"
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_DESCRIBE_PATH = "/v1/jobs/describe"
|
||||
_CONFLICTING_MESSAGE_CODE = "definition_validation_failure"
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
|
||||
# job-result tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_DELETED_REPLACE_KEYWORDS = (
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"version",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"replace",
|
||||
"expected_current_function_id",
|
||||
"alias",
|
||||
"lineage",
|
||||
)
|
||||
|
||||
_SPEC_KEYS = {
|
||||
"format_version",
|
||||
"name",
|
||||
"definition",
|
||||
"expected_current_function_id",
|
||||
}
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"text": pa.string(), "limit": pa.int32()},
|
||||
output=pa.string(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=True,
|
||||
capabilities=[
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
],
|
||||
)
|
||||
def packable_replace_normalize(text, limit):
|
||||
"""replace-source-marker-unique-xyz."""
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def _current_function_wire() -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": _CURRENT_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _definition_json(fn: object) -> dict[str, Any]:
|
||||
payload = _udf_mod._build_function_definition(fn)._to_json()
|
||||
if isinstance(payload, bytes):
|
||||
return json.loads(payload.decode("utf-8"))
|
||||
assert isinstance(payload, str)
|
||||
return json.loads(payload)
|
||||
|
||||
|
||||
def _expected_replace_spec(name: str, current_id: str, fn: object) -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"name": name,
|
||||
"definition": _definition_json(fn),
|
||||
"expected_current_function_id": current_id,
|
||||
}
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _assert_exact_replace_spec(
|
||||
body: dict[str, Any], expected: dict[str, Any], current_id: str
|
||||
) -> None:
|
||||
assert set(body) == _SPEC_KEYS
|
||||
assert body == expected
|
||||
assert body["format_version"] == 1
|
||||
assert body["expected_current_function_id"] == current_id
|
||||
assert body["expected_current_function_id"] is not None
|
||||
assert _SOURCE_MARKER in json.dumps(body["definition"])
|
||||
assert any(
|
||||
capability.get("reference") == _SECRET_REFERENCE
|
||||
for capability in body["definition"]["capabilities"]
|
||||
)
|
||||
|
||||
|
||||
def _lookup_success_handler(
|
||||
counters: dict[str, int],
|
||||
*,
|
||||
after_lookup: (
|
||||
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
|
||||
) = None,
|
||||
):
|
||||
"""Serve exact lookup; optionally continue for register/describe."""
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
counters["lookup"] = counters.get("lookup", 0) + 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body == {"name": _FUNCTION_NAME}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps({"function": _current_function_wire()}).encode("utf-8")
|
||||
)
|
||||
return
|
||||
if after_lookup is not None:
|
||||
after_lookup(request, raw)
|
||||
return
|
||||
counters["register"] = counters.get("register", 0) + 1
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected register")
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _observe_current(db) -> lancedb.Function:
|
||||
current = db.functions.get(_FUNCTION_NAME)
|
||||
assert type(current) is lancedb.Function
|
||||
assert current.id == _CURRENT_FUNCTION_ID
|
||||
assert not hasattr(current, "name")
|
||||
assert not hasattr(current, "replace")
|
||||
return current
|
||||
|
||||
|
||||
async def _observe_current_async(db) -> lancedb.Function:
|
||||
current = await db.functions.get(_FUNCTION_NAME)
|
||||
assert type(current) is lancedb.Function
|
||||
assert current.id == _CURRENT_FUNCTION_ID
|
||||
assert not hasattr(current, "name")
|
||||
assert not hasattr(current, "replace")
|
||||
return current
|
||||
|
||||
|
||||
def test_sync_remote_replace_exact_body_one_package_job_and_function_result():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
|
||||
expected_spec = _expected_replace_spec(
|
||||
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
|
||||
)
|
||||
register_attempts: list[dict[str, Any]] = []
|
||||
function_result_wire = {
|
||||
"kind": "function",
|
||||
"format_version": 1,
|
||||
"function": {
|
||||
"format_version": 1,
|
||||
"id": _REPLACED_FUNCTION_ID,
|
||||
"signature": expected_spec["definition"]["signature"],
|
||||
},
|
||||
}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
if request.path == _REGISTER_PATH:
|
||||
counters["register"] += 1
|
||||
register_attempts.append(
|
||||
{
|
||||
"request_id": request.headers.get("x-request-id"),
|
||||
"raw": payload,
|
||||
"body": json.loads(payload.decode("utf-8")),
|
||||
}
|
||||
)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps({"job_id": _JOB_ID_SYNC}).encode("utf-8"))
|
||||
return
|
||||
|
||||
assert request.path == _DESCRIBE_PATH
|
||||
counters["describe"] += 1
|
||||
body = json.loads(payload.decode("utf-8"))
|
||||
assert body["job_id"] == _JOB_ID_SYNC
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
"job_id": _JOB_ID_SYNC,
|
||||
"job_state": "DONE",
|
||||
"job_type": "register_function",
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
"result": function_result_wire,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
package_calls = {"n": 0}
|
||||
original_package = _udf_mod._package_udf
|
||||
|
||||
def counting_package(fn: object):
|
||||
package_calls["n"] += 1
|
||||
return original_package(fn)
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "replace_function")
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
|
||||
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
|
||||
job = db.functions.replace(
|
||||
_FUNCTION_NAME, current, packable_replace_normalize
|
||||
)
|
||||
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_SYNC
|
||||
waited = job.wait(timeout=timedelta(seconds=5))
|
||||
|
||||
assert package_calls["n"] == 1
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 1
|
||||
assert counters["describe"] == 1
|
||||
assert len(register_attempts) == 1
|
||||
attempt = register_attempts[0]
|
||||
assert isinstance(attempt["request_id"], str) and attempt["request_id"]
|
||||
assert attempt["raw"]
|
||||
_assert_exact_replace_spec(attempt["body"], expected_spec, current.id)
|
||||
assert type(waited) is lancedb.Function
|
||||
assert waited.id == _REPLACED_FUNCTION_ID
|
||||
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
|
||||
assert waited.output_type == pa.string()
|
||||
assert waited.output_nullable is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_replace_exact_body_returns_async_job():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
expected_spec = _expected_replace_spec(
|
||||
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
|
||||
)
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
|
||||
assert request.path == _REGISTER_PATH
|
||||
counters["register"] += 1
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
seen["request_id"] = request.headers.get("x-request-id")
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
|
||||
|
||||
async with _mock_remote_db_async(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "replace_function")
|
||||
current = await _observe_current_async(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
job = await db.functions.replace(
|
||||
_FUNCTION_NAME, current, packable_replace_normalize
|
||||
)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 1
|
||||
assert seen.get("raw")
|
||||
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
|
||||
_assert_exact_replace_spec(seen["body"], expected_spec, current.id)
|
||||
assert type(job) is lancedb.job.AsyncJob
|
||||
assert job.id == _JOB_ID_ASYNC
|
||||
|
||||
|
||||
def test_sync_remote_replace_failed_name_conflict_raises_job_failed_error_code():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
|
||||
|
||||
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
|
||||
if request.path == _REGISTER_PATH:
|
||||
counters["register"] += 1
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps({"job_id": _JOB_ID_CONFLICT}).encode("utf-8")
|
||||
)
|
||||
return
|
||||
|
||||
assert request.path == _DESCRIBE_PATH
|
||||
counters["describe"] += 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body["job_id"] == _JOB_ID_CONFLICT
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
"job_id": _JOB_ID_CONFLICT,
|
||||
"job_type": "register_function",
|
||||
"job_state": "FAILED",
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
"failure": {
|
||||
"phase": "validate",
|
||||
"message": (
|
||||
f"looks like {_CONFLICTING_MESSAGE_CODE} during CAS"
|
||||
),
|
||||
"retryable": False,
|
||||
"error_code": "name_conflict",
|
||||
},
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
job = db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
|
||||
assert type(job) is lancedb.job.Job
|
||||
with pytest.raises(JobFailedError) as exc_info:
|
||||
job.wait(timeout=timedelta(seconds=5))
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 1
|
||||
assert counters["describe"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, JobFailedError)
|
||||
assert err.error_code == "name_conflict"
|
||||
assert err.error_code != _CONFLICTING_MESSAGE_CODE
|
||||
|
||||
|
||||
def test_empty_name_rejects_before_register_transport():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.replace("", current, packable_replace_normalize)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_current",
|
||||
[
|
||||
_CURRENT_FUNCTION_ID,
|
||||
{"id": _CURRENT_FUNCTION_ID},
|
||||
object(),
|
||||
123,
|
||||
],
|
||||
)
|
||||
def test_raw_id_or_arbitrary_current_rejected_without_register(bad_current):
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
# Observe a real handle separately so the bad-current path is isolated.
|
||||
_ = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.replace(
|
||||
_FUNCTION_NAME, bad_current, packable_replace_normalize
|
||||
)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
|
||||
|
||||
def test_local_sync_replace_not_implemented_without_table_mutation(tmp_path):
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "replace_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_replace_not_implemented_without_table_mutation(tmp_path):
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "replace_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
|
||||
|
||||
@pytest.mark.parametrize("keyword", _DELETED_REPLACE_KEYWORDS)
|
||||
def test_replace_rejects_deleted_cas_retry_version_kwargs(keyword):
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.replace(
|
||||
_FUNCTION_NAME,
|
||||
current,
|
||||
packable_replace_normalize,
|
||||
**{keyword: True},
|
||||
)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
|
||||
|
||||
def test_no_direct_replace_function_methods_and_function_has_no_replace():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert not hasattr(db, "replace_function")
|
||||
assert not hasattr(db, "register_function")
|
||||
assert not hasattr(current, "replace")
|
||||
assert not hasattr(current, "replace_function")
|
||||
assert callable(getattr(db.functions, "replace", None))
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
@@ -0,0 +1,728 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python exact first-class Function revocation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.remote.errors import HttpError
|
||||
|
||||
_REVOKE_PATH = "/v1/functions/revoke"
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_REVOKE_CATALOG_NAME = "text.normalize.revoke-name"
|
||||
_REVOKE_FUNCTION_ID = "fn.exact.revoke-handle"
|
||||
_REVOKE_SERVER_MESSAGE_MARKER = (
|
||||
"SERVER_REVOKE_DIAGNOSTIC_MARKER id=fn.exact.revoke-handle "
|
||||
"name=text.normalize.revoke-name"
|
||||
)
|
||||
_SENSITIVE_BODY_MARKER = "SENSITIVE_REVOKE_BODY_MARKER"
|
||||
_CONFLICTING_MESSAGE_CODE = "revoked_function"
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
|
||||
# remove tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_DELETED_REVOKE_KEYWORDS = (
|
||||
"function_id",
|
||||
"name",
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"version",
|
||||
"reason",
|
||||
"expiry",
|
||||
"force",
|
||||
"if_exists",
|
||||
"remove",
|
||||
"delete",
|
||||
)
|
||||
|
||||
|
||||
def _function_error_cls(*, required: bool = True):
|
||||
"""Resolve FunctionError from the live module (records RED when absent)."""
|
||||
from lancedb import exceptions as exc_mod
|
||||
|
||||
cls = getattr(exc_mod, "FunctionError", None)
|
||||
if cls is None:
|
||||
if required:
|
||||
pytest.fail("lancedb.exceptions.FunctionError is missing")
|
||||
return type("MissingFunctionError", (), {})
|
||||
return cls
|
||||
|
||||
|
||||
def _sample_function_wire() -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": _REVOKE_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _lookup_success_body() -> bytes:
|
||||
return json.dumps({"function": _sample_function_wire()}).encode("utf-8")
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
def _close_db(db: Any) -> None:
|
||||
with contextlib.suppress(Exception):
|
||||
inner = getattr(db, "_conn", None)
|
||||
if inner is not None:
|
||||
inner.close()
|
||||
return
|
||||
close = getattr(db, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
db = None
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
if db is not None:
|
||||
_close_db(db)
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
db = None
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
if db is not None:
|
||||
_close_db(db)
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_chain_text(exc: BaseException) -> str:
|
||||
parts: list[str] = []
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(str(current))
|
||||
parts.append(repr(current))
|
||||
current = current.__cause__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _assert_payload_free(exc: BaseException) -> None:
|
||||
text = _exception_chain_text(exc)
|
||||
assert _REVOKE_SERVER_MESSAGE_MARKER not in text
|
||||
assert _REVOKE_CATALOG_NAME not in text
|
||||
assert _REVOKE_FUNCTION_ID not in text
|
||||
assert _SENSITIVE_BODY_MARKER not in text
|
||||
|
||||
|
||||
def _assert_exact_revoke_function(function: object) -> None:
|
||||
assert type(function) is lancedb.Function
|
||||
assert function.id == _REVOKE_FUNCTION_ID
|
||||
assert not hasattr(function, "name")
|
||||
assert function.parameters == (
|
||||
("text", pa.string()),
|
||||
("limit", pa.int32()),
|
||||
)
|
||||
assert function.output_type == pa.string()
|
||||
assert function.output_nullable is True
|
||||
assert _REVOKE_CATALOG_NAME not in repr(function)
|
||||
assert _REVOKE_CATALOG_NAME not in str(function)
|
||||
|
||||
|
||||
def _assert_exact_revoke_request(
|
||||
request: http.server.BaseHTTPRequestHandler,
|
||||
raw: bytes,
|
||||
body: dict[str, Any],
|
||||
*,
|
||||
expected_id: str,
|
||||
) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _REVOKE_PATH
|
||||
assert "?" not in request.path
|
||||
assert "remove" not in request.path
|
||||
assert raw
|
||||
assert body == {"function_id": expected_id}
|
||||
assert set(body) == {"function_id"}
|
||||
assert "name" not in body
|
||||
assert "expected_current_function_id" not in body
|
||||
assert "format_version" not in body
|
||||
assert "function" not in body
|
||||
assert "signature" not in body
|
||||
assert "job_id" not in body
|
||||
assert "idempotency_key" not in body
|
||||
assert "user_version" not in body
|
||||
assert "reason" not in body
|
||||
assert "expiry" not in body
|
||||
assert "force" not in body
|
||||
request_id = request.headers.get("x-request-id")
|
||||
assert isinstance(request_id, str) and request_id
|
||||
|
||||
|
||||
def _assert_native_revoke_method_present() -> None:
|
||||
assert hasattr(_native.Connection, "_revoke_function")
|
||||
assert callable(getattr(_native.Connection, "_revoke_function"))
|
||||
|
||||
|
||||
def _lookup_success_handler(
|
||||
counters: dict[str, int],
|
||||
*,
|
||||
after_lookup: (
|
||||
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
|
||||
) = None,
|
||||
):
|
||||
"""Serve exact name lookup; optionally continue for revoke."""
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
counters["lookup"] = counters.get("lookup", 0) + 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body == {"name": _REVOKE_CATALOG_NAME}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
if after_lookup is not None:
|
||||
after_lookup(request, raw)
|
||||
return
|
||||
counters["revoke"] = counters.get("revoke", 0) + 1
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected revoke")
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _observe_current(db) -> lancedb.Function:
|
||||
current = db.functions.get(_REVOKE_CATALOG_NAME)
|
||||
_assert_exact_revoke_function(current)
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
return current
|
||||
|
||||
|
||||
async def _observe_current_async(db) -> lancedb.Function:
|
||||
current = await db.functions.get(_REVOKE_CATALOG_NAME)
|
||||
_assert_exact_revoke_function(current)
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
return current
|
||||
|
||||
|
||||
def test_native_connection_exposes_private_revoke_function():
|
||||
_assert_native_revoke_method_present()
|
||||
|
||||
|
||||
def test_sync_remote_revoke_exact_body_path_request_id_returns_none():
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
revoke_attempts: list[dict[str, Any]] = []
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
body = json.loads(payload.decode("utf-8"))
|
||||
revoke_attempts.append(
|
||||
{
|
||||
"request": request,
|
||||
"raw": payload,
|
||||
"body": body,
|
||||
"request_id": request.headers.get("x-request-id"),
|
||||
}
|
||||
)
|
||||
# Illegal body on 204 must be ignored; success is status-driven only.
|
||||
request.send_response(204)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
"message": _REVOKE_SERVER_MESSAGE_MARKER,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "revoke_function")
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
|
||||
result = db.functions.revoke(current)
|
||||
|
||||
assert result is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
assert len(revoke_attempts) == 1
|
||||
attempt = revoke_attempts[0]
|
||||
_assert_exact_revoke_request(
|
||||
attempt["request"],
|
||||
attempt["raw"],
|
||||
attempt["body"],
|
||||
expected_id=current.id,
|
||||
)
|
||||
assert attempt["body"]["function_id"] == current.id
|
||||
_assert_exact_revoke_function(current)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_revoke_exact_body_returns_none():
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
seen["request"] = request
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
seen["request_id"] = request.headers.get("x-request-id")
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
async with _mock_remote_db_async(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "revoke_function")
|
||||
current = await _observe_current_async(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
result = await db.functions.revoke(current)
|
||||
|
||||
assert result is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
assert seen.get("raw")
|
||||
_assert_exact_revoke_request(
|
||||
seen["request"],
|
||||
seen["raw"],
|
||||
seen["body"],
|
||||
expected_id=current.id,
|
||||
)
|
||||
|
||||
|
||||
def test_repeated_remote_revoke_204_both_return_none():
|
||||
"""Two logical calls each receiving 204 both succeed (Python outcome only)."""
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
revoke_request_ids: list[str] = []
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
body = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_revoke_request(
|
||||
request, payload, body, expected_id=_REVOKE_FUNCTION_ID
|
||||
)
|
||||
request_id = request.headers.get("x-request-id")
|
||||
assert isinstance(request_id, str) and request_id
|
||||
revoke_request_ids.append(request_id)
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
|
||||
first = db.functions.revoke(current)
|
||||
second = db.functions.revoke(current)
|
||||
|
||||
assert first is None
|
||||
assert second is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 2
|
||||
assert len(revoke_request_ids) == 2
|
||||
_assert_exact_revoke_function(current)
|
||||
|
||||
|
||||
def test_after_revoke_name_and_id_lookup_still_return_function():
|
||||
"""Revoke does not unlink names; SDK-visible sequence only, not Sophon proof."""
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {
|
||||
"lookup_name": 0,
|
||||
"lookup_id": 0,
|
||||
"revoke": 0,
|
||||
}
|
||||
revoked = {"yes": False}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
if "name" in body:
|
||||
counters["lookup_name"] += 1
|
||||
assert body == {"name": _REVOKE_CATALOG_NAME}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
|
||||
counters["lookup_id"] += 1
|
||||
assert body == {"function_id": _REVOKE_FUNCTION_ID}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
_assert_exact_revoke_request(
|
||||
request, raw, body, expected_id=_REVOKE_FUNCTION_ID
|
||||
)
|
||||
revoked["yes"] = True
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup_name"] == 1
|
||||
assert counters["lookup_id"] == 0
|
||||
assert counters["revoke"] == 0
|
||||
assert not revoked["yes"]
|
||||
|
||||
result = db.functions.revoke(current)
|
||||
assert result is None
|
||||
assert counters["lookup_name"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
assert revoked["yes"]
|
||||
|
||||
by_name = db.functions.get(_REVOKE_CATALOG_NAME)
|
||||
by_id = db.functions.get_by_id(_REVOKE_FUNCTION_ID)
|
||||
|
||||
assert counters["lookup_name"] == 2
|
||||
assert counters["lookup_id"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
_assert_exact_revoke_function(by_name)
|
||||
_assert_exact_revoke_function(by_id)
|
||||
assert by_name.id == current.id
|
||||
assert by_id.id == current.id
|
||||
assert by_name.parameters == current.parameters
|
||||
assert by_id.parameters == current.parameters
|
||||
assert by_name.output_type == current.output_type
|
||||
assert by_id.output_type == current.output_type
|
||||
assert by_name.output_nullable is current.output_nullable
|
||||
assert by_id.output_nullable is current.output_nullable
|
||||
|
||||
|
||||
def test_explicit_name_or_function_not_found_is_function_error_payload_free():
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
body = {
|
||||
"error_code": "name_or_function_not_found",
|
||||
"message": (
|
||||
f"{_REVOKE_SERVER_MESSAGE_MARKER} looks_like {_CONFLICTING_MESSAGE_CODE} "
|
||||
"name_conflict"
|
||||
),
|
||||
"name": _REVOKE_CATALOG_NAME,
|
||||
"function_id": _REVOKE_FUNCTION_ID,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
parsed = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_revoke_request(
|
||||
request, payload, parsed, expected_id=_REVOKE_FUNCTION_ID
|
||||
)
|
||||
request.send_response(404)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.revoke(current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, function_error)
|
||||
assert err.code == "name_or_function_not_found"
|
||||
assert err.code != _CONFLICTING_MESSAGE_CODE
|
||||
assert err.code != "name_conflict"
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,status,response_body",
|
||||
[
|
||||
(
|
||||
"200_with_body",
|
||||
200,
|
||||
{
|
||||
"ok": True,
|
||||
"message": _REVOKE_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
"job_id": "must-not-infer-job",
|
||||
},
|
||||
),
|
||||
(
|
||||
"202_empty",
|
||||
202,
|
||||
f"{_REVOKE_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
|
||||
),
|
||||
("200_empty", 200, ""),
|
||||
],
|
||||
)
|
||||
def test_http_200_202_cannot_return_success(
|
||||
label: str, status: int, response_body: object
|
||||
):
|
||||
del label
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
parsed = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_revoke_request(
|
||||
request, payload, parsed, expected_id=_REVOKE_FUNCTION_ID
|
||||
)
|
||||
request.send_response(status)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
if isinstance(response_body, str):
|
||||
request.wfile.write(response_body.encode("utf-8"))
|
||||
else:
|
||||
request.wfile.write(json.dumps(response_body).encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.revoke(current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert not isinstance(err, _function_error_cls(required=False))
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_function",
|
||||
[
|
||||
_REVOKE_FUNCTION_ID,
|
||||
{"id": _REVOKE_FUNCTION_ID},
|
||||
object(),
|
||||
123,
|
||||
],
|
||||
)
|
||||
def test_raw_id_or_arbitrary_function_rejected_without_revoke(bad_function):
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
# Observe a real handle separately so the bad-function path is isolated.
|
||||
_ = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.revoke(bad_function)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
|
||||
|
||||
def test_local_sync_revoke_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "revoke_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.revoke(current)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
_close_db(db)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_revoke_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "revoke_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.revoke(current)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("keyword", _DELETED_REVOKE_KEYWORDS)
|
||||
def test_revoke_rejects_overdesigned_kwargs(keyword):
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.revoke(current, **{keyword: True})
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
|
||||
|
||||
def test_no_direct_revoke_methods_and_function_has_no_revoke_facade_private():
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert not hasattr(db, "revoke_function")
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
assert callable(getattr(db.functions, "revoke", None))
|
||||
assert not hasattr(lancedb, "_SyncFunctions")
|
||||
assert not hasattr(lancedb, "_AsyncFunctions")
|
||||
assert type(db.functions).__name__.startswith("_")
|
||||
assert "_SyncFunctions" not in getattr(lancedb, "__all__", [])
|
||||
assert "_AsyncFunctions" not in getattr(lancedb, "__all__", [])
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,899 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python ``table.add_generated_column`` (FF-032).
|
||||
|
||||
Public user shape under test:
|
||||
|
||||
job = table.add_generated_column(
|
||||
"normalized_text",
|
||||
normalize(text=col("text")),
|
||||
)
|
||||
job.wait()
|
||||
|
||||
These tests exercise the live worktree PyO3 extension and public sync/async
|
||||
wrappers. While the public methods and hidden native bridge are absent they
|
||||
fail against that extension; once present they freeze the public contract
|
||||
below. They must not fake success paths.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import inspect
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
import lancedb.job
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.expr import col
|
||||
from lancedb.remote.table import RemoteTable
|
||||
from lancedb.table import AsyncTable, LanceTable, Table
|
||||
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_JOB_DESCRIBE_PATH = "/v1/jobs/describe"
|
||||
_TABLE_NAME = "articles"
|
||||
_DESCRIBE_PATH = f"/v1/table/{_TABLE_NAME}/describe/"
|
||||
_CREATE_PATH = f"/v1/table/{_TABLE_NAME}/generated_columns/create/"
|
||||
_BRANCHES_CREATE_PATH = f"/v1/table/{_TABLE_NAME}/branches/create/"
|
||||
_BRANCHES_LIST_PATH = f"/v1/table/{_TABLE_NAME}/branches/list/"
|
||||
|
||||
_CATALOG_NAME = "text.normalize"
|
||||
_FUNCTION_ID = "fn.exact.normalize.gen-col"
|
||||
_JOB_ID_SYNC = "job-create-gen-col-sync-1"
|
||||
_JOB_ID_ASYNC = "job-create-gen-col-async-1"
|
||||
_JOB_ID_BRANCH = "job-create-gen-col-branch-1"
|
||||
_SOURCE_TABLE_VERSION = 42
|
||||
_TEXT_FIELD_ID = 7
|
||||
_BRANCH_NAME = "exp"
|
||||
_BRANCH_SOURCE_VERSION = 9
|
||||
_BRANCH_TEXT_FIELD_ID = 11
|
||||
|
||||
_DESCRIBE_BODY_MARKER = "SENSITIVE_DESCRIBE_BODY_MARKER_gen_col_xyz"
|
||||
_CREATE_RESPONSE_MARKER = "SENSITIVE_CREATE_RESPONSE_MARKER_gen_col_xyz"
|
||||
_LITERAL_PAYLOAD_SENTINEL = "LITERAL_PAYLOAD_SENTINEL_gen_col_xyz"
|
||||
|
||||
# Pinned Rust-canonical schema-only Utf8 type IPC (base64), shared with FF-028.
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_FORBIDDEN_PUBLIC_NAMES = (
|
||||
"FunctionCall",
|
||||
"BoundFunctionCall",
|
||||
"AuthoredFunctionCall",
|
||||
"CreateGeneratedColumnRequest",
|
||||
"CreateGeneratedColumnJobSpec",
|
||||
"GeneratedColumnBindingSnapshot",
|
||||
"GeneratedColumnCreateRequest",
|
||||
"geneva",
|
||||
"GenevaFunction",
|
||||
"VirtualColumnDefinition",
|
||||
)
|
||||
|
||||
_FORBIDDEN_METHOD_KWARGS = (
|
||||
"source_table_version",
|
||||
"version",
|
||||
"field_id",
|
||||
"field_ids",
|
||||
"output",
|
||||
"output_type",
|
||||
"output_nullable",
|
||||
"nullable",
|
||||
"spec",
|
||||
"retry_key",
|
||||
"idempotency_key",
|
||||
"request",
|
||||
"envelope",
|
||||
"table_ref",
|
||||
"branch",
|
||||
)
|
||||
|
||||
|
||||
def _sample_function_wire(
|
||||
*,
|
||||
function_id: str = _FUNCTION_ID,
|
||||
parameters: list[dict[str, str]] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": function_id,
|
||||
"signature": {
|
||||
"parameters": parameters
|
||||
or [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _text_schema_fields(
|
||||
*, arrow_type: str = "string", nullable: bool = True
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"fields": [
|
||||
{
|
||||
"name": "text",
|
||||
"type": {"type": arrow_type},
|
||||
"nullable": nullable,
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def _describe_body(
|
||||
*,
|
||||
version: int = _SOURCE_TABLE_VERSION,
|
||||
field_ids: list[int] | None = None,
|
||||
arrow_type: str = "string",
|
||||
include_marker: bool = True,
|
||||
) -> dict[str, Any]:
|
||||
body: dict[str, Any] = {
|
||||
"version": version,
|
||||
"schema": _text_schema_fields(arrow_type=arrow_type),
|
||||
"field_ids": field_ids if field_ids is not None else [_TEXT_FIELD_ID],
|
||||
}
|
||||
if include_marker:
|
||||
body["server_diagnostic"] = _DESCRIBE_BODY_MARKER
|
||||
return body
|
||||
|
||||
|
||||
def _create_gen_column_done_body(job_id: str) -> dict[str, Any]:
|
||||
# DONE with omitted result: create_gen_column projects JobResult::None.
|
||||
return {
|
||||
"job_id": job_id,
|
||||
"job_state": "DONE",
|
||||
"job_type": "create_gen_column",
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
}
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_text(exc: BaseException) -> str:
|
||||
parts = [str(exc), repr(exc)]
|
||||
current: BaseException | None = exc
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(f"{type(current).__name__}: {current}")
|
||||
current = current.__cause__ or current.__context__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _json_response(
|
||||
request: http.server.BaseHTTPRequestHandler, body: dict[str, Any]
|
||||
) -> None:
|
||||
payload = json.dumps(body).encode("utf-8")
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(payload)
|
||||
|
||||
|
||||
def _lookup_function(db: Any) -> lancedb.Function:
|
||||
return db.functions.get(_CATALOG_NAME)
|
||||
|
||||
|
||||
class _RequestLog:
|
||||
"""Track lookup/describe/create after setup; setup traffic is excluded."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.lookup: list[dict[str, Any]] = []
|
||||
self.describe: list[dict[str, Any]] = []
|
||||
self.create: list[dict[str, Any]] = []
|
||||
self.other_table: list[str] = []
|
||||
self.recording = False
|
||||
|
||||
def start(self) -> None:
|
||||
# Drop setup's explicit Function lookup and open_table describe so
|
||||
# operation accounting cannot be polluted by fixture traffic.
|
||||
self.lookup.clear()
|
||||
self.describe.clear()
|
||||
self.create.clear()
|
||||
self.other_table.clear()
|
||||
self.recording = True
|
||||
|
||||
def note(self, path: str, body: dict[str, Any] | None = None) -> None:
|
||||
if not self.recording:
|
||||
return
|
||||
if path == _LOOKUP_PATH:
|
||||
self.lookup.append(body or {})
|
||||
elif path == _DESCRIBE_PATH:
|
||||
self.describe.append(body or {})
|
||||
elif path == _CREATE_PATH:
|
||||
self.create.append(body or {})
|
||||
elif path.startswith(f"/v1/table/{_TABLE_NAME}/"):
|
||||
self.other_table.append(path)
|
||||
|
||||
|
||||
def _assert_no_operation_traffic(log: _RequestLog) -> None:
|
||||
assert log.lookup == []
|
||||
assert log.describe == []
|
||||
assert log.create == []
|
||||
assert log.other_table == []
|
||||
|
||||
|
||||
def _assert_exact_public_signature(method: Any) -> None:
|
||||
"""Freeze ``(self, column_name, call)`` with no varargs/kwargs escape hatches."""
|
||||
params = list(inspect.signature(method).parameters.values())
|
||||
assert [p.name for p in params] == ["self", "column_name", "call"]
|
||||
for param in params:
|
||||
assert param.kind in (
|
||||
inspect.Parameter.POSITIONAL_ONLY,
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
)
|
||||
assert param.default is inspect.Parameter.empty
|
||||
assert param.kind is not inspect.Parameter.VAR_POSITIONAL
|
||||
assert param.kind is not inspect.Parameter.VAR_KEYWORD
|
||||
assert param.kind is not inspect.Parameter.KEYWORD_ONLY
|
||||
|
||||
|
||||
def _open_table_and_function(
|
||||
*,
|
||||
describe_body: dict[str, Any] | None = None,
|
||||
on_create: Callable[[dict[str, Any], http.server.BaseHTTPRequestHandler], None]
|
||||
| None = None,
|
||||
job_id: str = _JOB_ID_SYNC,
|
||||
support_branch_create: bool = False,
|
||||
function_wire: dict[str, Any] | None = None,
|
||||
):
|
||||
"""Open remote table + immutable Function; return (db, table, function, log, cm)."""
|
||||
log = _RequestLog()
|
||||
binding_describe = describe_body or _describe_body()
|
||||
open_describe = {
|
||||
"version": 1,
|
||||
"schema": _text_schema_fields(),
|
||||
}
|
||||
state = {"opened": False}
|
||||
wire = function_wire or _sample_function_wire()
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _LOOKUP_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"function": wire})
|
||||
return
|
||||
|
||||
if request.path == _JOB_DESCRIBE_PATH:
|
||||
assert body["job_id"] == job_id
|
||||
_json_response(request, _create_gen_column_done_body(job_id))
|
||||
return
|
||||
|
||||
if support_branch_create and request.path == _BRANCHES_CREATE_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {})
|
||||
return
|
||||
|
||||
if support_branch_create and request.path == _BRANCHES_LIST_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(
|
||||
request,
|
||||
{
|
||||
"branches": {
|
||||
_BRANCH_NAME: {
|
||||
"parentBranch": None,
|
||||
"parentVersion": 1,
|
||||
"createAt": 1,
|
||||
"manifestSize": 1,
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
# First describe seeds open_table; later ones are binding snapshots.
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding_describe)
|
||||
return
|
||||
|
||||
if request.path == _CREATE_PATH:
|
||||
log.note(request.path, body)
|
||||
if on_create is not None:
|
||||
on_create(body, request)
|
||||
return
|
||||
_json_response(
|
||||
request,
|
||||
{
|
||||
"job_id": job_id,
|
||||
"server_extra": {"marker": _CREATE_RESPONSE_MARKER},
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
if request.path.startswith(f"/v1/table/{_TABLE_NAME}/"):
|
||||
log.note(request.path, body)
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected path")
|
||||
|
||||
cm = _mock_remote_db(handler)
|
||||
db = cm.__enter__()
|
||||
function = _lookup_function(db)
|
||||
table = db.open_table(_TABLE_NAME)
|
||||
assert isinstance(table, RemoteTable)
|
||||
# open_table consumed the seed describe; binding/create accounting starts now.
|
||||
# Setup's one explicit lookup is cleared here and must not pollute counts.
|
||||
log.start()
|
||||
return db, table, function, log, cm
|
||||
|
||||
|
||||
def _assert_exact_create_envelope(
|
||||
body: dict[str, Any],
|
||||
*,
|
||||
source_table_version: int,
|
||||
column_name: str,
|
||||
field_id: int,
|
||||
branch: str | None = None,
|
||||
) -> None:
|
||||
expected_keys = {"source_table_version", "spec"}
|
||||
if branch is not None:
|
||||
expected_keys.add("branch")
|
||||
assert set(body) == expected_keys
|
||||
assert body["source_table_version"] == source_table_version
|
||||
assert "table_ref" not in body
|
||||
if branch is None:
|
||||
assert "branch" not in body
|
||||
else:
|
||||
assert body["branch"] == branch
|
||||
|
||||
spec = body["spec"]
|
||||
assert set(spec) == {"format_version", "column_name", "function_call"}
|
||||
assert spec["format_version"] == 1
|
||||
assert spec["column_name"] == column_name
|
||||
for forbidden in (
|
||||
"table_ref",
|
||||
"source_table_version",
|
||||
"version",
|
||||
"output",
|
||||
"output_type",
|
||||
"output_field_id",
|
||||
"dependency_epoch",
|
||||
"materialized_epoch",
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"name",
|
||||
"handle",
|
||||
"artifact",
|
||||
"geneva",
|
||||
):
|
||||
assert forbidden not in spec
|
||||
|
||||
call = spec["function_call"]
|
||||
assert set(call) == {"function_id", "arguments"}
|
||||
assert call["function_id"] == _FUNCTION_ID
|
||||
assert len(call["arguments"]) == 1
|
||||
binding = call["arguments"][0]
|
||||
assert binding["parameter"] == "text"
|
||||
value = binding["value"]
|
||||
assert value["kind"] == "field"
|
||||
assert value["field_id"] == field_id
|
||||
assert value["data_type_ipc"] == _UTF8_TYPE_IPC_B64
|
||||
assert "name" not in value
|
||||
assert "column_name" not in value
|
||||
assert "text" not in value
|
||||
# Serialized call must not late-bind by column name anywhere relevant.
|
||||
dumped = json.dumps(call)
|
||||
assert '"column_name"' not in dumped
|
||||
assert "normalized_text" not in dumped
|
||||
|
||||
|
||||
def test_public_and_native_add_generated_column_seams_must_exist():
|
||||
"""Public sync/async methods and the private native bridge must exist."""
|
||||
assert hasattr(_native.Table, "_add_generated_column"), (
|
||||
"native private bridge Table._add_generated_column is missing"
|
||||
)
|
||||
assert hasattr(AsyncTable, "add_generated_column"), (
|
||||
"AsyncTable.add_generated_column is missing"
|
||||
)
|
||||
assert hasattr(Table, "add_generated_column"), (
|
||||
"Table.add_generated_column is missing"
|
||||
)
|
||||
assert hasattr(LanceTable, "add_generated_column"), (
|
||||
"LanceTable.add_generated_column is missing"
|
||||
)
|
||||
assert hasattr(RemoteTable, "add_generated_column"), (
|
||||
"RemoteTable.add_generated_column is missing"
|
||||
)
|
||||
|
||||
# Once present, freeze the exact public positional surface.
|
||||
_assert_exact_public_signature(Table.add_generated_column)
|
||||
_assert_exact_public_signature(LanceTable.add_generated_column)
|
||||
_assert_exact_public_signature(RemoteTable.add_generated_column)
|
||||
_assert_exact_public_signature(AsyncTable.add_generated_column)
|
||||
|
||||
|
||||
def test_sync_remote_add_generated_column_returns_job_without_eager_wrapper_mutation():
|
||||
db, table, normalize, log, cm = _open_table_and_function(job_id=_JOB_ID_SYNC)
|
||||
try:
|
||||
# Capture public wrapper state before the operation window.
|
||||
schema_before = table.schema
|
||||
version_before = table.version
|
||||
log.start()
|
||||
|
||||
call = normalize(text=col("text"))
|
||||
# Exact public argument order from the frozen user example.
|
||||
job = table.add_generated_column(
|
||||
"normalized_text",
|
||||
call,
|
||||
)
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_SYNC
|
||||
|
||||
# Exact success path stops after submit: one binding describe, one create,
|
||||
# and no catalog re-lookup. Do not wait yet.
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.create) == 1
|
||||
assert log.other_table == []
|
||||
|
||||
# Public schema/version through the existing wrapper must still reflect
|
||||
# the pre-submit table: generated column is not published by Job accept.
|
||||
# Access both before wait so eager wrapper cache invalidation / refresh /
|
||||
# version advancement is observable.
|
||||
schema_after = table.schema
|
||||
assert "normalized_text" not in schema_after.names
|
||||
assert schema_after == schema_before
|
||||
# Schema must be served from the existing wrapper cache — no extra
|
||||
# describe beyond the one binding snapshot.
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.create) == 1
|
||||
|
||||
version_after = table.version
|
||||
assert version_after == version_before
|
||||
# Public Remote ``version`` always describes once by design; that probe
|
||||
# must not drag a schema-cache miss, create, or catalog lookup with it.
|
||||
assert len(log.describe) == 2
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.create) == 1
|
||||
assert log.other_table == []
|
||||
|
||||
waited = job.wait()
|
||||
assert waited is None
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.create) == 1
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_add_generated_column_returns_async_job_and_wait_none():
|
||||
log = _RequestLog()
|
||||
state = {"opened": False}
|
||||
binding_describe = _describe_body()
|
||||
open_describe = {"version": 1, "schema": _text_schema_fields()}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _LOOKUP_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"function": _sample_function_wire()})
|
||||
return
|
||||
if request.path == _JOB_DESCRIBE_PATH:
|
||||
assert body["job_id"] == _JOB_ID_ASYNC
|
||||
_json_response(request, _create_gen_column_done_body(_JOB_ID_ASYNC))
|
||||
return
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding_describe)
|
||||
return
|
||||
if request.path == _CREATE_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"job_id": _JOB_ID_ASYNC})
|
||||
return
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
|
||||
async with _mock_remote_db_async(handler) as db:
|
||||
normalize = await db.functions.get(_CATALOG_NAME)
|
||||
table = await db.open_table(_TABLE_NAME)
|
||||
log.start()
|
||||
call = normalize(text=col("text"))
|
||||
job = await table.add_generated_column("normalized_text", call)
|
||||
assert type(job) is lancedb.job.AsyncJob
|
||||
assert job.id == _JOB_ID_ASYNC
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.create) == 1
|
||||
waited = await job.wait()
|
||||
assert waited is None
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.create) == 1
|
||||
|
||||
|
||||
def test_remote_add_generated_column_one_describe_one_create_exact_envelope():
|
||||
db, table, normalize, log, cm = _open_table_and_function(job_id=_JOB_ID_SYNC)
|
||||
try:
|
||||
call = normalize(text=col("text"))
|
||||
job = table.add_generated_column("normalized_text", call)
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_SYNC
|
||||
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.create) == 1
|
||||
assert log.other_table == []
|
||||
_assert_exact_create_envelope(
|
||||
log.create[0],
|
||||
source_table_version=_SOURCE_TABLE_VERSION,
|
||||
column_name="normalized_text",
|
||||
field_id=_TEXT_FIELD_ID,
|
||||
)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_remote_branch_add_generated_column_includes_exact_branch_identity():
|
||||
branch_describe = _describe_body(
|
||||
version=_BRANCH_SOURCE_VERSION,
|
||||
field_ids=[_BRANCH_TEXT_FIELD_ID],
|
||||
)
|
||||
db, table, normalize, log, cm = _open_table_and_function(
|
||||
describe_body=branch_describe,
|
||||
job_id=_JOB_ID_BRANCH,
|
||||
support_branch_create=True,
|
||||
)
|
||||
try:
|
||||
branched = table.branches.create(_BRANCH_NAME)
|
||||
assert isinstance(branched, RemoteTable)
|
||||
assert branched.current_branch() == _BRANCH_NAME
|
||||
log.start()
|
||||
|
||||
call = normalize(text=col("text"))
|
||||
job = branched.add_generated_column("normalized_text", call)
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_BRANCH
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.create) == 1
|
||||
assert log.describe[0].get("branch") == _BRANCH_NAME
|
||||
_assert_exact_create_envelope(
|
||||
log.create[0],
|
||||
source_table_version=_BRANCH_SOURCE_VERSION,
|
||||
column_name="normalized_text",
|
||||
field_id=_BRANCH_TEXT_FIELD_ID,
|
||||
branch=_BRANCH_NAME,
|
||||
)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_empty_column_name_fails_locally_with_zero_table_requests():
|
||||
db, table, normalize, log, cm = _open_table_and_function()
|
||||
try:
|
||||
# Authored call owns a real literal so payload-free failure is not vacuous.
|
||||
call = normalize(text=_LITERAL_PAYLOAD_SENTINEL)
|
||||
with pytest.raises((ValueError, TypeError)) as raised:
|
||||
table.add_generated_column("", call)
|
||||
text = _exception_text(raised.value)
|
||||
lowered = text.lower()
|
||||
assert "column" in lowered or "empty" in lowered or "non-empty" in lowered
|
||||
assert _DESCRIBE_BODY_MARKER not in text
|
||||
assert _CREATE_RESPONSE_MARKER not in text
|
||||
assert _LITERAL_PAYLOAD_SENTINEL not in text
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("column_ref", "expected_token"),
|
||||
[
|
||||
("missing_text", "missing_text"),
|
||||
("Text", "Text"), # exact-case mismatch against schema field "text"
|
||||
],
|
||||
)
|
||||
def test_missing_or_case_mismatch_column_one_describe_zero_create(
|
||||
column_ref: str, expected_token: str
|
||||
):
|
||||
db, table, normalize, log, cm = _open_table_and_function()
|
||||
try:
|
||||
call = normalize(text=col(column_ref))
|
||||
with pytest.raises(ValueError) as raised:
|
||||
table.add_generated_column("normalized_text", call)
|
||||
text = _exception_text(raised.value)
|
||||
assert expected_token in text
|
||||
assert "text" in text # parameter name from the Function signature
|
||||
assert "missing" in text.lower() or "field" in text.lower()
|
||||
assert _DESCRIBE_BODY_MARKER not in text
|
||||
assert _CREATE_RESPONSE_MARKER not in text
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert log.create == []
|
||||
assert log.other_table == []
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_type_mismatch_one_describe_zero_create_identifies_parameter():
|
||||
db, table, normalize, log, cm = _open_table_and_function(
|
||||
describe_body=_describe_body(arrow_type="int32"),
|
||||
)
|
||||
try:
|
||||
call = normalize(text=col("text"))
|
||||
with pytest.raises(ValueError) as raised:
|
||||
table.add_generated_column("normalized_text", call)
|
||||
text = _exception_text(raised.value)
|
||||
assert "text" in text
|
||||
assert "type" in text.lower() or "mismatch" in text.lower()
|
||||
assert _DESCRIBE_BODY_MARKER not in text
|
||||
assert _CREATE_RESPONSE_MARKER not in text
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert log.create == []
|
||||
assert log.other_table == []
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_literal_payload_stays_out_of_field_binding_failure():
|
||||
"""Authored call owns a real literal; later field binding fails payload-free."""
|
||||
wire = _sample_function_wire(
|
||||
parameters=[
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "prefix", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
]
|
||||
)
|
||||
db, table, normalize, log, cm = _open_table_and_function(function_wire=wire)
|
||||
try:
|
||||
call = normalize(text=col("missing_text"), prefix=_LITERAL_PAYLOAD_SENTINEL)
|
||||
with pytest.raises(ValueError) as raised:
|
||||
table.add_generated_column("normalized_text", call)
|
||||
text = _exception_text(raised.value)
|
||||
assert "missing_text" in text
|
||||
assert _LITERAL_PAYLOAD_SENTINEL not in text
|
||||
assert _DESCRIBE_BODY_MARKER not in text
|
||||
assert _CREATE_RESPONSE_MARKER not in text
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert log.create == []
|
||||
assert log.other_table == []
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_closed_async_table_fails_with_zero_operation_requests():
|
||||
log = _RequestLog()
|
||||
state = {"opened": False}
|
||||
binding_describe = _describe_body()
|
||||
open_describe = {"version": 1, "schema": _text_schema_fields()}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _LOOKUP_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"function": _sample_function_wire()})
|
||||
return
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding_describe)
|
||||
return
|
||||
if request.path == _CREATE_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"job_id": _JOB_ID_ASYNC})
|
||||
return
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
|
||||
async with _mock_remote_db_async(handler) as db:
|
||||
normalize = await db.functions.get(_CATALOG_NAME)
|
||||
table = await db.open_table(_TABLE_NAME)
|
||||
call = normalize(text=col("text"))
|
||||
# Public close only — do not mutate private implementation fields.
|
||||
table.close()
|
||||
log.start()
|
||||
try:
|
||||
await table.add_generated_column("normalized_text", call)
|
||||
except AttributeError:
|
||||
# Method missing: re-raise so the failure names the public seam.
|
||||
raise
|
||||
except Exception as exc:
|
||||
text = _exception_text(exc)
|
||||
assert "closed" in text.lower()
|
||||
else:
|
||||
pytest.fail("closed AsyncTable must fail before transport")
|
||||
_assert_no_operation_traffic(log)
|
||||
|
||||
|
||||
def test_rejects_non_authored_call_before_any_operation_request():
|
||||
db, table, normalize, log, cm = _open_table_and_function()
|
||||
try:
|
||||
bad_values = (
|
||||
normalize, # exact Function handle itself
|
||||
{"text": "x"},
|
||||
col("text"), # direct query Expr
|
||||
object(),
|
||||
)
|
||||
for bad in bad_values:
|
||||
with pytest.raises(TypeError):
|
||||
table.add_generated_column("normalized_text", bad)
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_native_valid_call_returns_not_supported_without_mutation(tmp_path):
|
||||
# Immutable Function handle is connection-free; obtain it via remote lookup.
|
||||
def lookup_only(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
_json_response(request, {"function": _sample_function_wire()})
|
||||
|
||||
with _mock_remote_db(lookup_only) as remote_db:
|
||||
normalize = _lookup_function(remote_db)
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table(_TABLE_NAME, [{"text": "Hello"}, {"text": "World"}])
|
||||
assert isinstance(table, LanceTable)
|
||||
version_before = table.version
|
||||
schema_before = table.schema
|
||||
rows_before = table.to_arrow().to_pylist()
|
||||
call = normalize(text=col("text"))
|
||||
|
||||
with pytest.raises(NotImplementedError) as raised:
|
||||
table.add_generated_column("normalized_text", call)
|
||||
text = _exception_text(raised.value)
|
||||
assert "not supported" in text.lower() or "submit_create_generated_column" in text
|
||||
assert "add_columns" not in text.lower()
|
||||
|
||||
assert table.version == version_before
|
||||
assert table.schema == schema_before
|
||||
assert "normalized_text" not in table.schema.names
|
||||
assert table.to_arrow().to_pylist() == rows_before
|
||||
|
||||
|
||||
def test_public_surface_is_minimal_and_private_call_stays_opaque():
|
||||
for name in _FORBIDDEN_PUBLIC_NAMES:
|
||||
assert name not in getattr(lancedb, "__all__", [])
|
||||
assert not hasattr(lancedb, name)
|
||||
|
||||
assert not hasattr(lancedb, "_FunctionCall")
|
||||
authored_type = getattr(_native, "_FunctionCall", None)
|
||||
assert authored_type is not None
|
||||
with pytest.raises(TypeError):
|
||||
authored_type()
|
||||
|
||||
# When the public method exists, reject overdesign kwargs and keep the frozen
|
||||
# positional surface: (self, column_name, call).
|
||||
if hasattr(Table, "add_generated_column"):
|
||||
_assert_exact_public_signature(Table.add_generated_column)
|
||||
for keyword in _FORBIDDEN_METHOD_KWARGS:
|
||||
assert (
|
||||
keyword not in inspect.signature(Table.add_generated_column).parameters
|
||||
)
|
||||
db, table, normalize, log, cm = _open_table_and_function()
|
||||
try:
|
||||
call = normalize(text=col("text"))
|
||||
for keyword in _FORBIDDEN_METHOD_KWARGS:
|
||||
with pytest.raises(TypeError):
|
||||
table.add_generated_column(
|
||||
"normalized_text",
|
||||
call,
|
||||
**{keyword: object()},
|
||||
)
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
if hasattr(LanceTable, "add_generated_column"):
|
||||
_assert_exact_public_signature(LanceTable.add_generated_column)
|
||||
if hasattr(RemoteTable, "add_generated_column"):
|
||||
_assert_exact_public_signature(RemoteTable.add_generated_column)
|
||||
if hasattr(AsyncTable, "add_generated_column"):
|
||||
_assert_exact_public_signature(AsyncTable.add_generated_column)
|
||||
for keyword in _FORBIDDEN_METHOD_KWARGS:
|
||||
assert (
|
||||
keyword
|
||||
not in inspect.signature(AsyncTable.add_generated_column).parameters
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,672 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python ``table.generated_column_status`` (B3d2).
|
||||
|
||||
Public user shape under test:
|
||||
|
||||
status = table.generated_column_status("complete_col") # "complete" | "incomplete"
|
||||
|
||||
These tests exercise the live worktree PyO3 extension and public sync/async
|
||||
wrappers. While the public methods and hidden native bridge are absent they
|
||||
fail against that extension; once present they freeze the public contract
|
||||
below. They must not fake success paths.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import inspect
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable, Literal, get_type_hints
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
import lancedb.table
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.remote.table import RemoteTable
|
||||
from lancedb.table import AsyncTable, LanceTable, Table
|
||||
|
||||
_TABLE_NAME = "articles"
|
||||
_DESCRIBE_PATH = f"/v1/table/{_TABLE_NAME}/describe/"
|
||||
|
||||
_ORDINARY_FIELD_ID = 1
|
||||
_COMPLETE_FIELD_ID = 5
|
||||
_INCOMPLETE_FIELD_ID = 7
|
||||
_STABLE_FIELD_IDS = [_ORDINARY_FIELD_ID, _COMPLETE_FIELD_ID, _INCOMPLETE_FIELD_ID]
|
||||
|
||||
_STATUS_FUNCTION_ID = "fn.exact.status.projection"
|
||||
_METADATA_KEY = "lancedb::generated_column"
|
||||
_RAW_METADATA_MARKER = "SENSITIVE_STATUS_METADATA_MARKER_b3d2_py_9f2e"
|
||||
|
||||
# Pinned Rust-canonical schema-only Utf8 type IPC (base64), shared with FF-028.
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_EXPECTED_RETURN = Literal["complete", "incomplete"]
|
||||
|
||||
_FORBIDDEN_PUBLIC_NAMES = (
|
||||
"GeneratedColumnStatus",
|
||||
"GeneratedColumnDefinition",
|
||||
"GeneratedColumnBindingSnapshot",
|
||||
"GeneratedColumnBindingEntry",
|
||||
)
|
||||
|
||||
_FORBIDDEN_BRIDGE_KWARGS = (
|
||||
"epoch",
|
||||
"dependency_epoch",
|
||||
"materialized_epoch",
|
||||
"function_id",
|
||||
"field_id",
|
||||
"field_ids",
|
||||
"version",
|
||||
"branch",
|
||||
"wait",
|
||||
"job",
|
||||
"request",
|
||||
"backend",
|
||||
)
|
||||
|
||||
|
||||
def _definition_metadata_json(
|
||||
output_field_id: int,
|
||||
dependency_epoch: int,
|
||||
materialized_epoch: int,
|
||||
*,
|
||||
text_field_id: int = _ORDINARY_FIELD_ID,
|
||||
) -> str:
|
||||
"""Exact JSON stored under Arrow field metadata ``lancedb::generated_column``."""
|
||||
return json.dumps(
|
||||
{
|
||||
"format_version": 1,
|
||||
"output_field_id": output_field_id,
|
||||
"function_call": {
|
||||
"function_id": _STATUS_FUNCTION_ID,
|
||||
"arguments": [
|
||||
{
|
||||
"parameter": "text",
|
||||
"value": {
|
||||
"kind": "field",
|
||||
"field_id": text_field_id,
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
"dependency_epoch": dependency_epoch,
|
||||
"materialized_epoch": materialized_epoch,
|
||||
},
|
||||
separators=(",", ":"),
|
||||
)
|
||||
|
||||
|
||||
def _field(
|
||||
name: str,
|
||||
*,
|
||||
arrow_type: str = "string",
|
||||
nullable: bool = True,
|
||||
metadata: dict[str, str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
body: dict[str, Any] = {
|
||||
"name": name,
|
||||
"type": {"type": arrow_type},
|
||||
"nullable": nullable,
|
||||
}
|
||||
if metadata is not None:
|
||||
body["metadata"] = metadata
|
||||
return body
|
||||
|
||||
|
||||
def _status_schema_fields(
|
||||
*,
|
||||
complete_meta: str | None = None,
|
||||
incomplete_meta: str | None = None,
|
||||
bad_name: str | None = None,
|
||||
bad_meta: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
fields = [
|
||||
_field("ordinary", arrow_type="string"),
|
||||
_field(
|
||||
"complete_col",
|
||||
arrow_type="int32",
|
||||
metadata={
|
||||
_METADATA_KEY: complete_meta
|
||||
if complete_meta is not None
|
||||
else _definition_metadata_json(_COMPLETE_FIELD_ID, 3, 3)
|
||||
},
|
||||
),
|
||||
_field(
|
||||
"incomplete_col",
|
||||
arrow_type="int32",
|
||||
metadata={
|
||||
_METADATA_KEY: incomplete_meta
|
||||
if incomplete_meta is not None
|
||||
else _definition_metadata_json(_INCOMPLETE_FIELD_ID, 4, 1)
|
||||
},
|
||||
),
|
||||
]
|
||||
if bad_name is not None and bad_meta is not None:
|
||||
fields.append(
|
||||
_field(
|
||||
bad_name,
|
||||
arrow_type="int32",
|
||||
metadata={_METADATA_KEY: bad_meta},
|
||||
)
|
||||
)
|
||||
return {"fields": fields}
|
||||
|
||||
|
||||
def _describe_body(
|
||||
*,
|
||||
version: int = 11,
|
||||
field_ids: list[int] | None = _STABLE_FIELD_IDS,
|
||||
schema: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
body: dict[str, Any] = {
|
||||
"version": version,
|
||||
"schema": schema if schema is not None else _status_schema_fields(),
|
||||
}
|
||||
if field_ids is not None:
|
||||
body["field_ids"] = field_ids
|
||||
return body
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_text(exc: BaseException) -> str:
|
||||
parts = [str(exc), repr(exc)]
|
||||
current: BaseException | None = exc
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(f"{type(current).__name__}: {current}")
|
||||
current = current.__cause__ or current.__context__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _json_response(
|
||||
request: http.server.BaseHTTPRequestHandler, body: dict[str, Any]
|
||||
) -> None:
|
||||
payload = json.dumps(body).encode("utf-8")
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(payload)
|
||||
|
||||
|
||||
class _RequestLog:
|
||||
"""Track post-open describe and any non-describe operation traffic."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.describe: list[dict[str, Any]] = []
|
||||
self.other: list[str] = []
|
||||
self.recording = False
|
||||
|
||||
def start(self) -> None:
|
||||
self.describe.clear()
|
||||
self.other.clear()
|
||||
self.recording = True
|
||||
|
||||
def note(self, path: str, body: dict[str, Any] | None = None) -> None:
|
||||
if not self.recording:
|
||||
return
|
||||
if path == _DESCRIBE_PATH:
|
||||
self.describe.append(body or {})
|
||||
else:
|
||||
self.other.append(path)
|
||||
|
||||
|
||||
def _assert_no_operation_traffic(log: _RequestLog) -> None:
|
||||
assert log.describe == []
|
||||
assert log.other == []
|
||||
|
||||
|
||||
def _assert_one_status_describe(log: _RequestLog) -> None:
|
||||
assert len(log.describe) == 1, f"expected one status describe, got {log.describe!r}"
|
||||
assert log.other == [], f"unexpected non-describe traffic: {log.other!r}"
|
||||
|
||||
|
||||
def _assert_exact_public_signature(method: Any) -> None:
|
||||
"""Freeze ``(self, column_name)`` with no varargs/kwargs/keyword-only escape."""
|
||||
params = list(inspect.signature(method).parameters.values())
|
||||
assert [p.name for p in params] == ["self", "column_name"]
|
||||
for param in params:
|
||||
assert param.kind in (
|
||||
inspect.Parameter.POSITIONAL_ONLY,
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
)
|
||||
assert param.default is inspect.Parameter.empty
|
||||
assert param.kind is not inspect.Parameter.VAR_POSITIONAL
|
||||
assert param.kind is not inspect.Parameter.VAR_KEYWORD
|
||||
assert param.kind is not inspect.Parameter.KEYWORD_ONLY
|
||||
|
||||
|
||||
def _assert_status_string(value: Any, expected: str) -> None:
|
||||
assert value == expected
|
||||
assert type(value) is str
|
||||
assert value in ("complete", "incomplete")
|
||||
|
||||
|
||||
def _open_remote_table(
|
||||
*,
|
||||
status_describe: dict[str, Any] | None = None,
|
||||
):
|
||||
"""Open sync RemoteTable; return (table, log, cm)."""
|
||||
log = _RequestLog()
|
||||
binding = status_describe if status_describe is not None else _describe_body()
|
||||
open_describe = {
|
||||
"version": 1,
|
||||
"schema": {"fields": [_field("ordinary", arrow_type="string")]},
|
||||
}
|
||||
state = {"opened": False}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding)
|
||||
return
|
||||
|
||||
log.note(request.path, body)
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected path")
|
||||
|
||||
cm = _mock_remote_db(handler)
|
||||
db = cm.__enter__()
|
||||
table = db.open_table(_TABLE_NAME)
|
||||
assert isinstance(table, RemoteTable)
|
||||
log.start()
|
||||
return table, log, cm
|
||||
|
||||
|
||||
async def _open_remote_table_async(
|
||||
*,
|
||||
status_describe: dict[str, Any] | None = None,
|
||||
):
|
||||
"""Open async table under a live mock server; return (table, log, cm)."""
|
||||
log = _RequestLog()
|
||||
binding = status_describe if status_describe is not None else _describe_body()
|
||||
open_describe = {
|
||||
"version": 1,
|
||||
"schema": {"fields": [_field("ordinary", arrow_type="string")]},
|
||||
}
|
||||
state = {"opened": False}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding)
|
||||
return
|
||||
|
||||
log.note(request.path, body)
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected path")
|
||||
|
||||
cm = _mock_remote_db_async(handler)
|
||||
db = await cm.__aenter__()
|
||||
table = await db.open_table(_TABLE_NAME)
|
||||
assert isinstance(table, AsyncTable)
|
||||
log.start()
|
||||
return table, log, cm
|
||||
|
||||
|
||||
def test_no_public_generated_column_status_resource_exported():
|
||||
"""Baseline: no public status class/enum/resource is exported."""
|
||||
for mod in (lancedb, lancedb.table, _native):
|
||||
for name in _FORBIDDEN_PUBLIC_NAMES:
|
||||
assert not hasattr(mod, name), f"{mod.__name__}.{name} must not be public"
|
||||
|
||||
|
||||
def test_public_surface_signatures_annotations_and_hidden_bridge():
|
||||
"""Four public methods + hidden native bridge must exist with frozen shape."""
|
||||
assert hasattr(_native.Table, "_generated_column_status"), (
|
||||
"native private bridge Table._generated_column_status is missing"
|
||||
)
|
||||
assert hasattr(Table, "generated_column_status"), (
|
||||
"Table.generated_column_status is missing"
|
||||
)
|
||||
assert hasattr(LanceTable, "generated_column_status"), (
|
||||
"LanceTable.generated_column_status is missing"
|
||||
)
|
||||
assert hasattr(RemoteTable, "generated_column_status"), (
|
||||
"RemoteTable.generated_column_status is missing"
|
||||
)
|
||||
assert hasattr(AsyncTable, "generated_column_status"), (
|
||||
"AsyncTable.generated_column_status is missing"
|
||||
)
|
||||
|
||||
bridge = _native.Table._generated_column_status
|
||||
_assert_exact_public_signature(bridge)
|
||||
for keyword in _FORBIDDEN_BRIDGE_KWARGS:
|
||||
assert keyword not in inspect.signature(bridge).parameters
|
||||
|
||||
for method in (
|
||||
Table.generated_column_status,
|
||||
LanceTable.generated_column_status,
|
||||
RemoteTable.generated_column_status,
|
||||
):
|
||||
_assert_exact_public_signature(method)
|
||||
assert not inspect.iscoroutinefunction(method)
|
||||
assert get_type_hints(method)["return"] == _EXPECTED_RETURN
|
||||
|
||||
async_method = AsyncTable.generated_column_status
|
||||
_assert_exact_public_signature(async_method)
|
||||
assert inspect.iscoroutinefunction(async_method)
|
||||
assert get_type_hints(async_method)["return"] == _EXPECTED_RETURN
|
||||
|
||||
|
||||
def test_sync_remote_complete_and_incomplete_one_describe_each():
|
||||
table, log, cm = _open_remote_table()
|
||||
try:
|
||||
complete = table.generated_column_status("complete_col")
|
||||
_assert_status_string(complete, "complete")
|
||||
_assert_one_status_describe(log)
|
||||
|
||||
log.start()
|
||||
incomplete = table.generated_column_status("incomplete_col")
|
||||
_assert_status_string(incomplete, "incomplete")
|
||||
_assert_one_status_describe(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_complete_and_incomplete_one_describe_each():
|
||||
table, log, cm = await _open_remote_table_async()
|
||||
try:
|
||||
complete = await table.generated_column_status("complete_col")
|
||||
_assert_status_string(complete, "complete")
|
||||
_assert_one_status_describe(log)
|
||||
|
||||
log.start()
|
||||
incomplete = await table.generated_column_status("incomplete_col")
|
||||
_assert_status_string(incomplete, "incomplete")
|
||||
_assert_one_status_describe(log)
|
||||
finally:
|
||||
await cm.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("column_name", "status_describe", "expected_exc"),
|
||||
[
|
||||
(
|
||||
"missing",
|
||||
_describe_body(),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"Complete_Col",
|
||||
_describe_body(),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"ordinary",
|
||||
_describe_body(),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"complete_col",
|
||||
_describe_body(
|
||||
schema=_status_schema_fields(
|
||||
complete_meta=_definition_metadata_json(
|
||||
_COMPLETE_FIELD_ID + 1, 3, 3
|
||||
)
|
||||
)
|
||||
),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"gen_bad",
|
||||
_describe_body(
|
||||
field_ids=[*_STABLE_FIELD_IDS, 9],
|
||||
schema=_status_schema_fields(
|
||||
bad_name="gen_bad",
|
||||
bad_meta=(
|
||||
'{"format_version":1,"output_field_id":9,'
|
||||
f'"function_call":{_RAW_METADATA_MARKER},'
|
||||
'"dependency_epoch":1,"materialized_epoch":1}'
|
||||
),
|
||||
),
|
||||
),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"complete_col",
|
||||
_describe_body(
|
||||
schema=_status_schema_fields(
|
||||
complete_meta=_definition_metadata_json(
|
||||
_COMPLETE_FIELD_ID, 1, 1
|
||||
).replace('"format_version":1', '"format_version":2')
|
||||
)
|
||||
),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"incomplete_col",
|
||||
_describe_body(
|
||||
schema=_status_schema_fields(
|
||||
incomplete_meta=_definition_metadata_json(
|
||||
_INCOMPLETE_FIELD_ID, 1, 2
|
||||
)
|
||||
)
|
||||
),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"complete_col",
|
||||
_describe_body(field_ids=None),
|
||||
NotImplementedError,
|
||||
),
|
||||
],
|
||||
ids=[
|
||||
"missing",
|
||||
"case_mismatch",
|
||||
"ordinary",
|
||||
"output_id_mismatch",
|
||||
"malformed_metadata",
|
||||
"unknown_format_version",
|
||||
"reversed_epochs",
|
||||
"old_server_missing_field_ids",
|
||||
],
|
||||
)
|
||||
def test_remote_fail_closed_matrix_one_describe(
|
||||
column_name: str,
|
||||
status_describe: dict[str, Any],
|
||||
expected_exc: type[BaseException],
|
||||
):
|
||||
table, log, cm = _open_remote_table(status_describe=status_describe)
|
||||
try:
|
||||
with pytest.raises(expected_exc) as raised:
|
||||
table.generated_column_status(column_name)
|
||||
text = _exception_text(raised.value)
|
||||
assert _RAW_METADATA_MARKER not in text
|
||||
_assert_one_status_describe(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_sync_empty_name_zero_post_open_requests():
|
||||
table, log, cm = _open_remote_table()
|
||||
try:
|
||||
with pytest.raises(ValueError):
|
||||
table.generated_column_status("")
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_empty_name_zero_post_open_requests():
|
||||
table, log, cm = await _open_remote_table_async()
|
||||
try:
|
||||
with pytest.raises(ValueError):
|
||||
await table.generated_column_status("")
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
await cm.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_closed_status_empty_validation_wins_and_nonempty_closed():
|
||||
"""Publicly closed AsyncTable: empty validates first; nonempty is closed."""
|
||||
table, log, cm = await _open_remote_table_async()
|
||||
try:
|
||||
table.close()
|
||||
|
||||
log.start()
|
||||
try:
|
||||
await table.generated_column_status("complete_col")
|
||||
except AttributeError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
text = _exception_text(exc)
|
||||
assert "closed" in text.lower()
|
||||
else:
|
||||
pytest.fail("closed AsyncTable must fail before transport")
|
||||
_assert_no_operation_traffic(log)
|
||||
|
||||
log.start()
|
||||
with pytest.raises(ValueError) as raised:
|
||||
await table.generated_column_status("")
|
||||
text = _exception_text(raised.value)
|
||||
assert "closed" not in text.lower()
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
await cm.__aexit__(None, None, None)
|
||||
|
||||
|
||||
def test_local_sync_ordinary_column_fails_without_side_effects(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table(
|
||||
"ordinary_only",
|
||||
[{"ordinary": "alpha", "id": 1}, {"ordinary": "beta", "id": 2}],
|
||||
)
|
||||
assert isinstance(table, LanceTable)
|
||||
version_before = table.version
|
||||
schema_before = table.schema
|
||||
data_before = table.to_arrow()
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
table.generated_column_status("ordinary")
|
||||
|
||||
assert table.version == version_before
|
||||
assert table.schema == schema_before
|
||||
assert table.to_arrow().equals(data_before)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_ordinary_column_fails_without_side_effects(tmp_path):
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
table = await db.create_table(
|
||||
"ordinary_only_async",
|
||||
[{"ordinary": "alpha", "id": 1}, {"ordinary": "beta", "id": 2}],
|
||||
)
|
||||
assert isinstance(table, AsyncTable)
|
||||
version_before = await table.version()
|
||||
schema_before = await table.schema()
|
||||
data_before = await table.to_arrow()
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
await table.generated_column_status("ordinary")
|
||||
|
||||
assert await table.version() == version_before
|
||||
assert await table.schema() == schema_before
|
||||
assert (await table.to_arrow()).equals(data_before)
|
||||
@@ -0,0 +1,291 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for the local @udf declaration surface."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import inspect
|
||||
import types
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import Function, Job, udf
|
||||
from lancedb._udf import _get_udf_config
|
||||
|
||||
_REMOVED_AUTHORING_KNOBS = (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_handling",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"error_policy",
|
||||
"FunctionVersion",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
)
|
||||
|
||||
|
||||
def _decorate(fn, **overrides):
|
||||
kwargs = {
|
||||
"inputs": {"x": pa.int32()},
|
||||
"output": pa.int64(),
|
||||
"python": "3.12",
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return udf(**kwargs)(fn)
|
||||
|
||||
|
||||
def test_udf_top_level_export_and_identity_metadata_behavior():
|
||||
assert "udf" in lancedb.__all__
|
||||
assert udf is lancedb.udf
|
||||
assert isinstance(importlib.import_module("lancedb._udf"), types.ModuleType)
|
||||
assert not isinstance(lancedb.udf, types.ModuleType)
|
||||
|
||||
def add(x, y=1):
|
||||
"""Add locally."""
|
||||
return x + y
|
||||
|
||||
original = add
|
||||
decorated = _decorate(
|
||||
add,
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
)
|
||||
|
||||
assert decorated is original
|
||||
assert decorated.__name__ == "add"
|
||||
assert decorated.__doc__ == "Add locally."
|
||||
assert str(inspect.signature(decorated)) == "(x, y=1)"
|
||||
assert decorated(2) == 3
|
||||
assert decorated(2, 5) == 7
|
||||
assert decorated(x=4, y=6) == 10
|
||||
|
||||
|
||||
def test_udf_config_snapshot_order_defaults_and_immutability():
|
||||
inputs = {"z": pa.string(), "a": pa.int32()}
|
||||
packages = ["pkg-b==2", "pkg-a==1"]
|
||||
|
||||
def combine(z, a):
|
||||
return f"{z}:{a}"
|
||||
|
||||
decorated = udf(
|
||||
inputs=inputs,
|
||||
output=pa.string(),
|
||||
python="3.11",
|
||||
packages=packages,
|
||||
output_nullable=False,
|
||||
)(combine)
|
||||
|
||||
config = _get_udf_config(decorated)
|
||||
assert config.inputs == (("z", pa.string()), ("a", pa.int32()))
|
||||
assert isinstance(config.inputs, tuple)
|
||||
assert config.output == pa.string()
|
||||
assert config.output_nullable is False
|
||||
assert config.python == "3.11"
|
||||
assert config.packages == ("pkg-b==2", "pkg-a==1")
|
||||
assert isinstance(config.packages, tuple)
|
||||
|
||||
inputs["extra"] = pa.bool_()
|
||||
del inputs["z"]
|
||||
packages.append("pkg-c==3")
|
||||
packages[0] = "mutated==0"
|
||||
assert config.inputs == (("z", pa.string()), ("a", pa.int32()))
|
||||
assert config.packages == ("pkg-b==2", "pkg-a==1")
|
||||
|
||||
for attr in ("inputs", "output", "output_nullable", "python", "packages"):
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(config, attr, None)
|
||||
|
||||
def defaults_only(x):
|
||||
return x
|
||||
|
||||
defaulted = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
)(defaults_only)
|
||||
default_config = _get_udf_config(defaulted)
|
||||
assert default_config.packages == ()
|
||||
assert default_config.output_nullable is True
|
||||
|
||||
|
||||
def test_udf_accepts_lambda_and_closure_for_local_declaration():
|
||||
ambient = "ambient-secret-value-xyz"
|
||||
|
||||
lam = udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(lambda n: n + 1)
|
||||
assert lam(3) == 4
|
||||
assert _get_udf_config(lam).inputs == (("n", pa.int32()),)
|
||||
|
||||
def factory(offset):
|
||||
@udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
packages=["demo==0.1"],
|
||||
)
|
||||
def closed(n):
|
||||
return n + offset + len(ambient)
|
||||
|
||||
return closed
|
||||
|
||||
closed = factory(10)
|
||||
assert closed(2) == 12 + len(ambient)
|
||||
assert _get_udf_config(closed).packages == ("demo==0.1",)
|
||||
|
||||
|
||||
def test_udf_declaration_defers_signature_and_implementation_packaging():
|
||||
"""Declaration must not validate callable signature or embed implementation."""
|
||||
|
||||
def local_add(left, right=1):
|
||||
return left + right
|
||||
|
||||
decorated = udf(
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(local_add)
|
||||
|
||||
assert decorated is local_add
|
||||
assert str(inspect.signature(decorated)) == "(left, right=1)"
|
||||
assert decorated(2) == 3
|
||||
assert decorated(2, 5) == 7
|
||||
|
||||
config = _get_udf_config(decorated)
|
||||
assert config.inputs == (("x", pa.int32()), ("y", pa.int32()))
|
||||
for attr in (
|
||||
"source",
|
||||
"module",
|
||||
"callable",
|
||||
"function",
|
||||
"implementation",
|
||||
"bundle",
|
||||
"artifact",
|
||||
"digest",
|
||||
):
|
||||
assert not hasattr(config, attr)
|
||||
|
||||
|
||||
def test_udf_lookup_double_decoration_and_non_function_target():
|
||||
def plain(x):
|
||||
return x
|
||||
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
_get_udf_config(plain)
|
||||
|
||||
decorated = _decorate(plain)
|
||||
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
_decorate(decorated)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(object())
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(42)
|
||||
|
||||
|
||||
def test_udf_config_validation_errors():
|
||||
def target(x):
|
||||
return x
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
udf({"x": pa.int32()}, pa.int32(), "3.12")(target)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, inputs=[("x", pa.int32())])
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, inputs={1: pa.int32()})
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_decorate(target, inputs={"": pa.int32()})
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, inputs={"x": "int32"})
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, output="int64")
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, python=3.12)
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_decorate(target, python="")
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, packages="pkg==1")
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_decorate(target, packages=["pkg==1", ""])
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_decorate(target, packages=["pkg==1", "pkg==1"])
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, packages=["pkg==1", 2])
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, output_nullable=1)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, output_nullable="true")
|
||||
|
||||
|
||||
def test_udf_rejects_removed_overdesign_and_has_no_durable_side_effects():
|
||||
params = inspect.signature(udf).parameters
|
||||
for name in _REMOVED_AUTHORING_KNOBS:
|
||||
assert name not in params
|
||||
|
||||
def score(x):
|
||||
"""score body marker unique-xyz."""
|
||||
ambient = "ambient-secret-value-xyz"
|
||||
return f"{ambient}:{x}"
|
||||
|
||||
decorated = _decorate(
|
||||
score,
|
||||
packages=["score==1.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
config = _get_udf_config(decorated)
|
||||
text = repr(config).lower()
|
||||
|
||||
assert "score body marker unique-xyz" not in text
|
||||
assert "ambient-secret-value-xyz" not in text
|
||||
for token in (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"functionversion",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
):
|
||||
assert token not in text
|
||||
|
||||
for attr in _REMOVED_AUTHORING_KNOBS:
|
||||
assert not hasattr(config, attr)
|
||||
|
||||
assert not isinstance(decorated, Function)
|
||||
assert not isinstance(decorated, Job)
|
||||
for attr in ("id", "function_id", "job", "job_id", "registration"):
|
||||
assert not hasattr(decorated, attr)
|
||||
@@ -0,0 +1,490 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for local FunctionCapability authoring and @udf capabilities."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import Function, FunctionCapability, Job, udf
|
||||
from lancedb._udf import _get_udf_config, _package_udf
|
||||
|
||||
_SECRET_REFERENCE = "secret://team/capability-redact-token-xyz"
|
||||
_SECRET_ENV = "API_TOKEN"
|
||||
_NETWORK_ORIGIN = "https://api.example.com"
|
||||
|
||||
_OVERDESIGN_ATTRS = (
|
||||
"id",
|
||||
"function_id",
|
||||
"FunctionVersion",
|
||||
"function_version",
|
||||
"artifact",
|
||||
"digest",
|
||||
"authorization",
|
||||
"authorized",
|
||||
"value",
|
||||
"plaintext",
|
||||
"plaintext_secret",
|
||||
"secret_value",
|
||||
"job",
|
||||
"job_id",
|
||||
"catalog",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_handling",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"error_policy",
|
||||
"geneva",
|
||||
)
|
||||
|
||||
|
||||
def _decorate(fn, **overrides):
|
||||
kwargs = {
|
||||
"inputs": {"x": pa.int32()},
|
||||
"output": pa.int64(),
|
||||
"python": "3.12",
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return udf(**kwargs)(fn)
|
||||
|
||||
|
||||
def _network(origin: str = _NETWORK_ORIGIN) -> FunctionCapability:
|
||||
return FunctionCapability.network(origin)
|
||||
|
||||
|
||||
def _secret(
|
||||
reference: str = _SECRET_REFERENCE,
|
||||
*,
|
||||
environment_variable: str = _SECRET_ENV,
|
||||
) -> FunctionCapability:
|
||||
return FunctionCapability.secret(
|
||||
reference,
|
||||
environment_variable=environment_variable,
|
||||
)
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=False,
|
||||
)
|
||||
def packable_without_capabilities(x):
|
||||
return x + 1
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=False,
|
||||
capabilities=[
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
],
|
||||
)
|
||||
def packable_with_capabilities(x):
|
||||
return x + 1
|
||||
|
||||
|
||||
def test_function_capability_export_factories_projection_equality_immutability():
|
||||
assert "FunctionCapability" in lancedb.__all__
|
||||
assert FunctionCapability is lancedb.FunctionCapability
|
||||
|
||||
network = _network()
|
||||
secret = _secret()
|
||||
|
||||
assert network.kind == "network"
|
||||
assert network.origin == _NETWORK_ORIGIN
|
||||
assert network.reference is None
|
||||
assert network.environment_variable is None
|
||||
|
||||
assert secret.kind == "secret"
|
||||
assert secret.reference == _SECRET_REFERENCE
|
||||
assert secret.environment_variable == _SECRET_ENV
|
||||
assert secret.origin is None
|
||||
|
||||
assert network == FunctionCapability.network(_NETWORK_ORIGIN)
|
||||
assert secret == FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
)
|
||||
assert network != secret
|
||||
assert network != FunctionCapability.network("https://other.example.com")
|
||||
assert secret != FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable="OTHER_TOKEN",
|
||||
)
|
||||
|
||||
public_attrs = ("kind", "origin", "reference", "environment_variable")
|
||||
internal_slots = ("_kind", "_origin", "_reference", "_environment_variable")
|
||||
immutable_attrs = public_attrs + internal_slots
|
||||
|
||||
for attr in public_attrs:
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(network, attr, None)
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(secret, attr, None)
|
||||
|
||||
for attr in immutable_attrs:
|
||||
# Fresh instances per attempt so a RED slot mutation cannot corrupt
|
||||
# shared fixtures used by later assertions in this test.
|
||||
fresh_network = _network("https://fresh-immutability.example.com")
|
||||
fresh_secret = _secret(
|
||||
"secret://team/fresh-immutability-token",
|
||||
environment_variable="FRESH_IMMUTABILITY_TOKEN",
|
||||
)
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(fresh_network, attr, None)
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(fresh_secret, attr, None)
|
||||
with pytest.raises(AttributeError):
|
||||
delattr(fresh_network, attr)
|
||||
with pytest.raises(AttributeError):
|
||||
delattr(fresh_secret, attr)
|
||||
|
||||
retained_origin = "https://config-retain.example.com"
|
||||
retained_reference = "secret://team/config-retain-token"
|
||||
retained_env = "CONFIG_RETAIN_TOKEN"
|
||||
retained_network = FunctionCapability.network(retained_origin)
|
||||
retained_secret = FunctionCapability.secret(
|
||||
retained_reference,
|
||||
environment_variable=retained_env,
|
||||
)
|
||||
expected_capabilities = (
|
||||
FunctionCapability.network(retained_origin),
|
||||
FunctionCapability.secret(
|
||||
retained_reference,
|
||||
environment_variable=retained_env,
|
||||
),
|
||||
)
|
||||
|
||||
def retain_target(x):
|
||||
return x
|
||||
|
||||
retained = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
capabilities=[retained_network, retained_secret],
|
||||
)(retain_target)
|
||||
retained_config = _get_udf_config(retained)
|
||||
assert retained_config.capabilities == expected_capabilities
|
||||
|
||||
for attr in immutable_attrs:
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(retained_network, attr, "mutated")
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(retained_secret, attr, "mutated")
|
||||
with pytest.raises(AttributeError):
|
||||
delattr(retained_network, attr)
|
||||
with pytest.raises(AttributeError):
|
||||
delattr(retained_secret, attr)
|
||||
|
||||
assert retained_config.capabilities == expected_capabilities
|
||||
assert retained_config.capabilities[0] is retained_network
|
||||
assert retained_config.capabilities[1] is retained_secret
|
||||
assert retained_config.capabilities[0].kind == "network"
|
||||
assert retained_config.capabilities[0].origin == retained_origin
|
||||
assert retained_config.capabilities[0].reference is None
|
||||
assert retained_config.capabilities[0].environment_variable is None
|
||||
assert retained_config.capabilities[1].kind == "secret"
|
||||
assert retained_config.capabilities[1].reference == retained_reference
|
||||
assert retained_config.capabilities[1].environment_variable == retained_env
|
||||
assert retained_config.capabilities[1].origin is None
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability()
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability( # type: ignore[call-arg]
|
||||
kind="network",
|
||||
origin=_NETWORK_ORIGIN,
|
||||
)
|
||||
|
||||
assert not isinstance(network, Function)
|
||||
assert not isinstance(secret, Function)
|
||||
assert not isinstance(network, Job)
|
||||
assert not isinstance(secret, Job)
|
||||
for attr in _OVERDESIGN_ATTRS:
|
||||
assert not hasattr(network, attr)
|
||||
assert not hasattr(secret, attr)
|
||||
|
||||
|
||||
def test_function_capability_validation_and_secret_redaction():
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.network(None) # type: ignore[arg-type]
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.network(123) # type: ignore[arg-type]
|
||||
with pytest.raises(ValueError):
|
||||
FunctionCapability.network("")
|
||||
|
||||
# Backend authorization owns URL/scheme policy; non-empty is enough here.
|
||||
loose = FunctionCapability.network("example.com")
|
||||
assert loose.kind == "network"
|
||||
assert loose.origin == "example.com"
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[misc]
|
||||
_SECRET_REFERENCE,
|
||||
_SECRET_ENV,
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret(None, environment_variable=_SECRET_ENV) # type: ignore[arg-type]
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret(123, environment_variable=_SECRET_ENV) # type: ignore[arg-type]
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable=None) # type: ignore[arg-type]
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable=1) # type: ignore[arg-type]
|
||||
|
||||
with pytest.raises(ValueError) as empty_ref:
|
||||
FunctionCapability.secret("", environment_variable=_SECRET_ENV)
|
||||
assert _SECRET_REFERENCE not in str(empty_ref.value)
|
||||
assert _SECRET_REFERENCE not in repr(empty_ref.value)
|
||||
|
||||
with pytest.raises(ValueError) as empty_env:
|
||||
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable="")
|
||||
assert _SECRET_REFERENCE not in str(empty_env.value)
|
||||
assert _SECRET_REFERENCE not in repr(empty_env.value)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[call-arg]
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
value="super-secret",
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[call-arg]
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
plaintext_secret="super-secret",
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[call-arg]
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
environment={_SECRET_ENV: "super-secret"},
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[call-arg]
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
headers={"Authorization": "Bearer super-secret"},
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.network( # type: ignore[call-arg]
|
||||
_NETWORK_ORIGIN,
|
||||
headers={"X-Trace": "1"},
|
||||
)
|
||||
|
||||
secret = _secret()
|
||||
assert not hasattr(secret, "value")
|
||||
assert not hasattr(secret, "plaintext")
|
||||
assert not hasattr(secret, "plaintext_secret")
|
||||
assert not hasattr(secret, "secret_value")
|
||||
|
||||
secret_text = repr(secret)
|
||||
assert "secret" in secret_text.lower()
|
||||
assert _SECRET_ENV in secret_text
|
||||
assert _SECRET_REFERENCE not in secret_text
|
||||
assert "super-secret" not in secret_text
|
||||
|
||||
network_text = repr(_network())
|
||||
assert "network" in network_text.lower()
|
||||
assert _NETWORK_ORIGIN in network_text
|
||||
|
||||
|
||||
def test_udf_capabilities_ordered_immutable_config_default_and_validation():
|
||||
params = inspect.signature(udf).parameters
|
||||
assert "capabilities" in params
|
||||
assert params["capabilities"].kind is inspect.Parameter.KEYWORD_ONLY
|
||||
assert params["capabilities"].default == ()
|
||||
|
||||
def identity_target(x):
|
||||
"""capabilities identity marker."""
|
||||
return x + 1
|
||||
|
||||
original = identity_target
|
||||
decorated = _decorate(identity_target)
|
||||
assert decorated is original
|
||||
assert decorated.__name__ == "identity_target"
|
||||
assert decorated.__doc__ == "capabilities identity marker."
|
||||
assert decorated(2) == 3
|
||||
assert _get_udf_config(decorated).capabilities == ()
|
||||
|
||||
first = _network("https://b.example.com")
|
||||
second = _network("https://a.example.com")
|
||||
third = _network("https://b.example.com")
|
||||
secret = _secret()
|
||||
capabilities = [first, second, third, secret]
|
||||
|
||||
def combine(x):
|
||||
return x
|
||||
|
||||
with_caps = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
packages=["pkg-b==2", "pkg-a==1"],
|
||||
capabilities=capabilities,
|
||||
)(combine)
|
||||
config = _get_udf_config(with_caps)
|
||||
assert config.capabilities == (first, second, third, secret)
|
||||
assert isinstance(config.capabilities, tuple)
|
||||
assert config.packages == ("pkg-b==2", "pkg-a==1")
|
||||
assert config.inputs == (("x", pa.int32()),)
|
||||
|
||||
capabilities.append(_network("https://mutated.example.com"))
|
||||
capabilities[0] = _network("https://replaced.example.com")
|
||||
assert config.capabilities == (first, second, third, secret)
|
||||
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(config, "capabilities", ())
|
||||
|
||||
def target(x):
|
||||
return x
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, capabilities="https://api.example.com")
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, capabilities=b"https://api.example.com")
|
||||
|
||||
class _BadCapability:
|
||||
def __repr__(self) -> str:
|
||||
return "unique-bad-capability-repr-xyz"
|
||||
|
||||
with pytest.raises(TypeError) as bad_item:
|
||||
_decorate(target, capabilities=[_BadCapability()])
|
||||
assert "unique-bad-capability-repr-xyz" not in str(bad_item.value)
|
||||
assert "unique-bad-capability-repr-xyz" not in repr(bad_item.value)
|
||||
|
||||
with pytest.raises(TypeError) as bad_mixed:
|
||||
_decorate(
|
||||
target,
|
||||
capabilities=[_network(), "unique-bad-capability-string-xyz"],
|
||||
)
|
||||
assert "unique-bad-capability-string-xyz" not in str(bad_mixed.value)
|
||||
assert "unique-bad-capability-string-xyz" not in repr(bad_mixed.value)
|
||||
|
||||
|
||||
def test_udf_capabilities_rejects_function_capability_subclass_before_property_access():
|
||||
marker = "unique-hostile-capability-subclass-marker-xyz"
|
||||
|
||||
class _HostileFunctionCapability(FunctionCapability):
|
||||
@property
|
||||
def kind(self) -> str:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def origin(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def reference(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def environment_variable(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
hostile = object.__new__(_HostileFunctionCapability)
|
||||
assert isinstance(hostile, FunctionCapability)
|
||||
assert type(hostile) is not FunctionCapability
|
||||
|
||||
def target(x):
|
||||
return x
|
||||
|
||||
with pytest.raises(TypeError) as exc_info:
|
||||
_decorate(target, capabilities=[hostile])
|
||||
assert marker not in str(exc_info.value)
|
||||
assert marker not in repr(exc_info.value)
|
||||
assert _SECRET_REFERENCE not in str(exc_info.value)
|
||||
assert _SECRET_REFERENCE not in repr(exc_info.value)
|
||||
assert exc_info.value.__cause__ is None
|
||||
assert exc_info.value.__context__ is None
|
||||
|
||||
|
||||
def test_package_udf_preserves_capabilities_and_redacts_secret_reference():
|
||||
packaged = _package_udf(packable_with_capabilities)
|
||||
config = packaged.config
|
||||
|
||||
assert packaged.config is _get_udf_config(packable_with_capabilities)
|
||||
assert config.capabilities == (
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
)
|
||||
assert config.capabilities[0].kind == "network"
|
||||
assert config.capabilities[0].origin == _NETWORK_ORIGIN
|
||||
assert config.capabilities[1].kind == "secret"
|
||||
assert config.capabilities[1].reference == _SECRET_REFERENCE
|
||||
assert config.capabilities[1].environment_variable == _SECRET_ENV
|
||||
assert config.packages == ("pkg-a==1",)
|
||||
assert config.python == "3.12"
|
||||
assert config.output_nullable is False
|
||||
|
||||
nested = (
|
||||
f"{packaged!r}\n{config!r}\n{config.capabilities!r}\n{config.capabilities[1]!r}"
|
||||
)
|
||||
assert _SECRET_REFERENCE not in nested
|
||||
assert _SECRET_ENV in repr(config.capabilities[1])
|
||||
|
||||
|
||||
def test_capabilities_are_additive_to_existing_declaration_and_packaging():
|
||||
def score(x):
|
||||
return x
|
||||
|
||||
decorated = _decorate(
|
||||
score,
|
||||
packages=["score==1.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
config = _get_udf_config(decorated)
|
||||
assert config.inputs == (("x", pa.int32()),)
|
||||
assert config.output == pa.int64()
|
||||
assert config.output_nullable is True
|
||||
assert config.python == "3.12"
|
||||
assert config.packages == ("score==1.0",)
|
||||
assert config.capabilities == ()
|
||||
assert decorated is score
|
||||
assert decorated(4) == 4
|
||||
|
||||
packaged = _package_udf(packable_without_capabilities)
|
||||
assert packaged.config is _get_udf_config(packable_without_capabilities)
|
||||
assert packaged.callable_name == "packable_without_capabilities"
|
||||
assert packaged.config.capabilities == ()
|
||||
assert packaged.config.packages == ("pkg-a==1",)
|
||||
assert packaged.config.output_nullable is False
|
||||
assert packable_without_capabilities(1) == 2
|
||||
|
||||
params = inspect.signature(udf).parameters
|
||||
for name in (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_handling",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"error_policy",
|
||||
"FunctionVersion",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
):
|
||||
assert name not in params
|
||||
@@ -0,0 +1,506 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for the private UDF -> FunctionDefinition bridge."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import io
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import FunctionCapability, udf
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb import _udf as _udf_mod
|
||||
|
||||
_SOURCE_MARKER = "bridge-source-marker-unique-xyz"
|
||||
_SECRET_REFERENCE = "secret://team/bridge-redact-token-xyz"
|
||||
_SECRET_ENV = "BRIDGE_API_TOKEN"
|
||||
_NETWORK_ORIGIN = "https://api.bridge-example.com"
|
||||
_NETWORK_ORIGIN_B = "https://other.bridge-example.com"
|
||||
|
||||
_FORBIDDEN_WIRE_KEYS = (
|
||||
"id",
|
||||
"function_id",
|
||||
"FunctionId",
|
||||
"catalog",
|
||||
"catalog_name",
|
||||
"version",
|
||||
"function_version",
|
||||
"FunctionVersion",
|
||||
"lineage",
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"digest",
|
||||
"artifact",
|
||||
"artifact_digest",
|
||||
"storage",
|
||||
"storage_location",
|
||||
"location",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"nullPolicy",
|
||||
"timestamp",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"worker",
|
||||
"scheduler",
|
||||
"attempt",
|
||||
"attempt_id",
|
||||
"replica",
|
||||
"placement",
|
||||
"job",
|
||||
"job_id",
|
||||
"retry_key",
|
||||
"registration",
|
||||
)
|
||||
|
||||
_OVERDESIGN_ATTRS = (
|
||||
"id",
|
||||
"function_id",
|
||||
"FunctionVersion",
|
||||
"function_version",
|
||||
"artifact",
|
||||
"digest",
|
||||
"job",
|
||||
"job_id",
|
||||
"catalog",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"null_handling",
|
||||
)
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"text": pa.string(), "limit": pa.int32()},
|
||||
output=pa.string(),
|
||||
python="3.12",
|
||||
packages=["pkg-b==2", "pkg-a==1"],
|
||||
output_nullable=True,
|
||||
capabilities=[
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
FunctionCapability.network(_NETWORK_ORIGIN_B),
|
||||
],
|
||||
)
|
||||
def packable_bridge_normalize(text, limit):
|
||||
"""bridge-source-marker-unique-xyz."""
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def _build_function_definition(fn: object):
|
||||
return _udf_mod._build_function_definition(fn)
|
||||
|
||||
|
||||
def _function_definition_type():
|
||||
return _native._FunctionDefinition
|
||||
|
||||
|
||||
def _new_function_definition(**kwargs):
|
||||
return _native._new_function_definition(**kwargs)
|
||||
|
||||
|
||||
def _json_bytes(definition) -> bytes:
|
||||
payload = definition._to_json()
|
||||
if isinstance(payload, bytes):
|
||||
return payload
|
||||
assert isinstance(payload, str)
|
||||
return payload.encode("utf-8")
|
||||
|
||||
|
||||
def _decode_type_ipc(encoded: str) -> pa.DataType:
|
||||
raw = base64.b64decode(encoded)
|
||||
reader = pa.ipc.open_file(io.BytesIO(raw))
|
||||
assert reader.num_record_batches == 0
|
||||
assert len(reader.schema) == 1
|
||||
return reader.schema.field(0).type
|
||||
|
||||
|
||||
def _assert_exact_object_keys(value: dict, expected: set[str], *, context: str) -> None:
|
||||
assert isinstance(value, dict), f"{context} must be an object"
|
||||
assert set(value) == expected, f"{context} keys must match exactly: {set(value)!r}"
|
||||
|
||||
|
||||
def _assert_forbidden_keys_absent(value: object, *, context: str) -> None:
|
||||
if isinstance(value, dict):
|
||||
for key in value:
|
||||
assert key not in _FORBIDDEN_WIRE_KEYS, (
|
||||
f"forbidden key {key!r} at {context}: {value!r}"
|
||||
)
|
||||
if key == "name" and context in {
|
||||
"definition",
|
||||
"signature",
|
||||
"signature.output",
|
||||
"implementation",
|
||||
}:
|
||||
raise AssertionError(
|
||||
f"catalog/function identity key `name` must be absent at {context}"
|
||||
)
|
||||
child_context = f"{context}.{key}"
|
||||
if key == "parameters" and context == "signature":
|
||||
child_context = "signature.parameters"
|
||||
_assert_forbidden_keys_absent(value[key], context=child_context)
|
||||
elif isinstance(value, list):
|
||||
for idx, item in enumerate(value):
|
||||
item_context = (
|
||||
f"signature.parameters[{idx}]"
|
||||
if context == "signature.parameters"
|
||||
else f"{context}[{idx}]"
|
||||
)
|
||||
if context == "signature.parameters":
|
||||
assert isinstance(item, dict)
|
||||
assert "name" in item
|
||||
for key in item:
|
||||
assert key not in _FORBIDDEN_WIRE_KEYS
|
||||
assert key != "catalog_name"
|
||||
_assert_forbidden_keys_absent(
|
||||
{k: v for k, v in item.items() if k != "name"},
|
||||
context=item_context,
|
||||
)
|
||||
else:
|
||||
_assert_forbidden_keys_absent(item, context=item_context)
|
||||
|
||||
|
||||
def _assert_sanitized_text(*parts: object) -> None:
|
||||
combined = "\n".join(str(part) for part in parts)
|
||||
lowered = combined.lower()
|
||||
assert _SOURCE_MARKER.lower() not in lowered
|
||||
assert _SECRET_REFERENCE.lower() not in lowered
|
||||
assert str(Path(__file__).resolve()).lower() not in lowered
|
||||
assert Path(__file__).resolve().as_posix().lower() not in lowered
|
||||
|
||||
|
||||
def _assert_clean_validation_error(exc_info) -> None:
|
||||
_assert_sanitized_text(exc_info.value, repr(exc_info.value))
|
||||
assert exc_info.value.__cause__ is None
|
||||
assert exc_info.value.__context__ is None
|
||||
|
||||
|
||||
def _valid_builder_kwargs(**overrides):
|
||||
kwargs = {
|
||||
"parameters": [("text", pa.string()), ("limit", pa.int32())],
|
||||
"output_type": pa.string(),
|
||||
"output_nullable": True,
|
||||
"module": "bridge_mod",
|
||||
"callable_name": "normalize",
|
||||
"source": (
|
||||
"def normalize(text, limit):\n"
|
||||
f" # {_SOURCE_MARKER}\n"
|
||||
" return text[:limit]\n"
|
||||
),
|
||||
"python": "3.12",
|
||||
"packages": ["pkg-b==2", "pkg-a==1"],
|
||||
"capabilities": [
|
||||
("network", _NETWORK_ORIGIN, None),
|
||||
("secret", _SECRET_REFERENCE, _SECRET_ENV),
|
||||
("network", _NETWORK_ORIGIN_B, None),
|
||||
],
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return kwargs
|
||||
|
||||
|
||||
def test_build_function_definition_private_native_immutability_and_export_surface():
|
||||
assert "_build_function_definition" not in getattr(lancedb, "__all__", [])
|
||||
assert "_FunctionDefinition" not in lancedb.__all__
|
||||
assert not hasattr(lancedb, "_FunctionDefinition")
|
||||
assert not hasattr(lancedb, "_build_function_definition")
|
||||
assert not hasattr(lancedb, "_new_function_definition")
|
||||
|
||||
definition = _build_function_definition(packable_bridge_normalize)
|
||||
definition_type = _function_definition_type()
|
||||
assert type(definition) is definition_type
|
||||
assert definition_type.__module__ == "lancedb._lancedb"
|
||||
assert definition_type.__name__ == "_FunctionDefinition"
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
definition_type()
|
||||
|
||||
for attr in _OVERDESIGN_ATTRS:
|
||||
assert not hasattr(definition, attr)
|
||||
|
||||
for attr in ("signature", "module", "source", "capabilities"):
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(definition, attr, None)
|
||||
|
||||
|
||||
def test_build_function_definition_json_wire_ordered_contract_without_identity():
|
||||
definition = _build_function_definition(packable_bridge_normalize)
|
||||
encoded_a = _json_bytes(definition)
|
||||
encoded_b = _json_bytes(definition)
|
||||
assert encoded_a == encoded_b
|
||||
|
||||
wire = json.loads(encoded_a.decode("utf-8"))
|
||||
_assert_exact_object_keys(
|
||||
wire,
|
||||
{"format_version", "signature", "implementation", "capabilities"},
|
||||
context="definition",
|
||||
)
|
||||
assert wire["format_version"] == 1
|
||||
_assert_forbidden_keys_absent(wire, context="definition")
|
||||
|
||||
signature = wire["signature"]
|
||||
_assert_exact_object_keys(signature, {"parameters", "output"}, context="signature")
|
||||
parameters = signature["parameters"]
|
||||
assert [parameter["name"] for parameter in parameters] == ["text", "limit"]
|
||||
for parameter in parameters:
|
||||
_assert_exact_object_keys(
|
||||
parameter, {"name", "data_type_ipc"}, context="parameter"
|
||||
)
|
||||
assert isinstance(parameter["data_type_ipc"], str)
|
||||
assert parameter["data_type_ipc"]
|
||||
assert _decode_type_ipc(parameters[0]["data_type_ipc"]) == pa.string()
|
||||
assert _decode_type_ipc(parameters[1]["data_type_ipc"]) == pa.int32()
|
||||
|
||||
output = signature["output"]
|
||||
_assert_exact_object_keys(
|
||||
output, {"data_type_ipc", "nullable"}, context="signature.output"
|
||||
)
|
||||
assert output["nullable"] is True
|
||||
assert _decode_type_ipc(output["data_type_ipc"]) == pa.string()
|
||||
|
||||
implementation = wire["implementation"]
|
||||
_assert_exact_object_keys(
|
||||
implementation,
|
||||
{"kind", "module", "callable", "source", "python", "packages"},
|
||||
context="implementation",
|
||||
)
|
||||
assert implementation["kind"] == "python"
|
||||
assert implementation["module"] == __name__
|
||||
assert implementation["callable"] == "packable_bridge_normalize"
|
||||
assert implementation["source"] == Path(__file__).read_text(encoding="utf-8")
|
||||
assert _SOURCE_MARKER in implementation["source"]
|
||||
assert implementation["python"] == "3.12"
|
||||
assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"]
|
||||
|
||||
capabilities = wire["capabilities"]
|
||||
assert capabilities == [
|
||||
{"kind": "network", "origin": _NETWORK_ORIGIN},
|
||||
{
|
||||
"kind": "secret",
|
||||
"reference": _SECRET_REFERENCE,
|
||||
"environment_variable": _SECRET_ENV,
|
||||
},
|
||||
{"kind": "network", "origin": _NETWORK_ORIGIN_B},
|
||||
]
|
||||
for capability in capabilities:
|
||||
assert "value" not in capability
|
||||
assert "plaintext" not in capability
|
||||
assert "plaintext_secret" not in capability
|
||||
assert "secret_value" not in capability
|
||||
|
||||
|
||||
def test_native_definition_repr_includes_safe_structure_and_redacts_sensitive_text():
|
||||
definition = _build_function_definition(packable_bridge_normalize)
|
||||
rendered = repr(definition)
|
||||
assert "_FunctionDefinition" in rendered or "FunctionDefinition" in rendered
|
||||
assert __name__ in rendered
|
||||
assert "packable_bridge_normalize" in rendered
|
||||
assert "3.12" in rendered
|
||||
_assert_sanitized_text(rendered)
|
||||
|
||||
|
||||
def test_new_function_definition_builder_preserves_normalized_wire():
|
||||
definition = _new_function_definition(**_valid_builder_kwargs())
|
||||
assert type(definition) is _function_definition_type()
|
||||
|
||||
encoded_a = _json_bytes(definition)
|
||||
encoded_b = _json_bytes(definition)
|
||||
assert encoded_a == encoded_b
|
||||
|
||||
wire = json.loads(encoded_a.decode("utf-8"))
|
||||
assert wire["format_version"] == 1
|
||||
assert [parameter["name"] for parameter in wire["signature"]["parameters"]] == [
|
||||
"text",
|
||||
"limit",
|
||||
]
|
||||
assert _decode_type_ipc(wire["signature"]["parameters"][0]["data_type_ipc"]) == (
|
||||
pa.string()
|
||||
)
|
||||
assert _decode_type_ipc(wire["signature"]["parameters"][1]["data_type_ipc"]) == (
|
||||
pa.int32()
|
||||
)
|
||||
assert wire["signature"]["output"]["nullable"] is True
|
||||
assert _decode_type_ipc(wire["signature"]["output"]["data_type_ipc"]) == pa.string()
|
||||
|
||||
implementation = wire["implementation"]
|
||||
assert implementation["kind"] == "python"
|
||||
assert implementation["module"] == "bridge_mod"
|
||||
assert implementation["callable"] == "normalize"
|
||||
assert implementation["source"] == _valid_builder_kwargs()["source"]
|
||||
assert _SOURCE_MARKER in implementation["source"]
|
||||
assert implementation["python"] == "3.12"
|
||||
assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"]
|
||||
assert wire["capabilities"] == [
|
||||
{"kind": "network", "origin": _NETWORK_ORIGIN},
|
||||
{
|
||||
"kind": "secret",
|
||||
"reference": _SECRET_REFERENCE,
|
||||
"environment_variable": _SECRET_ENV,
|
||||
},
|
||||
{"kind": "network", "origin": _NETWORK_ORIGIN_B},
|
||||
]
|
||||
_assert_forbidden_keys_absent(wire, context="definition")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides",),
|
||||
[
|
||||
({"parameters": [("text", pa.string()), ("text", pa.int32())]},),
|
||||
({"parameters": [("", pa.string())]},),
|
||||
({"module": ""},),
|
||||
({"callable_name": ""},),
|
||||
({"source": ""},),
|
||||
({"python": ""},),
|
||||
({"packages": ["pkg-a==1", ""]},),
|
||||
({"packages": ["pkg-a==1", "pkg-a==1"]},),
|
||||
({"capabilities": [("filesystem", _NETWORK_ORIGIN, None)]},),
|
||||
({"capabilities": [("network", _NETWORK_ORIGIN, _SECRET_ENV)]},),
|
||||
({"capabilities": [("secret", _SECRET_REFERENCE, None)]},),
|
||||
({"capabilities": [("secret", _SECRET_REFERENCE, "")]},),
|
||||
({"capabilities": [("network", "", None)]},),
|
||||
({"capabilities": [("secret", "", _SECRET_ENV)]},),
|
||||
],
|
||||
)
|
||||
def test_new_function_definition_strict_validation_rejections(overrides):
|
||||
kwargs = _valid_builder_kwargs(**overrides)
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_new_function_definition(**kwargs)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
|
||||
def test_new_function_definition_validation_does_not_echo_secret_or_source_marker():
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_new_function_definition(**_valid_builder_kwargs(module=""))
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_new_function_definition(
|
||||
**_valid_builder_kwargs(packages=["pkg-a==1", "pkg-a==1"])
|
||||
)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_new_function_definition(
|
||||
**_valid_builder_kwargs(
|
||||
capabilities=[("secret", _SECRET_REFERENCE, None)],
|
||||
)
|
||||
)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides",),
|
||||
[
|
||||
({"parameters": [("text", "not-a-datatype")]},),
|
||||
({"parameters": [(123, pa.string())]},),
|
||||
({"output_type": "not-a-datatype"},),
|
||||
({"output_type": None},),
|
||||
({"output_nullable": "yes"},),
|
||||
({"packages": "pkg-a==1"},),
|
||||
({"capabilities": "network"},),
|
||||
({"capabilities": [("network", _NETWORK_ORIGIN)]},),
|
||||
({"capabilities": [("network", _NETWORK_ORIGIN, None, "extra")]},),
|
||||
],
|
||||
)
|
||||
def test_new_function_definition_wrong_pyarrow_and_shape_values_fail_closed(overrides):
|
||||
kwargs = _valid_builder_kwargs(**overrides)
|
||||
with pytest.raises((TypeError, ValueError)) as exc_info:
|
||||
_new_function_definition(**kwargs)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
|
||||
class _HostileRaisingIterable:
|
||||
def __iter__(self):
|
||||
raise RuntimeError(f"{_SECRET_REFERENCE} {_SOURCE_MARKER}")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides",),
|
||||
[
|
||||
({"parameters": _HostileRaisingIterable()},),
|
||||
({"packages": _HostileRaisingIterable()},),
|
||||
({"capabilities": _HostileRaisingIterable()},),
|
||||
(
|
||||
{
|
||||
"capabilities": [
|
||||
("network", _NETWORK_ORIGIN, None),
|
||||
_HostileRaisingIterable(),
|
||||
("network", _NETWORK_ORIGIN_B, None),
|
||||
]
|
||||
},
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_new_function_definition_hostile_iterable_iter_raises_fail_closed(overrides):
|
||||
kwargs = _valid_builder_kwargs(**overrides)
|
||||
with pytest.raises((TypeError, ValueError)) as exc_info:
|
||||
_new_function_definition(**kwargs)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=False,
|
||||
)
|
||||
def packable_bridge_capability_exact_type(x):
|
||||
return x + 1
|
||||
|
||||
|
||||
def test_build_function_definition_rejects_forged_function_capability_subclass():
|
||||
marker = f"{_SECRET_REFERENCE} {_SOURCE_MARKER}"
|
||||
|
||||
class _HostileFunctionCapability(FunctionCapability):
|
||||
@property
|
||||
def kind(self) -> str:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def origin(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def reference(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def environment_variable(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
hostile = object.__new__(_HostileFunctionCapability)
|
||||
assert isinstance(hostile, FunctionCapability)
|
||||
assert type(hostile) is not FunctionCapability
|
||||
|
||||
config_attr = _udf_mod._CONFIG_ATTR
|
||||
original = getattr(packable_bridge_capability_exact_type, config_attr)
|
||||
forged = _udf_mod._UdfConfig(
|
||||
inputs=original.inputs,
|
||||
output=original.output,
|
||||
output_nullable=original.output_nullable,
|
||||
python=original.python,
|
||||
packages=original.packages,
|
||||
capabilities=(hostile,),
|
||||
)
|
||||
setattr(packable_bridge_capability_exact_type, config_attr, forged)
|
||||
try:
|
||||
with pytest.raises((TypeError, ValueError)) as exc_info:
|
||||
_build_function_definition(packable_bridge_capability_exact_type)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
assert marker not in str(exc_info.value)
|
||||
assert marker not in repr(exc_info.value)
|
||||
finally:
|
||||
setattr(packable_bridge_capability_exact_type, config_attr, original)
|
||||
@@ -0,0 +1,486 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for private UDF packaging validation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import inspect
|
||||
import json
|
||||
import sys
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
from lancedb import Function, Job, udf
|
||||
from lancedb._udf import _get_udf_config, _package_udf
|
||||
|
||||
_BODY_MARKER = "packaging body marker unique-xyz"
|
||||
_AMBIENT_SECRET = "ambient-secret-value-xyz"
|
||||
_BUILTIN_SHADOW_SECRET = "builtin-shadow-secret-xyz"
|
||||
_SOURCE_MISMATCH_SECRET = "source-mismatch-secret-xyz"
|
||||
_INVALID_UTF8_SECRET = "invalid-utf8-secret-xyz"
|
||||
|
||||
_OVERDESIGN_ATTRS = (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_handling",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"error_policy",
|
||||
"FunctionVersion",
|
||||
"function_version",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
"id",
|
||||
"function_id",
|
||||
"job",
|
||||
"job_id",
|
||||
"registration",
|
||||
"catalog",
|
||||
"retry_key",
|
||||
"source_path",
|
||||
"path",
|
||||
"function",
|
||||
)
|
||||
|
||||
_PACKAGING_CONSTANT = 41
|
||||
|
||||
|
||||
def _packaging_helper(value: int) -> int:
|
||||
return value + _PACKAGING_CONSTANT
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=False,
|
||||
)
|
||||
def packable_add(x):
|
||||
"""packaging body marker unique-xyz."""
|
||||
return _packaging_helper(x) + len(json.dumps({"k": 1}))
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def packable_kwonly(x, *, y=2):
|
||||
return x + y
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def packable_rebind_target(x):
|
||||
return x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def uses_injected_ambient(x):
|
||||
return x + len(INJECTED_AMBIENT_GLOBAL) # noqa: F821
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def uses_shadowed_builtin_len(x):
|
||||
return x + len((1, 2, 3))
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def mismatch_names(left, right):
|
||||
return left + right
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"y": pa.int32(), "x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def mismatch_order(x, y):
|
||||
return x + y
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def positional_only(x, /, y):
|
||||
return x + y
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def varargs_fn(x, *args):
|
||||
return x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def kwargs_fn(x, **kwargs):
|
||||
return x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
async def async_fn(x):
|
||||
return x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
async def async_gen_fn(x):
|
||||
yield x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def generator_fn(x):
|
||||
yield x
|
||||
|
||||
|
||||
def _assert_sanitized_text(*parts: object, secret: str = _AMBIENT_SECRET) -> None:
|
||||
combined = "\n".join(str(part) for part in parts)
|
||||
lowered = combined.lower()
|
||||
assert _BODY_MARKER.lower() not in lowered
|
||||
assert secret.lower() not in lowered
|
||||
assert str(Path(__file__).resolve()).lower() not in lowered
|
||||
assert Path(__file__).resolve().as_posix().lower() not in lowered
|
||||
|
||||
|
||||
def _assert_packaging_rejection(exc_info, *, secret: str = _AMBIENT_SECRET) -> None:
|
||||
_assert_sanitized_text(exc_info.value, repr(exc_info.value), secret=secret)
|
||||
assert exc_info.value.__cause__ is None
|
||||
assert exc_info.value.__context__ is None
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _temporary_imported_module(
|
||||
directory: Path, module_name: str, source: str
|
||||
) -> Iterator[tuple[Path, object]]:
|
||||
path = directory / f"{module_name}.py"
|
||||
path.write_text(source, encoding="utf-8")
|
||||
inserted = str(directory)
|
||||
sys.path.insert(0, inserted)
|
||||
try:
|
||||
sys.modules.pop(module_name, None)
|
||||
module = importlib.import_module(module_name)
|
||||
yield path, module
|
||||
finally:
|
||||
sys.modules.pop(module_name, None)
|
||||
try:
|
||||
sys.path.remove(inserted)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
|
||||
def _temp_udf_module_source(*, body: str, secret: str | None = None) -> str:
|
||||
secret_line = f"_SECRET = {secret!r}\n" if secret is not None else ""
|
||||
return (
|
||||
"import pyarrow as pa\n"
|
||||
"from lancedb import udf\n"
|
||||
f"{secret_line}\n"
|
||||
"@udf(\n"
|
||||
' inputs={"x": pa.int32()},\n'
|
||||
" output=pa.int32(),\n"
|
||||
' python="3.12",\n'
|
||||
")\n"
|
||||
"def temp_pack_target(x):\n"
|
||||
f" {body}\n"
|
||||
)
|
||||
|
||||
|
||||
def test_package_udf_success_snapshot_source_module_callable_config_and_repr():
|
||||
packaged = _package_udf(packable_add)
|
||||
source = Path(__file__).read_text(encoding="utf-8")
|
||||
|
||||
assert packaged.source == source
|
||||
assert packaged.module == __name__
|
||||
assert packaged.module != "__main__"
|
||||
assert packaged.callable_name == "packable_add"
|
||||
assert packable_add.__qualname__ == "packable_add"
|
||||
assert packaged.config is _get_udf_config(packable_add)
|
||||
assert packaged.config.inputs == (("x", pa.int32()),)
|
||||
assert packaged.config.output == pa.int64()
|
||||
assert packaged.config.output_nullable is False
|
||||
assert packaged.config.python == "3.12"
|
||||
assert packaged.config.packages == ("pkg-a==1",)
|
||||
|
||||
for attr in ("source", "module", "callable_name", "config"):
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(packaged, attr, None)
|
||||
|
||||
text = repr(packaged)
|
||||
_assert_sanitized_text(text)
|
||||
assert _BODY_MARKER not in text
|
||||
|
||||
|
||||
def test_package_udf_allows_source_bound_import_constant_and_helper():
|
||||
packaged = _package_udf(packable_add)
|
||||
assert packaged.callable_name == "packable_add"
|
||||
assert "import json" in packaged.source
|
||||
assert "_PACKAGING_CONSTANT" in packaged.source
|
||||
assert "_packaging_helper" in packaged.source
|
||||
assert packable_add(1) == _packaging_helper(1) + len(json.dumps({"k": 1}))
|
||||
|
||||
|
||||
def test_package_udf_accepts_positional_or_keyword_and_keyword_only_defaults():
|
||||
packaged = _package_udf(packable_kwonly)
|
||||
assert packaged.callable_name == "packable_kwonly"
|
||||
assert packaged.config.inputs == (("x", pa.int32()), ("y", pa.int32()))
|
||||
assert str(inspect.signature(packable_kwonly)) == "(x, *, y=2)"
|
||||
assert packable_kwonly(3) == 5
|
||||
assert packable_kwonly(3, y=7) == 10
|
||||
|
||||
|
||||
def test_package_udf_rejects_lambda_and_closure():
|
||||
lam = udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(lambda n: n + 1)
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(lam)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
ambient = _AMBIENT_SECRET
|
||||
|
||||
def factory(offset):
|
||||
@udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def closed(n):
|
||||
return n + offset + len(ambient)
|
||||
|
||||
return closed
|
||||
|
||||
closed = factory(10)
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(closed)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
def outer():
|
||||
total = 0
|
||||
|
||||
@udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def nested(n):
|
||||
nonlocal total
|
||||
total += n
|
||||
return total
|
||||
|
||||
return nested
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(outer())
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_signature_mismatches_and_unsupported_parameter_kinds():
|
||||
for target in (
|
||||
mismatch_names,
|
||||
mismatch_order,
|
||||
positional_only,
|
||||
varargs_fn,
|
||||
kwargs_fn,
|
||||
):
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(target)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_async_and_generator_functions():
|
||||
for target in (async_fn, async_gen_fn, generator_fn):
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(target)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_dynamic_exec_source():
|
||||
namespace: dict[str, object] = {}
|
||||
exec(
|
||||
"def dynamic_pack_target(x):\n return x + 1\n",
|
||||
namespace,
|
||||
)
|
||||
dynamic = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(namespace["dynamic_pack_target"])
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(dynamic)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_undecorated_and_wrong_input_types():
|
||||
def plain(x):
|
||||
return x
|
||||
|
||||
with pytest.raises(TypeError) as exc_info:
|
||||
_package_udf(plain)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
with pytest.raises(TypeError) as exc_info:
|
||||
_package_udf(object())
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
with pytest.raises(TypeError) as exc_info:
|
||||
_package_udf(42)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_rebound_module_attribute():
|
||||
module = sys.modules[__name__]
|
||||
original = module.packable_rebind_target
|
||||
replacement = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(lambda x: x)
|
||||
module.packable_rebind_target = replacement
|
||||
try:
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(original)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
finally:
|
||||
module.packable_rebind_target = original
|
||||
|
||||
|
||||
def test_package_udf_rejects_injected_ambient_global():
|
||||
module = sys.modules[__name__]
|
||||
secret = _AMBIENT_SECRET
|
||||
module.INJECTED_AMBIENT_GLOBAL = secret
|
||||
try:
|
||||
assert uses_injected_ambient(3) == 3 + len(secret)
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(uses_injected_ambient)
|
||||
_assert_packaging_rejection(exc_info, secret=secret)
|
||||
finally:
|
||||
delattr(module, "INJECTED_AMBIENT_GLOBAL")
|
||||
|
||||
|
||||
def test_package_udf_rejects_builtin_shadow_injection():
|
||||
module = sys.modules[__name__]
|
||||
secret = _BUILTIN_SHADOW_SECRET
|
||||
assert not hasattr(module, "len")
|
||||
module.len = secret
|
||||
try:
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(uses_shadowed_builtin_len)
|
||||
_assert_packaging_rejection(exc_info, secret=secret)
|
||||
finally:
|
||||
delattr(module, "len")
|
||||
|
||||
|
||||
def test_package_udf_rejects_loaded_code_source_mismatch(tmp_path: Path):
|
||||
secret = _SOURCE_MISMATCH_SECRET
|
||||
module_name = "udf_pkg_source_mismatch_mod"
|
||||
original = _temp_udf_module_source(body="return x + 1")
|
||||
replacement = _temp_udf_module_source(
|
||||
body=f"return x + 99 # {secret}",
|
||||
secret=secret,
|
||||
)
|
||||
with _temporary_imported_module(tmp_path, module_name, original) as (
|
||||
path,
|
||||
module,
|
||||
):
|
||||
target = module.temp_pack_target
|
||||
assert target(1) == 2
|
||||
path.write_text(replacement, encoding="utf-8")
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(target)
|
||||
_assert_packaging_rejection(exc_info, secret=secret)
|
||||
err_text = f"{exc_info.value}\n{exc_info.value!r}"
|
||||
assert str(path.resolve()) not in err_text
|
||||
assert path.resolve().as_posix() not in err_text
|
||||
|
||||
|
||||
def test_package_udf_rejects_invalid_utf8_after_import(tmp_path: Path):
|
||||
secret = _INVALID_UTF8_SECRET
|
||||
module_name = "udf_pkg_invalid_utf8_mod"
|
||||
original = _temp_udf_module_source(body="return x + 1")
|
||||
with _temporary_imported_module(tmp_path, module_name, original) as (
|
||||
path,
|
||||
module,
|
||||
):
|
||||
target = module.temp_pack_target
|
||||
assert target(1) == 2
|
||||
path.write_bytes(secret.encode("utf-8") + b"\xff\xfe invalid-bytes")
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(target)
|
||||
assert type(exc_info.value) is ValueError
|
||||
assert exc_info.value.__cause__ is None
|
||||
assert exc_info.value.__context__ is None
|
||||
err_text = f"{exc_info.value}\n{exc_info.value!r}"
|
||||
assert secret not in err_text
|
||||
assert "b'" not in err_text
|
||||
assert r"\xff" not in err_text
|
||||
assert str(path.resolve()) not in err_text
|
||||
assert path.resolve().as_posix() not in err_text
|
||||
|
||||
|
||||
def test_package_udf_snapshot_has_no_durable_overdesign_and_is_not_function_or_job():
|
||||
packaged = _package_udf(packable_add)
|
||||
assert not isinstance(packaged, Function)
|
||||
assert not isinstance(packaged, Job)
|
||||
for attr in _OVERDESIGN_ATTRS:
|
||||
assert not hasattr(packaged, attr)
|
||||
|
||||
text = repr(packaged).lower()
|
||||
for token in (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"functionversion",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
"retry_key",
|
||||
):
|
||||
assert token not in text
|
||||
_assert_sanitized_text(text)
|
||||
@@ -2308,34 +2308,226 @@ def test_remote_connection_jobs_surface():
|
||||
job.wait(timeout=timedelta(seconds=5))
|
||||
|
||||
|
||||
def test_remote_add_bases_posts_the_bases_array():
|
||||
captured_body = {}
|
||||
# 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):
|
||||
if request.path == "/v1/table/test/describe/":
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(BLOB_DESCRIBE_RESPONSE).encode())
|
||||
elif request.path == "/v1/table/test/bases/":
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
captured_body.update(json.loads(request.rfile.read(content_len)))
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(b'{"version": 2}')
|
||||
else:
|
||||
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())
|
||||
|
||||
with mock_lancedb_connection(handler) as db:
|
||||
table = db.open_table("test")
|
||||
table.add_bases(lancedb.TableBase(path="s3://bucket/media/"))
|
||||
return handler
|
||||
|
||||
assert captured_body["bases"] == [
|
||||
{
|
||||
"path": "s3://bucket/media/",
|
||||
"isDatasetRoot": False,
|
||||
}
|
||||
]
|
||||
|
||||
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,65 +3854,3 @@ async def test_async_search_runs_embedding_on_dedicated_executor(
|
||||
assert all(name.startswith("lancedb-embedding") for name in captured_threads), (
|
||||
f"embedding ran off the dedicated executor: {captured_threads}"
|
||||
)
|
||||
|
||||
|
||||
def test_computed_column_declare_and_refresh(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("computed", [{"x": 1}, {"x": 2}])
|
||||
|
||||
table.add_columns(computed={"doubled": "x * 2"})
|
||||
assert table.to_arrow()["doubled"].to_pylist() == [None, None]
|
||||
|
||||
result = table.refresh_column("doubled")
|
||||
assert result.rows_filled == 2
|
||||
assert sorted(table.to_arrow()["doubled"].to_pylist()) == [2, 4]
|
||||
|
||||
table.add([{"x": 5}])
|
||||
assert table.refresh_column("doubled").rows_filled == 1
|
||||
assert sorted(table.to_arrow()["doubled"].to_pylist()) == [2, 4, 10]
|
||||
|
||||
|
||||
def test_computed_column_rejects_transforms_and_computed_together(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("computed_mixed", [{"x": 1}])
|
||||
with pytest.raises(ValueError):
|
||||
table.add_columns({"a": "x + 1"}, computed={"b": "x * 2"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_computed_column_async(tmp_path):
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
table = await db.create_table("computed_async", [{"x": 3}])
|
||||
|
||||
await table.add_columns(computed={"tripled": "x * 3"})
|
||||
await table.refresh_column("tripled")
|
||||
|
||||
assert (await table.to_arrow())["tripled"].to_pylist() == [9]
|
||||
|
||||
|
||||
def test_refresh_column_async_returns_job(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("computed_job", [{"x": 1}, {"x": 2}])
|
||||
table.add_columns(computed={"doubled": "x * 2"})
|
||||
|
||||
job = table.refresh_column_async("doubled")
|
||||
assert job.id is None # in-process jobs have no server id
|
||||
job.wait()
|
||||
assert job.status() == "finished"
|
||||
assert sorted(table.to_arrow()["doubled"].to_pylist()) == [2, 4]
|
||||
|
||||
# Bad input raises at the call, not through the job.
|
||||
with pytest.raises(Exception, match="not a computed column"):
|
||||
table.refresh_column_async("x")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_column_async_job_async_table(tmp_path):
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
table = await db.create_table("computed_job_async", [{"x": 3}])
|
||||
await table.add_columns(computed={"tripled": "x * 3"})
|
||||
|
||||
job = await table.refresh_column_async("tripled")
|
||||
await job.wait()
|
||||
assert await job.status() == "finished"
|
||||
assert (await table.to_arrow())["tripled"].to_pylist() == [9]
|
||||
|
||||
+116
-17
@@ -23,6 +23,7 @@ use lancedb::{
|
||||
connection::NamespaceClientPushdownOperation,
|
||||
database::namespace::LanceNamespaceDatabase,
|
||||
database::{CreateTableMode, Database, ReadConsistency},
|
||||
function::{FunctionId, RegisterFunctionJobSpec},
|
||||
};
|
||||
use pyo3::{
|
||||
Bound, FromPyObject, Py, PyAny, PyRef, PyResult, Python,
|
||||
@@ -346,23 +347,6 @@ impl Connection {
|
||||
})
|
||||
}
|
||||
|
||||
#[pyo3(signature = (name, namespace_path=None))]
|
||||
pub fn drop_table_async(
|
||||
self_: PyRef<'_, Self>,
|
||||
name: String,
|
||||
namespace_path: Option<Vec<String>>,
|
||||
) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
let ns_path = namespace_path.unwrap_or_default();
|
||||
future_into_py(self_.py(), async move {
|
||||
inner
|
||||
.drop_table_async(name, &ns_path)
|
||||
.await
|
||||
.infer_error()
|
||||
.map(crate::job::Job::new)
|
||||
})
|
||||
}
|
||||
|
||||
#[pyo3(signature = (namespace_path=None,))]
|
||||
pub fn drop_all_tables(
|
||||
self_: PyRef<'_, Self>,
|
||||
@@ -606,6 +590,121 @@ impl Connection {
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
/// Submit a first-class Function registration job.
|
||||
///
|
||||
/// Accepts the exact private [`crate::function::PyFunctionDefinition`] and
|
||||
/// builds [`RegisterFunctionJobSpec`] with `expected_current_function_id =
|
||||
/// None` (create-if-absent). Does not JSON round-trip the definition.
|
||||
pub fn _register_function<'py>(
|
||||
self_: PyRef<'py, Self>,
|
||||
name: String,
|
||||
definition: Bound<'_, crate::function::PyFunctionDefinition>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
let definition = definition.get().inner().clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let spec = RegisterFunctionJobSpec::try_new(name, definition, None).infer_error()?;
|
||||
let job = inner.register_function(spec).await.infer_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
})
|
||||
}
|
||||
|
||||
/// Submit a first-class Function conditional replace job.
|
||||
///
|
||||
/// Accepts the observed native [`crate::function::Function`] handle and the
|
||||
/// exact private [`crate::function::PyFunctionDefinition`], then builds
|
||||
/// [`RegisterFunctionJobSpec`] with `expected_current_function_id =
|
||||
/// Some(current.id)`. Reads only `current.inner().id().clone()`. Does not
|
||||
/// JSON round-trip the definition.
|
||||
pub fn _replace_function<'py>(
|
||||
self_: PyRef<'py, Self>,
|
||||
name: String,
|
||||
current: Bound<'_, crate::function::Function>,
|
||||
definition: Bound<'_, crate::function::PyFunctionDefinition>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
let definition = definition.get().inner().clone();
|
||||
let current_id = current.get().inner().id().clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let spec = RegisterFunctionJobSpec::try_new(name, definition, Some(current_id))
|
||||
.infer_error()?;
|
||||
let job = inner.register_function(spec).await.infer_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
})
|
||||
}
|
||||
|
||||
/// Look up the Function currently bound to a database-scoped name.
|
||||
///
|
||||
/// Wraps the exact Rust [`lancedb::function::Function`] once. Empty names
|
||||
/// fail as [`PyValueError`] before transport via the Rust connection.
|
||||
pub fn _lookup_function_by_name<'py>(
|
||||
self_: PyRef<'py, Self>,
|
||||
name: String,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let function = inner.lookup_function_by_name(&name).await.infer_error()?;
|
||||
Ok(crate::function::Function::new(function))
|
||||
})
|
||||
}
|
||||
|
||||
/// Look up an immutable Function by exact opaque Function ID string.
|
||||
///
|
||||
/// Constructs [`FunctionId`] with [`FunctionId::try_new`] before dispatch so
|
||||
/// empty IDs fail as [`PyValueError`] before transport.
|
||||
pub fn _lookup_function_by_id<'py>(
|
||||
self_: PyRef<'py, Self>,
|
||||
function_id: String,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let id = FunctionId::try_new(function_id).infer_error()?;
|
||||
let function = inner.lookup_function_by_id(&id).await.infer_error()?;
|
||||
Ok(crate::function::Function::new(function))
|
||||
})
|
||||
}
|
||||
|
||||
/// Conditionally remove a database-scoped Function catalog name.
|
||||
///
|
||||
/// Clones the observed native [`crate::function::Function`] once and
|
||||
/// delegates to Rust [`lancedb::Connection::remove_function_name`]. Empty
|
||||
/// names fail as [`PyValueError`] before transport via the Rust connection.
|
||||
pub fn _remove_function_name<'py>(
|
||||
self_: PyRef<'py, Self>,
|
||||
name: String,
|
||||
current: Bound<'_, crate::function::Function>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
let current = current.get().inner().clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
inner
|
||||
.remove_function_name(&name, ¤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]
|
||||
|
||||
+15
-2
@@ -102,11 +102,14 @@ impl<T> PythonErrorExt<T> for std::result::Result<T, LanceError> {
|
||||
err.setattr(intern!(py, "__cause__"), cause_err)?;
|
||||
Err(PyErr::from_value(err))
|
||||
}),
|
||||
LanceError::JobFailed { .. } => Python::attach(|py| {
|
||||
LanceError::JobFailed { failure, .. } => Python::attach(|py| {
|
||||
let cls = py
|
||||
.import(intern!(py, "lancedb.exceptions"))?
|
||||
.getattr(intern!(py, "JobFailedError"))?;
|
||||
Err(PyErr::from_value(cls.call1((err.to_string(),))?))
|
||||
// Structural projection only: failure.error_code.as_str().
|
||||
// Never infer a code from message, phase, retryable, or source.
|
||||
let error_code = failure.error_code.as_ref().map(|code| code.as_str());
|
||||
Err(PyErr::from_value(cls.call1((err.to_string(), error_code))?))
|
||||
}),
|
||||
LanceError::JobCancelled { .. } => Python::attach(|py| {
|
||||
let cls = py
|
||||
@@ -114,6 +117,16 @@ impl<T> PythonErrorExt<T> for std::result::Result<T, LanceError> {
|
||||
.getattr(intern!(py, "JobCancelledError"))?;
|
||||
Err(PyErr::from_value(cls.call1((err.to_string(),))?))
|
||||
}),
|
||||
LanceError::Function { code, message } => Python::attach(|py| {
|
||||
let cls = py
|
||||
.import(intern!(py, "lancedb.exceptions"))?
|
||||
.getattr(intern!(py, "FunctionError"))?;
|
||||
// Structural projection only: code.as_str() + sanitized message.
|
||||
// Never infer a code from HTTP status or diagnostic text.
|
||||
Err(PyErr::from_value(
|
||||
cls.call1((message.as_str(), code.as_str()))?,
|
||||
))
|
||||
}),
|
||||
_ => self.runtime_error(),
|
||||
},
|
||||
}
|
||||
|
||||
+28
-1
@@ -10,7 +10,7 @@
|
||||
use std::ops::{Add, Div, Mul, Not, Sub};
|
||||
|
||||
use arrow::{datatypes::DataType, pyarrow::PyArrowType};
|
||||
use datafusion_common::ScalarValue;
|
||||
use datafusion_common::{Column, ScalarValue};
|
||||
use lancedb::expr::{
|
||||
DfExpr, col as ldb_col, contains, expr_cast, is_in, lit as df_lit, lower, upper,
|
||||
};
|
||||
@@ -27,6 +27,33 @@ use pyo3::{Bound, PyAny, PyResult, exceptions::PyValueError, prelude::*, pyfunct
|
||||
#[derive(Clone)]
|
||||
pub struct PyExpr(pub DfExpr);
|
||||
|
||||
/// Crate-private inspection result for Function call authoring (FF-028).
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) enum DirectExprView<'a> {
|
||||
/// Direct unqualified DataFusion Column; name is case-sensitive.
|
||||
UnqualifiedColumn(&'a str),
|
||||
/// Direct Literal scalar; Arrow type is owned by the scalar value.
|
||||
Literal(&'a ScalarValue),
|
||||
}
|
||||
|
||||
impl PyExpr {
|
||||
/// Inspect a direct Column/Literal node for Function call authoring.
|
||||
///
|
||||
/// Returns `None` for every other expression shape (arithmetic, cast,
|
||||
/// scalar function, predicate, alias, qualified column, etc.).
|
||||
pub(crate) fn as_direct_column_or_literal(&self) -> Option<DirectExprView<'_>> {
|
||||
match &self.0 {
|
||||
DfExpr::Column(Column {
|
||||
relation: None,
|
||||
name,
|
||||
..
|
||||
}) => Some(DirectExprView::UnqualifiedColumn(name.as_str())),
|
||||
DfExpr::Literal(value, _) => Some(DirectExprView::Literal(value)),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl PyExpr {
|
||||
// ── comparisons ──────────────────────────────────────────────────────────
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+31
-4
@@ -3,6 +3,7 @@
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::function::Function;
|
||||
use crate::runtime::future_into_py;
|
||||
use pyo3::{Bound, PyAny, PyRef, PyResult, pyclass, pymethods};
|
||||
|
||||
@@ -21,6 +22,23 @@ impl Job {
|
||||
}
|
||||
}
|
||||
|
||||
/// Project a Rust [`lancedb::JobResult`] onto the Python success surface.
|
||||
///
|
||||
/// Delegates variant interpretation to [`lancedb::JobResult::into_function`]:
|
||||
/// no nested Function collapses to Python `None`; an exact Function becomes
|
||||
/// the corresponding [`Function`] handle.
|
||||
fn project_wait_result(result: lancedb::JobResult) -> Option<Function> {
|
||||
result.into_function().map(Function::new)
|
||||
}
|
||||
|
||||
/// Project a describe `result` onto Python `Optional[Function]`.
|
||||
///
|
||||
/// Rust `None`, `Some(JobResult::None)`, and JSON null all become Python
|
||||
/// `None`. Only `Some(JobResult::Function)` becomes a [`Function`] handle.
|
||||
fn project_description_result(result: Option<lancedb::JobResult>) -> Option<Function> {
|
||||
result.and_then(project_wait_result)
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl Job {
|
||||
#[getter]
|
||||
@@ -39,8 +57,8 @@ impl Job {
|
||||
pub fn wait(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
inner.wait().await.infer_error()?;
|
||||
Ok(())
|
||||
let result = inner.wait().await.infer_error()?;
|
||||
Ok(project_wait_result(result))
|
||||
})
|
||||
}
|
||||
|
||||
@@ -93,14 +111,16 @@ pub struct JobFailureInfo {
|
||||
phase: Option<String>,
|
||||
message: Option<String>,
|
||||
retryable: Option<bool>,
|
||||
/// Exact wire `error_code` string when Rust decoded one; never inferred.
|
||||
error_code: Option<String>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl JobFailureInfo {
|
||||
fn __repr__(&self) -> String {
|
||||
format!(
|
||||
"JobFailureInfo(phase={:?}, message={:?}, retryable={:?})",
|
||||
self.phase, self.message, self.retryable
|
||||
"JobFailureInfo(phase={:?}, message={:?}, retryable={:?}, error_code={:?})",
|
||||
self.phase, self.message, self.retryable, self.error_code
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -115,6 +135,7 @@ pub struct JobDescription {
|
||||
creation_ms: i64,
|
||||
spec_json: Option<String>,
|
||||
failure: Option<JobFailureInfo>,
|
||||
result: Option<Function>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
@@ -139,7 +160,13 @@ impl From<lancedb::database::JobDescription> for JobDescription {
|
||||
phase: failure.phase,
|
||||
message: failure.message,
|
||||
retryable: failure.retryable,
|
||||
// Structural projection only: exact as_str(); never infer.
|
||||
error_code: failure
|
||||
.error_code
|
||||
.as_ref()
|
||||
.map(|code| code.as_str().to_string()),
|
||||
}),
|
||||
result: project_description_result(description.result),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+9
-3
@@ -16,14 +16,14 @@ use query::{FTSQuery, HybridQuery, Query, VectorQuery};
|
||||
use session::Session;
|
||||
use table::{
|
||||
AddColumnsResult, AddResult, AlterColumnsResult, DeleteResult, DropColumnsResult, FtsToken,
|
||||
LsmWriteSpec, MergeResult, PyBlobFile, RefreshColumnResult, Table, UpdateFieldMetadataResult,
|
||||
UpdateResult,
|
||||
LsmWriteSpec, MergeResult, PyBlobFile, Table, UpdateFieldMetadataResult, UpdateResult,
|
||||
};
|
||||
|
||||
pub mod arrow;
|
||||
pub mod connection;
|
||||
pub mod error;
|
||||
pub mod expr;
|
||||
pub mod function;
|
||||
pub mod header;
|
||||
pub mod index;
|
||||
pub mod job;
|
||||
@@ -46,6 +46,9 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
m.add_class::<Connection>()?;
|
||||
m.add_class::<Session>()?;
|
||||
m.add_class::<Table>()?;
|
||||
m.add_class::<crate::function::Function>()?;
|
||||
m.add_class::<crate::function::PyFunctionDefinition>()?;
|
||||
m.add_class::<crate::function::AuthoredFunctionCall>()?;
|
||||
m.add_class::<crate::job::Job>()?;
|
||||
m.add_class::<crate::job::JobInfo>()?;
|
||||
m.add_class::<crate::job::JobDescription>()?;
|
||||
@@ -58,7 +61,6 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
m.add_class::<VectorQuery>()?;
|
||||
m.add_class::<RecordBatchStream>()?;
|
||||
m.add_class::<AddColumnsResult>()?;
|
||||
m.add_class::<RefreshColumnResult>()?;
|
||||
m.add_class::<AlterColumnsResult>()?;
|
||||
m.add_class::<UpdateFieldMetadataResult>()?;
|
||||
m.add_class::<AddResult>()?;
|
||||
@@ -90,6 +92,10 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
m.add_function(wrap_pyfunction!(expr_col, m)?)?;
|
||||
m.add_function(wrap_pyfunction!(expr_lit, m)?)?;
|
||||
m.add_function(wrap_pyfunction!(expr_func, m)?)?;
|
||||
m.add_function(wrap_pyfunction!(
|
||||
crate::function::_new_function_definition,
|
||||
m
|
||||
)?)?;
|
||||
m.add("__version__", env!("CARGO_PKG_VERSION"))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
+141
-88
@@ -22,12 +22,11 @@ use lancedb::index::scalar::FtsIndexBuilder;
|
||||
use lancedb::table::{
|
||||
AddDataMode, ColumnAlteration, Duration, FieldMetadataUpdate, FtsToken as LanceDbFtsToken,
|
||||
NewColumnTransform, OptimizeAction, OptimizeOptions, Ref, Table as LanceDbTable,
|
||||
TableBase as LanceTableBase,
|
||||
};
|
||||
use lancedb::tokenize as lancedb_tokenize;
|
||||
use pyo3::{
|
||||
Bound, FromPyObject, Py, PyAny, PyRef, PyResult, Python,
|
||||
exceptions::{PyRuntimeError, PyValueError},
|
||||
exceptions::{PyNotImplementedError, PyRuntimeError, PyValueError},
|
||||
pyclass, pyfunction, pymethods,
|
||||
types::{IntoPyDict, PyAnyMethods, PyBytes, PyDict, PyDictMethods, PyList, PyListMethods},
|
||||
};
|
||||
@@ -95,13 +94,6 @@ fn lsm_stats_to_py(py: Python<'_>, stats: &lancedb::table::LsmStats) -> PyResult
|
||||
Ok(out.unbind())
|
||||
}
|
||||
|
||||
#[derive(FromPyObject)]
|
||||
pub(crate) struct PyTableBase {
|
||||
path: String,
|
||||
name: Option<String>,
|
||||
is_dataset_root: bool,
|
||||
}
|
||||
|
||||
#[derive(FromPyObject)]
|
||||
enum PredicateArg {
|
||||
Expr(PyExpr),
|
||||
@@ -423,32 +415,6 @@ pub struct AddColumnsResult {
|
||||
pub version: u64,
|
||||
}
|
||||
|
||||
#[pyclass(get_all, from_py_object)]
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct RefreshColumnResult {
|
||||
pub rows_filled: u64,
|
||||
pub version: u64,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl RefreshColumnResult {
|
||||
pub fn __repr__(&self) -> String {
|
||||
format!(
|
||||
"RefreshColumnResult(rows_filled={}, version={})",
|
||||
self.rows_filled, self.version
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<lancedb::table::RefreshColumnResult> for RefreshColumnResult {
|
||||
fn from(result: lancedb::table::RefreshColumnResult) -> Self {
|
||||
Self {
|
||||
rows_filled: result.rows_filled,
|
||||
version: result.version,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl AddColumnsResult {
|
||||
pub fn __repr__(&self) -> String {
|
||||
@@ -964,6 +930,146 @@ impl Table {
|
||||
})
|
||||
}
|
||||
|
||||
/// Hidden bridge: bind an authored Function call once and submit create.
|
||||
///
|
||||
/// Private native path for Python ``table.add_generated_column``. Rejects an
|
||||
/// empty ``column_name`` before reading the table handle. Does not expose
|
||||
/// source version, stable field IDs, the operation spec, or request envelope.
|
||||
#[doc(hidden)]
|
||||
pub fn _add_generated_column<'a>(
|
||||
self_: PyRef<'a, Self>,
|
||||
column_name: String,
|
||||
call: Bound<'_, crate::function::AuthoredFunctionCall>,
|
||||
) -> PyResult<Bound<'a, PyAny>> {
|
||||
if column_name.is_empty() {
|
||||
return Err(PyValueError::new_err("column_name must be non-empty"));
|
||||
}
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
let authored = call.get().clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let (source_table_version, bound_call) =
|
||||
authored.bind_to_table(&inner).await.infer_error()?;
|
||||
let spec = lancedb::function::CreateGeneratedColumnJobSpec::try_new(
|
||||
column_name,
|
||||
authored.function(),
|
||||
bound_call,
|
||||
)
|
||||
.infer_error()?;
|
||||
let job = inner
|
||||
.submit_create_generated_column(source_table_version, spec)
|
||||
.await
|
||||
.infer_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
})
|
||||
}
|
||||
|
||||
/// Hidden bridge: project generated-column completeness for one column name.
|
||||
///
|
||||
/// Private native path for Python ``table.generated_column_status``. Rejects
|
||||
/// an empty ``column_name`` before reading the table handle. Maps only the
|
||||
/// known Rust status variants to ``"complete"`` / ``"incomplete"``.
|
||||
#[doc(hidden)]
|
||||
pub fn _generated_column_status<'a>(
|
||||
self_: PyRef<'a, Self>,
|
||||
column_name: String,
|
||||
) -> PyResult<Bound<'a, PyAny>> {
|
||||
if column_name.is_empty() {
|
||||
return Err(PyValueError::new_err("column_name must be non-empty"));
|
||||
}
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let status = inner
|
||||
.generated_column_status(column_name)
|
||||
.await
|
||||
.infer_error()?;
|
||||
match status {
|
||||
lancedb::function::GeneratedColumnStatus::Complete => Ok("complete"),
|
||||
lancedb::function::GeneratedColumnStatus::Incomplete => Ok("incomplete"),
|
||||
_ => Err(PyNotImplementedError::new_err(
|
||||
"unsupported generated column status",
|
||||
)),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/// Hidden bridge: load exact definition, resolve Function by ID, submit refresh.
|
||||
///
|
||||
/// Private native path for Python ``table.refresh_generated_column``. Rejects
|
||||
/// an empty ``column_name`` before reading the table handle. Does not expose
|
||||
/// source version, Function, field IDs, epochs, specs, or request envelope.
|
||||
#[doc(hidden)]
|
||||
pub fn _refresh_generated_column<'a>(
|
||||
self_: PyRef<'a, Self>,
|
||||
column_name: String,
|
||||
) -> PyResult<Bound<'a, PyAny>> {
|
||||
if column_name.is_empty() {
|
||||
return Err(PyValueError::new_err("column_name must be non-empty"));
|
||||
}
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let (source_table_version, definition) = inner
|
||||
.generated_column_definition_snapshot(column_name)
|
||||
.await
|
||||
.infer_error()?;
|
||||
let function_id = definition.function_call().function_id().clone();
|
||||
let function = inner
|
||||
.resolve_function_for_generated_column(&function_id)
|
||||
.await
|
||||
.infer_error()?;
|
||||
let spec =
|
||||
lancedb::function::RefreshGeneratedColumnJobSpec::try_new(&function, definition)
|
||||
.infer_error()?;
|
||||
let job = inner
|
||||
.submit_refresh_generated_column(source_table_version, spec)
|
||||
.await
|
||||
.infer_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
})
|
||||
}
|
||||
|
||||
/// Hidden bridge: one binding snapshot, bind new call, submit change.
|
||||
///
|
||||
/// Private native path for Python ``table.alter_generated_column``. Rejects
|
||||
/// an empty ``column_name`` before reading the table handle. Fetches exactly
|
||||
/// one binding snapshot, loads the expected definition from that same
|
||||
/// object, binds the authored call against it, and submits change. Does not
|
||||
/// expose source version, Function handles, field IDs, epochs, specs, or
|
||||
/// request envelope.
|
||||
#[doc(hidden)]
|
||||
pub fn _alter_generated_column<'a>(
|
||||
self_: PyRef<'a, Self>,
|
||||
column_name: String,
|
||||
new_call: Bound<'_, crate::function::AuthoredFunctionCall>,
|
||||
) -> PyResult<Bound<'a, PyAny>> {
|
||||
if column_name.is_empty() {
|
||||
return Err(PyValueError::new_err("column_name must be non-empty"));
|
||||
}
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
let authored = new_call.get().clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let snapshot = inner
|
||||
.generated_column_binding_snapshot()
|
||||
.await
|
||||
.infer_error()?;
|
||||
let expected_definition = snapshot
|
||||
.generated_column_definition(&column_name)
|
||||
.infer_error()?;
|
||||
let (source_table_version, bound_new_call) =
|
||||
authored.bind_against_snapshot(&snapshot).infer_error()?;
|
||||
let spec = lancedb::function::ChangeGeneratedColumnJobSpec::try_new(
|
||||
expected_definition,
|
||||
authored.function(),
|
||||
bound_new_call,
|
||||
)
|
||||
.infer_error()?;
|
||||
let job = inner
|
||||
.submit_change_generated_column(source_table_version, spec)
|
||||
.await
|
||||
.infer_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
})
|
||||
}
|
||||
|
||||
pub fn drop_index(self_: PyRef<'_, Self>, index_name: String) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
@@ -1246,25 +1352,6 @@ impl Table {
|
||||
})
|
||||
}
|
||||
|
||||
#[pyo3(signature = (bases))]
|
||||
pub fn add_bases(
|
||||
self_: PyRef<'_, Self>,
|
||||
bases: Vec<PyTableBase>,
|
||||
) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
let bases: Vec<LanceTableBase> = bases
|
||||
.into_iter()
|
||||
.map(|base| LanceTableBase {
|
||||
path: base.path,
|
||||
name: base.name,
|
||||
is_dataset_root: base.is_dataset_root,
|
||||
})
|
||||
.collect();
|
||||
future_into_py(self_.py(), async move {
|
||||
inner.add_bases(bases).await.infer_error()
|
||||
})
|
||||
}
|
||||
|
||||
/// Read blob bytes for `row_ids` from blob v2 column `column`.
|
||||
#[pyo3(signature = (column, row_ids))]
|
||||
pub fn fetch_blobs(
|
||||
@@ -1563,40 +1650,6 @@ impl Table {
|
||||
})
|
||||
}
|
||||
|
||||
pub fn add_computed_columns(
|
||||
self_: PyRef<'_, Self>,
|
||||
columns: Vec<(String, String)>,
|
||||
) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let mut builder = inner.add_columns();
|
||||
for (name, expression) in columns {
|
||||
builder = builder.computed(name, expression);
|
||||
}
|
||||
let result = builder.execute().await.infer_error()?;
|
||||
Ok(AddColumnsResult::from(result))
|
||||
})
|
||||
}
|
||||
|
||||
pub fn refresh_column(self_: PyRef<'_, Self>, column: String) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let result = inner.refresh_column(column).await.infer_error()?;
|
||||
Ok(RefreshColumnResult::from(result))
|
||||
})
|
||||
}
|
||||
|
||||
pub fn refresh_column_async(
|
||||
self_: PyRef<'_, Self>,
|
||||
column: String,
|
||||
) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let job = inner.refresh_column_async(column).await.infer_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
})
|
||||
}
|
||||
|
||||
pub fn add_columns_with_schema(
|
||||
self_: PyRef<'_, Self>,
|
||||
schema: PyArrowType<Schema>,
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "lancedb"
|
||||
version = "0.38.0-beta.0"
|
||||
version = "0.37.1-beta.1"
|
||||
edition.workspace = true
|
||||
description = "LanceDB: A serverless, low-latency vector database for AI applications"
|
||||
license.workspace = true
|
||||
@@ -12,6 +12,7 @@ rust-version.workspace = true
|
||||
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
|
||||
[dependencies]
|
||||
ahash = { workspace = true }
|
||||
base64 = "0.22"
|
||||
arrow = { workspace = true }
|
||||
arrow-array = { workspace = true }
|
||||
arrow-buffer = { workspace = true }
|
||||
@@ -188,9 +189,6 @@ required-features = ["bedrock"]
|
||||
[[example]]
|
||||
name = "bench_streaming_dataloader"
|
||||
|
||||
[[example]]
|
||||
name = "bench_open_missing_table"
|
||||
|
||||
[[example]]
|
||||
name = "simple"
|
||||
|
||||
|
||||
@@ -1,150 +0,0 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
// Release benchmark for opening a missing table as sibling-table cardinality grows.
|
||||
//
|
||||
// The fixture uses real `.lance` directories and marker files. Fixture creation is
|
||||
// outside the timed section. Defaults intentionally cover 1k, 10k, and 100k siblings
|
||||
// with 10 warmups and 100 distinct missing-table opens per scale:
|
||||
//
|
||||
// ```text
|
||||
// cargo run --release -p lancedb --example bench_open_missing_table
|
||||
// ```
|
||||
//
|
||||
// `BENCH_SIBLINGS`, `BENCH_WARMUPS`, and `BENCH_TRIALS` override those defaults.
|
||||
// Reduced settings are useful only as a smoke test. Performance comparisons require
|
||||
// the same machine, filesystem, fixture sizes, settings, lockfile, and alternating
|
||||
// baseline/candidate execution order.
|
||||
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use anyhow::{Context, Result, bail};
|
||||
use lancedb::connection::Connection;
|
||||
use lancedb::{Error, connect};
|
||||
use object_store::ObjectStoreExt as _;
|
||||
use object_store::path::Path;
|
||||
|
||||
const MAX_SIBLINGS: usize = 1_000_000;
|
||||
const MAX_WARMUPS: usize = 10_000;
|
||||
const MAX_TRIALS: usize = 100_000;
|
||||
|
||||
fn env_usize(key: &str, default: usize, max: usize) -> Result<usize> {
|
||||
let value = match std::env::var(key) {
|
||||
Ok(value) => value
|
||||
.parse()
|
||||
.with_context(|| format!("invalid {key} value: {value}"))?,
|
||||
Err(std::env::VarError::NotPresent) => default,
|
||||
Err(error) => return Err(error).with_context(|| format!("reading {key}")),
|
||||
};
|
||||
if value == 0 || value > max {
|
||||
bail!("{key} must be between 1 and {max}");
|
||||
}
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
fn sibling_counts() -> Result<Vec<usize>> {
|
||||
let raw = std::env::var("BENCH_SIBLINGS").unwrap_or_else(|_| "1000,10000,100000".into());
|
||||
let mut counts = raw
|
||||
.split(',')
|
||||
.map(|value| {
|
||||
value
|
||||
.trim()
|
||||
.parse::<usize>()
|
||||
.with_context(|| format!("invalid BENCH_SIBLINGS value: {value}"))
|
||||
})
|
||||
.collect::<Result<Vec<_>>>()?;
|
||||
counts.sort_unstable();
|
||||
counts.dedup();
|
||||
if counts.is_empty() || counts[0] == 0 || counts[counts.len() - 1] > MAX_SIBLINGS {
|
||||
bail!("BENCH_SIBLINGS values must be between 1 and {MAX_SIBLINGS}");
|
||||
}
|
||||
Ok(counts)
|
||||
}
|
||||
|
||||
async fn add_siblings(
|
||||
store: &object_store::local::LocalFileSystem,
|
||||
start: usize,
|
||||
end: usize,
|
||||
) -> Result<()> {
|
||||
for index in start..end {
|
||||
let marker = Path::from(format!("sibling_{index:06}.lance/_marker"));
|
||||
store
|
||||
.put(&marker, bytes::Bytes::new().into())
|
||||
.await
|
||||
.with_context(|| format!("creating benchmark marker {marker}"))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn time_missing_open(db: &Connection, name: &str) -> Result<Duration> {
|
||||
let started = Instant::now();
|
||||
let result = db.open_table(name).execute().await;
|
||||
let elapsed = started.elapsed();
|
||||
match result {
|
||||
Err(Error::TableNotFound { .. }) => Ok(elapsed),
|
||||
Err(error) => bail!("expected TableNotFound for {name}, got {error:?}"),
|
||||
Ok(_) => bail!("benchmark missing-table name unexpectedly exists: {name}"),
|
||||
}
|
||||
}
|
||||
|
||||
fn percentile(sorted: &[Duration], percentile: usize) -> Duration {
|
||||
let rank = (sorted.len() * percentile).div_ceil(100).saturating_sub(1);
|
||||
sorted[rank]
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
let counts = sibling_counts()?;
|
||||
let warmups = env_usize("BENCH_WARMUPS", 10, MAX_WARMUPS)?;
|
||||
let trials = env_usize("BENCH_TRIALS", 100, MAX_TRIALS)?;
|
||||
|
||||
let fixture = tempfile::tempdir().context("creating benchmark fixture")?;
|
||||
let database_path = fixture.path();
|
||||
let fixture_store = object_store::local::LocalFileSystem::new_with_prefix(database_path)
|
||||
.context("creating benchmark object store")?;
|
||||
let db = connect(database_path.to_str().context("non-UTF-8 fixture path")?)
|
||||
.execute()
|
||||
.await?;
|
||||
|
||||
println!(
|
||||
"config: siblings={counts:?} warmups={warmups} trials={trials} profile={} os={} arch={}",
|
||||
if cfg!(debug_assertions) {
|
||||
"debug"
|
||||
} else {
|
||||
"release"
|
||||
},
|
||||
std::env::consts::OS,
|
||||
std::env::consts::ARCH,
|
||||
);
|
||||
println!("lower is better; fixture setup and teardown are excluded");
|
||||
println!("| siblings | samples | p50 | p95 | max |");
|
||||
println!("| ---: | ---: | ---: | ---: | ---: |");
|
||||
|
||||
let mut created = 0;
|
||||
for sibling_count in counts {
|
||||
add_siblings(&fixture_store, created, sibling_count).await?;
|
||||
created = sibling_count;
|
||||
|
||||
for index in 0..warmups {
|
||||
let name = format!("__missing_warmup_{sibling_count}_{index}");
|
||||
let _ = time_missing_open(&db, &name).await?;
|
||||
}
|
||||
|
||||
let mut samples = Vec::with_capacity(trials);
|
||||
for index in 0..trials {
|
||||
let name = format!("__missing_trial_{sibling_count}_{index}");
|
||||
samples.push(time_missing_open(&db, &name).await?);
|
||||
}
|
||||
samples.sort_unstable();
|
||||
|
||||
println!(
|
||||
"| {sibling_count} | {} | {:?} | {:?} | {:?} |",
|
||||
samples.len(),
|
||||
percentile(&samples, 50),
|
||||
percentile(&samples, 95),
|
||||
samples[samples.len() - 1],
|
||||
);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -333,11 +333,13 @@ pub(crate) fn ensure_blob_storage_version(schema: &Schema, params: &mut WritePar
|
||||
.data_storage_version
|
||||
.unwrap_or(LanceFileVersion::Stable)
|
||||
.resolve();
|
||||
if matches!(
|
||||
resolved,
|
||||
ConcreteFileVersion::V1 | ConcreteFileVersion::V2_0 | ConcreteFileVersion::V2_1
|
||||
) {
|
||||
params.data_storage_version = Some(LanceFileVersion::V2_2);
|
||||
// Exact formats deliberately have no Ord: capability is not implied by
|
||||
// release order. Enumerate every current concrete variant explicitly.
|
||||
match resolved {
|
||||
ConcreteFileVersion::V1 | ConcreteFileVersion::V2_0 | ConcreteFileVersion::V2_1 => {
|
||||
params.data_storage_version = Some(LanceFileVersion::V2_2);
|
||||
}
|
||||
ConcreteFileVersion::V2_2 | ConcreteFileVersion::V2_3 => {}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -502,7 +504,7 @@ mod tests {
|
||||
ensure_blob_storage_version(&blob_schema(), &mut params);
|
||||
assert_eq!(
|
||||
params.data_storage_version.unwrap().resolve(),
|
||||
ConcreteFileVersion::V2_2
|
||||
LanceFileVersion::V2_2.resolve()
|
||||
);
|
||||
}
|
||||
|
||||
@@ -515,7 +517,7 @@ mod tests {
|
||||
ensure_blob_storage_version(&blob_schema(), &mut params);
|
||||
assert_eq!(
|
||||
params.data_storage_version.unwrap().resolve(),
|
||||
ConcreteFileVersion::V2_2
|
||||
LanceFileVersion::V2_2.resolve()
|
||||
);
|
||||
}
|
||||
|
||||
@@ -526,7 +528,10 @@ mod tests {
|
||||
..Default::default()
|
||||
};
|
||||
ensure_blob_storage_version(&blob_schema(), &mut params);
|
||||
assert_eq!(params.data_storage_version.unwrap(), LanceFileVersion::V2_3);
|
||||
assert_eq!(
|
||||
params.data_storage_version.unwrap().resolve(),
|
||||
LanceFileVersion::V2_3.resolve()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -28,6 +28,7 @@ use crate::database::{
|
||||
};
|
||||
use crate::embeddings::{EmbeddingRegistry, MemoryRegistry};
|
||||
use crate::error::{Error, Result};
|
||||
use crate::function::{Function, FunctionId, RegisterFunctionJobSpec};
|
||||
#[cfg(feature = "remote")]
|
||||
use crate::remote::{
|
||||
client::ClientConfig,
|
||||
@@ -409,11 +410,6 @@ impl Connection {
|
||||
///
|
||||
/// The names will be returned in lexicographical order (ascending)
|
||||
///
|
||||
/// Listing databases discover physical `*.lance` entries without opening every
|
||||
/// dataset. The result is a point-in-time discovery snapshot: an entry may still be
|
||||
/// under creation, may contain only uncommitted storage, or may be concurrently
|
||||
/// dropped before it is opened.
|
||||
///
|
||||
/// The parameters `page_token` and `limit` can be used to paginate the results
|
||||
pub fn table_names(&self) -> TableNamesBuilder {
|
||||
TableNamesBuilder::new(self.internal.clone())
|
||||
@@ -461,9 +457,10 @@ impl Connection {
|
||||
///
|
||||
/// # Returns
|
||||
/// Created [`TableRef`], or [`Error::TableNotFound`] if the table does not exist.
|
||||
/// On listing databases, a committed Lance manifest is authoritative for table
|
||||
/// existence. Uncommitted files or a physical `<name>.lance` directory alone do not
|
||||
/// make a table openable.
|
||||
/// If the table's storage is present but holds no readable dataset (for example a
|
||||
/// `<name>.lance` directory left behind by an interrupted drop and re-create, which
|
||||
/// [`Self::table_names`] still lists) this returns [`Error::TableCorrupted`]
|
||||
/// instead.
|
||||
pub fn open_table(&self, name: impl Into<String>) -> OpenTableBuilder {
|
||||
OpenTableBuilder::new(
|
||||
self.internal.clone(),
|
||||
@@ -554,6 +551,88 @@ impl Connection {
|
||||
self.internal.job_history(job_id).await
|
||||
}
|
||||
|
||||
/// Submit a first-class Function registration job.
|
||||
///
|
||||
/// Returns a [`crate::job::Job`] handle for the accepted server-side job.
|
||||
/// Only remote databases support registration; local databases return
|
||||
/// [`Error::NotSupported`].
|
||||
pub async fn register_function(
|
||||
&self,
|
||||
spec: RegisterFunctionJobSpec,
|
||||
) -> Result<crate::job::Job> {
|
||||
self.internal.register_function(spec).await
|
||||
}
|
||||
|
||||
/// Look up the Function currently bound to a database-scoped name.
|
||||
///
|
||||
/// The name is lookup indirection only and is never part of the returned
|
||||
/// [`Function`]. Empty names return [`Error::InvalidInput`] before backend
|
||||
/// dispatch. Only remote databases support enterprise catalog lookup;
|
||||
/// nonempty local lookups return [`Error::NotSupported`].
|
||||
pub async fn lookup_function_by_name(&self, name: impl AsRef<str>) -> Result<Function> {
|
||||
let name = name.as_ref();
|
||||
// Public nonempty invariant: validate before any Database backend sees
|
||||
// the call so local and remote Connections agree on InvalidInput.
|
||||
if name.is_empty() {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "function lookup name must be non-empty".into(),
|
||||
});
|
||||
}
|
||||
self.internal.lookup_function_by_name(name).await
|
||||
}
|
||||
|
||||
/// Look up an immutable Function by exact opaque [`FunctionId`].
|
||||
///
|
||||
/// Exact-ID lookup is independent of later catalog name changes. Only
|
||||
/// remote databases support enterprise catalog lookup; local databases
|
||||
/// return [`Error::NotSupported`].
|
||||
pub async fn lookup_function_by_id(&self, function_id: &FunctionId) -> Result<Function> {
|
||||
self.internal.lookup_function_by_id(function_id).await
|
||||
}
|
||||
|
||||
/// Conditionally remove a database-scoped Function catalog name.
|
||||
///
|
||||
/// This is a direct synchronous catalog compare-and-swap (CAS), not a
|
||||
/// [`crate::job::Job`], not physical [`Function`] deletion, and not
|
||||
/// revocation. The caller supplies an observed immutable [`Function`]
|
||||
/// handle; only [`Function::id`] is authority for the CAS precondition.
|
||||
///
|
||||
/// Empty names return [`Error::InvalidInput`] before backend dispatch.
|
||||
/// Nonempty names on local/default backends return [`Error::NotSupported`].
|
||||
/// Remote backends complete only when the server reports durable CAS
|
||||
/// success for the `(name, current.id)` pair.
|
||||
pub async fn remove_function_name(
|
||||
&self,
|
||||
name: impl AsRef<str>,
|
||||
current: &Function,
|
||||
) -> Result<()> {
|
||||
let name = name.as_ref();
|
||||
// Public nonempty invariant: validate before any Database backend sees
|
||||
// the call so local and remote Connections agree on InvalidInput.
|
||||
if name.is_empty() {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "function name removal name must be non-empty".into(),
|
||||
});
|
||||
}
|
||||
self.internal.remove_function_name(name, current).await
|
||||
}
|
||||
|
||||
/// Revoke an exact immutable [`Function`] by opaque id.
|
||||
///
|
||||
/// This is a direct synchronous administrator catalog set-bit, not a
|
||||
/// [`crate::job::Job`], not catalog name removal, not physical deletion,
|
||||
/// and not [`Function`] or generated-column mutation. The caller supplies
|
||||
/// an already-validated exact [`Function`] handle; only [`Function::id`]
|
||||
/// is sent on the wire.
|
||||
///
|
||||
/// Local/default backends return [`Error::NotSupported`]. Remote backends
|
||||
/// complete only when the server reports durable success for that exact
|
||||
/// id. Repeated logical calls that each receive success succeed; there is
|
||||
/// no client-side already-revoked branch.
|
||||
pub async fn revoke_function(&self, function: &Function) -> Result<()> {
|
||||
self.internal.revoke_function(function).await
|
||||
}
|
||||
|
||||
/// Drop a table in the database.
|
||||
///
|
||||
/// # Arguments
|
||||
@@ -565,21 +644,6 @@ impl Connection {
|
||||
.await
|
||||
}
|
||||
|
||||
/// Start dropping a table and return a handle to the cleanup job.
|
||||
///
|
||||
/// The table may become unavailable before its physical data is removed.
|
||||
/// Call [`crate::job::Job::wait`] to wait for cleanup to finish. Local
|
||||
/// backends may complete the drop before returning the handle.
|
||||
pub async fn drop_table_async(
|
||||
&self,
|
||||
name: impl AsRef<str>,
|
||||
namespace_path: &[String],
|
||||
) -> Result<crate::job::Job> {
|
||||
self.internal
|
||||
.drop_table_async(name.as_ref(), namespace_path)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Drop the database
|
||||
///
|
||||
/// This is the same as dropping all of the tables
|
||||
|
||||
@@ -439,7 +439,7 @@ mod tests {
|
||||
.unwrap()
|
||||
.data_storage_format
|
||||
.lance_file_format();
|
||||
// Compare resolved versions since Stable/Next are aliases that resolve at storage time
|
||||
// Compare concrete stored format to the resolved requested alias.
|
||||
assert_eq!(storage_format, data_storage_version.resolve());
|
||||
}
|
||||
|
||||
|
||||
@@ -30,12 +30,16 @@ use lance_namespace::models::{
|
||||
|
||||
use crate::data::scannable::Scannable;
|
||||
use crate::error::Result;
|
||||
use crate::function::{Function, FunctionId, RegisterFunctionJobSpec};
|
||||
use crate::table::{BaseTable, WriteOptions};
|
||||
|
||||
pub mod listing;
|
||||
pub mod namespace;
|
||||
pub(crate) mod read_freshness;
|
||||
|
||||
#[cfg(test)]
|
||||
mod create_table_generated_column_schema_admission_contract;
|
||||
|
||||
pub trait DatabaseOptions {
|
||||
fn serialize_into_map(&self, map: &mut HashMap<String, String>);
|
||||
}
|
||||
@@ -230,6 +234,12 @@ pub struct JobDescription {
|
||||
pub creation_ms: i64,
|
||||
/// The job-type-specific specification. Null when the server omits it.
|
||||
pub spec: serde_json::Value,
|
||||
/// Explicit success result from the describe envelope, when present.
|
||||
///
|
||||
/// Missing or JSON `null` wire `result` is [`None`]. An explicit
|
||||
/// [`crate::JobResult::None`] object is `Some(JobResult::None)`. An exact
|
||||
/// Function result is `Some(JobResult::Function(...))`.
|
||||
pub result: Option<crate::job::JobResult>,
|
||||
/// Why the job failed, when the job is failed and the server reports a
|
||||
/// reason.
|
||||
pub failure: Option<crate::error::JobFailure>,
|
||||
@@ -311,6 +321,64 @@ pub trait Database:
|
||||
async fn job_history(&self, _job_id: Option<&str>) -> Result<Vec<RecordBatch>> {
|
||||
job_op_not_supported("job_history")
|
||||
}
|
||||
/// Submit a first-class Function registration job.
|
||||
///
|
||||
/// Returns a [`crate::job::Job`] handle for the accepted server-side job.
|
||||
/// Local databases do not support registration.
|
||||
async fn register_function(&self, _spec: RegisterFunctionJobSpec) -> Result<crate::job::Job> {
|
||||
job_op_not_supported("register_function")
|
||||
}
|
||||
/// Look up the Function currently bound to a database-scoped name.
|
||||
///
|
||||
/// The name is lookup indirection only and is never part of the returned
|
||||
/// [`Function`]. Empty names return [`crate::Error::InvalidInput`] before
|
||||
/// the unsupported fallback so local and remote backends agree. Nonempty
|
||||
/// names on databases without enterprise catalog lookup return
|
||||
/// [`crate::Error::NotSupported`].
|
||||
async fn lookup_function_by_name(&self, name: &str) -> Result<Function> {
|
||||
// Public nonempty invariant on the Database trait seam itself:
|
||||
// Connection::database() exposes Arc<dyn Database>, so empty-name
|
||||
// rejection cannot rely solely on Connection prevalidation.
|
||||
if name.is_empty() {
|
||||
return Err(crate::error::Error::InvalidInput {
|
||||
message: "function lookup name must be non-empty".into(),
|
||||
});
|
||||
}
|
||||
job_op_not_supported("lookup_function_by_name")
|
||||
}
|
||||
/// Look up an immutable Function by exact opaque [`FunctionId`].
|
||||
///
|
||||
/// Exact-ID lookup is independent of later catalog name changes. Local
|
||||
/// databases do not support enterprise catalog lookup.
|
||||
async fn lookup_function_by_id(&self, _function_id: &FunctionId) -> Result<Function> {
|
||||
job_op_not_supported("lookup_function_by_id")
|
||||
}
|
||||
/// Conditionally remove a database-scoped Function catalog name.
|
||||
///
|
||||
/// Direct synchronous catalog CAS, not a Job and not physical Function
|
||||
/// deletion. Empty names return [`crate::Error::InvalidInput`] before the
|
||||
/// unsupported fallback so local and remote backends agree. Nonempty names
|
||||
/// on databases without enterprise catalog mutation return
|
||||
/// [`crate::Error::NotSupported`].
|
||||
async fn remove_function_name(&self, name: &str, _current: &Function) -> Result<()> {
|
||||
// Public nonempty invariant on the Database trait seam itself:
|
||||
// Connection::database() exposes Arc<dyn Database>, so empty-name
|
||||
// rejection cannot rely solely on Connection prevalidation.
|
||||
if name.is_empty() {
|
||||
return Err(crate::error::Error::InvalidInput {
|
||||
message: "function name removal name must be non-empty".into(),
|
||||
});
|
||||
}
|
||||
job_op_not_supported("remove_function_name")
|
||||
}
|
||||
/// Revoke an exact immutable [`Function`] by opaque id.
|
||||
///
|
||||
/// Direct synchronous administrator catalog set-bit, not a Job, not name
|
||||
/// removal, and not physical Function deletion. Databases without
|
||||
/// enterprise catalog mutation return [`crate::Error::NotSupported`].
|
||||
async fn revoke_function(&self, _function: &Function) -> Result<()> {
|
||||
job_op_not_supported("revoke_function")
|
||||
}
|
||||
/// Open a table in the database
|
||||
async fn open_table(&self, request: OpenTableRequest) -> Result<Arc<dyn BaseTable>>;
|
||||
/// Rename a table in the database
|
||||
@@ -323,18 +391,6 @@ pub trait Database:
|
||||
) -> Result<()>;
|
||||
/// Drop a table in the database
|
||||
async fn drop_table(&self, name: &str, namespace_path: &[String]) -> Result<()>;
|
||||
/// Start dropping a table and return a handle to the cleanup job.
|
||||
///
|
||||
/// Backends without asynchronous cleanup complete the drop before
|
||||
/// returning an already-finished job.
|
||||
async fn drop_table_async(
|
||||
&self,
|
||||
name: &str,
|
||||
namespace_path: &[String],
|
||||
) -> Result<crate::job::Job> {
|
||||
self.drop_table(name, namespace_path).await?;
|
||||
Ok(crate::job::Job::new_done())
|
||||
}
|
||||
/// Drop all tables in the database
|
||||
async fn drop_all_tables(&self, namespace_path: &[String]) -> Result<()>;
|
||||
fn as_any(&self) -> &dyn std::any::Any;
|
||||
|
||||
@@ -0,0 +1,788 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! RED runtime contract tests for create-table schema admission (B4g).
|
||||
//!
|
||||
//! Caller-authored Arrow field metadata under
|
||||
//! [`crate::function::GENERATED_COLUMN_METADATA_KEY`] must not enter table
|
||||
//! schema state through general-purpose `Database::create_table`. Generated
|
||||
//! definitions are Job-owned. This module proves the missing admission guard
|
||||
//! on Native listing, Native namespace, and Remote create paths.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use arrow_array::{Int32Array, RecordBatch, StringArray};
|
||||
use arrow_schema::{DataType, Field, Schema, SchemaRef};
|
||||
use tempfile::TempDir;
|
||||
|
||||
use crate::arrow::SendableRecordBatchStream;
|
||||
use crate::data::scannable::Scannable;
|
||||
use crate::database::listing::ListingDatabase;
|
||||
use crate::database::{CreateTableMode, CreateTableRequest, Database, TableNamesRequest};
|
||||
use crate::error::Error;
|
||||
use crate::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
|
||||
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition,
|
||||
};
|
||||
|
||||
const ID: &str = "id";
|
||||
const ORDINARY: &str = "ordinary";
|
||||
const GEN_OUT: &str = "gen_out";
|
||||
const ORDINARY_META_KEY: &str = "unit";
|
||||
const ORDINARY_META_VALUE: &str = "label";
|
||||
const FN_ID: &str = "fn.exact.b4g.create_table.literal";
|
||||
const MALFORMED_MARKER: &str = "SENSITIVE_B4G_CREATE_TABLE_METADATA_MARKER_9d2e_a7c1";
|
||||
|
||||
/// Counts [`Scannable::scan_as_stream`] calls. [`Scannable::schema`] is free.
|
||||
struct ObservableScannable {
|
||||
batch: RecordBatch,
|
||||
scan_calls: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
impl ObservableScannable {
|
||||
fn new(batch: RecordBatch, scan_calls: Arc<AtomicUsize>) -> Self {
|
||||
Self { batch, scan_calls }
|
||||
}
|
||||
}
|
||||
|
||||
impl Scannable for ObservableScannable {
|
||||
fn schema(&self) -> SchemaRef {
|
||||
self.batch.schema()
|
||||
}
|
||||
|
||||
fn scan_as_stream(&mut self) -> SendableRecordBatchStream {
|
||||
self.scan_calls.fetch_add(1, Ordering::SeqCst);
|
||||
self.batch.scan_as_stream()
|
||||
}
|
||||
|
||||
fn num_rows(&self) -> Option<usize> {
|
||||
Some(self.batch.num_rows())
|
||||
}
|
||||
|
||||
fn rescannable(&self) -> bool {
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
fn literal_definition(output_field_id: i32) -> GeneratedColumnDefinition {
|
||||
let function = Function::new(
|
||||
FunctionId::try_new(FN_ID).unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("label", DataType::Utf8)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"label".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(StringArray::from(vec![Some("literal-only")])) as arrow_array::ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
GeneratedColumnDefinition::try_new(output_field_id, call, 1, 1).unwrap()
|
||||
}
|
||||
|
||||
fn valid_reserved_payload() -> String {
|
||||
literal_definition(1).to_metadata_json().unwrap()
|
||||
}
|
||||
|
||||
fn malformed_reserved_payload() -> String {
|
||||
format!(
|
||||
r#"{{"format_version":1,"output_field_id":1,"function_call":"{MALFORMED_MARKER}","dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
)
|
||||
}
|
||||
|
||||
fn batch_with_field_metadata(metadata: HashMap<String, String>) -> RecordBatch {
|
||||
let gen_field = Field::new(GEN_OUT, DataType::Int32, true).with_metadata(metadata);
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new(ID, DataType::Int32, false),
|
||||
Field::new(ORDINARY, DataType::Utf8, true),
|
||||
gen_field,
|
||||
]));
|
||||
RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(Int32Array::from(vec![1])),
|
||||
Arc::new(StringArray::from(vec![Some("seed")])),
|
||||
Arc::new(Int32Array::from(vec![10])),
|
||||
],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn reserved_batch(payload: &str) -> RecordBatch {
|
||||
batch_with_field_metadata(
|
||||
[(
|
||||
GENERATED_COLUMN_METADATA_KEY.to_string(),
|
||||
payload.to_string(),
|
||||
)]
|
||||
.into(),
|
||||
)
|
||||
}
|
||||
|
||||
fn ordinary_metadata_batch() -> RecordBatch {
|
||||
batch_with_field_metadata(
|
||||
[(
|
||||
ORDINARY_META_KEY.to_string(),
|
||||
ORDINARY_META_VALUE.to_string(),
|
||||
)]
|
||||
.into(),
|
||||
)
|
||||
}
|
||||
|
||||
fn plain_seed_batch() -> RecordBatch {
|
||||
batch_with_field_metadata(HashMap::new())
|
||||
}
|
||||
|
||||
fn assert_not_supported_redacted(err: &Error, label: &str, forbidden_substrings: &[&str]) {
|
||||
match err {
|
||||
Error::NotSupported { message } => {
|
||||
let rendered = format!("{err}\n{err:?}\n{message}");
|
||||
assert!(
|
||||
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
|
||||
"{label}: leaked metadata wire key: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(FN_ID),
|
||||
"{label}: leaked Function ID: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(MALFORMED_MARKER),
|
||||
"{label}: leaked malformed marker: {rendered}"
|
||||
);
|
||||
for needle in forbidden_substrings {
|
||||
assert!(
|
||||
!rendered.contains(needle),
|
||||
"{label}: leaked forbidden substring `{needle}`: {rendered}"
|
||||
);
|
||||
}
|
||||
assert!(
|
||||
message.to_lowercase().contains("generated")
|
||||
|| message.to_lowercase().contains("job"),
|
||||
"{label}: message must describe Job-owned generated-column boundary: {message}"
|
||||
);
|
||||
}
|
||||
other => panic!("{label}: expected Error::NotSupported, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
async fn listing_db() -> (TempDir, ListingDatabase) {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let uri = tmp.path().to_str().unwrap();
|
||||
let request = crate::connection::ConnectRequest {
|
||||
uri: uri.to_string(),
|
||||
#[cfg(feature = "remote")]
|
||||
client_config: Default::default(),
|
||||
options: Default::default(),
|
||||
namespace_client_properties: Default::default(),
|
||||
manifest_enabled: false,
|
||||
read_consistency_interval: None,
|
||||
session: None,
|
||||
};
|
||||
let db = ListingDatabase::connect_with_options(&request)
|
||||
.await
|
||||
.unwrap();
|
||||
(tmp, db)
|
||||
}
|
||||
|
||||
fn listing_table_dir(tmp: &TempDir, name: &str) -> std::path::PathBuf {
|
||||
tmp.path().join(format!("{name}.lance"))
|
||||
}
|
||||
|
||||
async fn listing_create(
|
||||
db: &ListingDatabase,
|
||||
name: &str,
|
||||
data: Box<dyn Scannable>,
|
||||
mode: CreateTableMode,
|
||||
) -> crate::Result<Arc<dyn crate::table::BaseTable>> {
|
||||
db.create_table(CreateTableRequest {
|
||||
name: name.to_string(),
|
||||
namespace_path: vec![],
|
||||
data,
|
||||
mode,
|
||||
write_options: Default::default(),
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn assert_listing_absent(db: &ListingDatabase, tmp: &TempDir, name: &str) {
|
||||
#[allow(deprecated)]
|
||||
let names = db.table_names(TableNamesRequest::default()).await.unwrap();
|
||||
assert!(
|
||||
!names.contains(&name.to_string()),
|
||||
"rejected create must leave no listed table `{name}`; got {names:?}"
|
||||
);
|
||||
assert!(
|
||||
!listing_table_dir(tmp, name).exists(),
|
||||
"rejected create must leave no storage directory for `{name}`"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn listing_create_rejects_reserved_generated_column_metadata_before_scan() {
|
||||
let (tmp, db) = listing_db().await;
|
||||
let payload = valid_reserved_payload();
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = listing_create(&db, "b4g_listing_create", data, CreateTableMode::Create)
|
||||
.await
|
||||
.expect_err("listing Create must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"listing Create reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(
|
||||
scan_calls.load(Ordering::SeqCst),
|
||||
0,
|
||||
"rejection must occur before Scannable::scan_as_stream"
|
||||
);
|
||||
assert_listing_absent(&db, &tmp, "b4g_listing_create").await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn listing_overwrite_rejects_reserved_generated_column_metadata_and_preserves_table() {
|
||||
let (tmp, db) = listing_db().await;
|
||||
let seed = listing_create(
|
||||
&db,
|
||||
"b4g_listing_overwrite",
|
||||
Box::new(plain_seed_batch()) as Box<dyn Scannable>,
|
||||
CreateTableMode::Create,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let version_before = seed.version().await.unwrap();
|
||||
let schema_before = seed.schema().await.unwrap();
|
||||
assert!(
|
||||
!schema_before
|
||||
.field_with_name(GEN_OUT)
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.contains_key(GENERATED_COLUMN_METADATA_KEY)
|
||||
);
|
||||
|
||||
let payload = malformed_reserved_payload();
|
||||
assert!(payload.contains(MALFORMED_MARKER));
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = listing_create(
|
||||
&db,
|
||||
"b4g_listing_overwrite",
|
||||
data,
|
||||
CreateTableMode::Overwrite,
|
||||
)
|
||||
.await
|
||||
.expect_err("listing Overwrite must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"listing Overwrite reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
|
||||
|
||||
let reopened = db
|
||||
.open_table(crate::database::OpenTableRequest {
|
||||
name: "b4g_listing_overwrite".to_string(),
|
||||
namespace_path: vec![],
|
||||
index_cache_size: None,
|
||||
lance_read_params: None,
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
managed_versioning: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(reopened.version().await.unwrap(), version_before);
|
||||
let schema_after = reopened.schema().await.unwrap();
|
||||
assert_eq!(schema_after.as_ref(), schema_before.as_ref());
|
||||
assert!(
|
||||
!schema_after
|
||||
.field_with_name(GEN_OUT)
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.contains_key(GENERATED_COLUMN_METADATA_KEY)
|
||||
);
|
||||
assert!(listing_table_dir(&tmp, "b4g_listing_overwrite").exists());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn listing_exist_ok_absent_rejects_reserved_generated_column_metadata_before_scan() {
|
||||
let (tmp, db) = listing_db().await;
|
||||
let payload = valid_reserved_payload();
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = listing_create(
|
||||
&db,
|
||||
"b4g_listing_exist_ok",
|
||||
data,
|
||||
CreateTableMode::exist_ok(|req| req),
|
||||
)
|
||||
.await
|
||||
.expect_err("listing ExistOk (absent) must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"listing ExistOk absent reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
|
||||
assert_listing_absent(&db, &tmp, "b4g_listing_exist_ok").await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn listing_ordinary_field_metadata_is_accepted_and_preserved() {
|
||||
let (_tmp, db) = listing_db().await;
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
ordinary_metadata_batch(),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let table = listing_create(&db, "b4g_listing_ordinary", data, CreateTableMode::Create)
|
||||
.await
|
||||
.expect("ordinary field metadata must remain accepted");
|
||||
assert!(
|
||||
scan_calls.load(Ordering::SeqCst) > 0,
|
||||
"successful create may consume the Scannable"
|
||||
);
|
||||
|
||||
let schema = table.schema().await.unwrap();
|
||||
let md = schema.field_with_name(GEN_OUT).unwrap().metadata();
|
||||
assert_eq!(
|
||||
md.get(ORDINARY_META_KEY).map(String::as_str),
|
||||
Some(ORDINARY_META_VALUE)
|
||||
);
|
||||
assert!(!md.contains_key(GENERATED_COLUMN_METADATA_KEY));
|
||||
}
|
||||
|
||||
#[cfg(not(windows))] // directory namespace tests are unix-only in this crate
|
||||
mod namespace_admission {
|
||||
use super::*;
|
||||
use crate::connect_namespace;
|
||||
use lance_namespace::models::{CreateNamespaceRequest, DescribeTableRequest};
|
||||
|
||||
async fn namespace_conn() -> (TempDir, crate::Connection) {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let root = tmp.path().to_str().unwrap().to_string();
|
||||
let mut properties = HashMap::new();
|
||||
properties.insert("root".to_string(), root);
|
||||
let conn = connect_namespace("dir", properties)
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
conn.create_namespace(CreateNamespaceRequest {
|
||||
id: Some(vec!["b4g_ns".into()]),
|
||||
..Default::default()
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
(tmp, conn)
|
||||
}
|
||||
|
||||
async fn assert_namespace_undeclared(conn: &crate::Connection, name: &str) {
|
||||
let names = conn
|
||||
.table_names()
|
||||
.namespace(vec!["b4g_ns".into()])
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(
|
||||
!names.contains(&name.to_string()),
|
||||
"rejected namespace create must leave no declared/listed table `{name}`; got {names:?}"
|
||||
);
|
||||
let ns = conn.namespace_client().await.unwrap();
|
||||
let describe = ns
|
||||
.describe_table(DescribeTableRequest {
|
||||
id: Some(vec!["b4g_ns".into(), name.into()]),
|
||||
..Default::default()
|
||||
})
|
||||
.await;
|
||||
assert!(
|
||||
describe.is_err(),
|
||||
"rejected namespace create must leave no describable table `{name}`"
|
||||
);
|
||||
}
|
||||
|
||||
async fn namespace_create(
|
||||
conn: &crate::Connection,
|
||||
name: &str,
|
||||
data: Box<dyn Scannable>,
|
||||
mode: CreateTableMode,
|
||||
) -> crate::Result<Arc<dyn crate::table::BaseTable>> {
|
||||
conn.database()
|
||||
.create_table(CreateTableRequest {
|
||||
name: name.to_string(),
|
||||
namespace_path: vec!["b4g_ns".into()],
|
||||
data,
|
||||
mode,
|
||||
write_options: Default::default(),
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn namespace_create_rejects_reserved_before_declare_describe_or_storage() {
|
||||
let (_tmp, conn) = namespace_conn().await;
|
||||
let payload = valid_reserved_payload();
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = namespace_create(&conn, "b4g_ns_create", data, CreateTableMode::Create)
|
||||
.await
|
||||
.expect_err("namespace Create must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"namespace Create reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
|
||||
assert_namespace_undeclared(&conn, "b4g_ns_create").await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn namespace_overwrite_rejects_reserved_before_declare_describe_or_storage() {
|
||||
let (_tmp, conn) = namespace_conn().await;
|
||||
let seed = namespace_create(
|
||||
&conn,
|
||||
"b4g_ns_overwrite",
|
||||
Box::new(plain_seed_batch()) as Box<dyn Scannable>,
|
||||
CreateTableMode::Create,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let version_before = seed.version().await.unwrap();
|
||||
let schema_before = seed.schema().await.unwrap();
|
||||
|
||||
let payload = malformed_reserved_payload();
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = namespace_create(&conn, "b4g_ns_overwrite", data, CreateTableMode::Overwrite)
|
||||
.await
|
||||
.expect_err("namespace Overwrite must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"namespace Overwrite reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
|
||||
|
||||
let reopened = conn
|
||||
.database()
|
||||
.open_table(crate::database::OpenTableRequest {
|
||||
name: "b4g_ns_overwrite".to_string(),
|
||||
namespace_path: vec!["b4g_ns".into()],
|
||||
index_cache_size: None,
|
||||
lance_read_params: None,
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
managed_versioning: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(reopened.version().await.unwrap(), version_before);
|
||||
assert_eq!(
|
||||
reopened.schema().await.unwrap().as_ref(),
|
||||
schema_before.as_ref()
|
||||
);
|
||||
assert!(
|
||||
!reopened
|
||||
.schema()
|
||||
.await
|
||||
.unwrap()
|
||||
.field_with_name(GEN_OUT)
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.contains_key(GENERATED_COLUMN_METADATA_KEY)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn namespace_exist_ok_absent_rejects_reserved_before_declare_describe_or_storage() {
|
||||
let (_tmp, conn) = namespace_conn().await;
|
||||
let payload = valid_reserved_payload();
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = namespace_create(
|
||||
&conn,
|
||||
"b4g_ns_exist_ok",
|
||||
data,
|
||||
CreateTableMode::exist_ok(|req| req),
|
||||
)
|
||||
.await
|
||||
.expect_err("namespace ExistOk (absent) must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"namespace ExistOk absent reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
|
||||
assert_namespace_undeclared(&conn, "b4g_ns_exist_ok").await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn namespace_exist_ok_existing_rejects_reserved_even_when_mode_would_ignore_data() {
|
||||
let (_tmp, conn) = namespace_conn().await;
|
||||
let seed = namespace_create(
|
||||
&conn,
|
||||
"b4g_ns_exist_ok_existing",
|
||||
Box::new(plain_seed_batch()) as Box<dyn Scannable>,
|
||||
CreateTableMode::Create,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let version_before = seed.version().await.unwrap();
|
||||
let schema_before = seed.schema().await.unwrap();
|
||||
|
||||
let payload = valid_reserved_payload();
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = namespace_create(
|
||||
&conn,
|
||||
"b4g_ns_exist_ok_existing",
|
||||
data,
|
||||
CreateTableMode::exist_ok(|req| req),
|
||||
)
|
||||
.await
|
||||
.expect_err(
|
||||
"namespace ExistOk must not accept reserved metadata merely because data is ignored",
|
||||
);
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"namespace ExistOk existing reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
|
||||
|
||||
let reopened = conn
|
||||
.database()
|
||||
.open_table(crate::database::OpenTableRequest {
|
||||
name: "b4g_ns_exist_ok_existing".to_string(),
|
||||
namespace_path: vec!["b4g_ns".into()],
|
||||
index_cache_size: None,
|
||||
lance_read_params: None,
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
managed_versioning: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(reopened.version().await.unwrap(), version_before);
|
||||
assert_eq!(
|
||||
reopened.schema().await.unwrap().as_ref(),
|
||||
schema_before.as_ref()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "remote")]
|
||||
mod remote_admission {
|
||||
use super::*;
|
||||
use std::io::Cursor;
|
||||
|
||||
use arrow_ipc::reader::StreamReader;
|
||||
use async_trait::async_trait;
|
||||
|
||||
use crate::Connection;
|
||||
use crate::remote::{ClientConfig, HeaderProvider};
|
||||
|
||||
#[derive(Debug)]
|
||||
struct CountingHeaderProvider {
|
||||
calls: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl HeaderProvider for CountingHeaderProvider {
|
||||
async fn get_headers(&self) -> crate::Result<HashMap<String, String>> {
|
||||
self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(HashMap::from([(
|
||||
"X-B4g-Test".to_string(),
|
||||
"must-not-be-requested".to_string(),
|
||||
)]))
|
||||
}
|
||||
}
|
||||
|
||||
fn counting_handler(
|
||||
calls: Arc<AtomicUsize>,
|
||||
) -> impl Fn(reqwest::Request) -> http::Response<String> + Clone + Send + Sync + 'static {
|
||||
move |_request| {
|
||||
calls.fetch_add(1, Ordering::SeqCst);
|
||||
http::Response::builder()
|
||||
.status(200)
|
||||
.body(String::new())
|
||||
.unwrap()
|
||||
}
|
||||
}
|
||||
|
||||
async fn remote_create(
|
||||
conn: &Connection,
|
||||
name: &str,
|
||||
data: Box<dyn Scannable>,
|
||||
mode: CreateTableMode,
|
||||
) -> crate::Result<Arc<dyn crate::table::BaseTable>> {
|
||||
// Direct Database trait path used by Connection::create_table.
|
||||
conn.database()
|
||||
.create_table(CreateTableRequest {
|
||||
name: name.to_string(),
|
||||
namespace_path: vec![],
|
||||
data,
|
||||
mode,
|
||||
write_options: Default::default(),
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn assert_remote_rejects(
|
||||
mode: CreateTableMode,
|
||||
table_name: &str,
|
||||
payload: &str,
|
||||
label: &str,
|
||||
) {
|
||||
let handler_calls = Arc::new(AtomicUsize::new(0));
|
||||
let header_calls = Arc::new(AtomicUsize::new(0));
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let config = ClientConfig {
|
||||
header_provider: Some(Arc::new(CountingHeaderProvider {
|
||||
calls: header_calls.clone(),
|
||||
}) as Arc<dyn HeaderProvider>),
|
||||
..Default::default()
|
||||
};
|
||||
let conn = Connection::new_with_handler_and_config(
|
||||
counting_handler(handler_calls.clone()),
|
||||
config,
|
||||
);
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = remote_create(&conn, table_name, data, mode)
|
||||
.await
|
||||
.expect_err("remote create must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(&err, label, &[payload]);
|
||||
assert_eq!(
|
||||
scan_calls.load(Ordering::SeqCst),
|
||||
0,
|
||||
"{label}: rejection must occur before scan_as_stream"
|
||||
);
|
||||
assert_eq!(
|
||||
header_calls.load(Ordering::SeqCst),
|
||||
0,
|
||||
"{label}: rejection must occur before header-provider invocation"
|
||||
);
|
||||
assert_eq!(
|
||||
handler_calls.load(Ordering::SeqCst),
|
||||
0,
|
||||
"{label}: rejection must occur before HTTP handler"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remote_create_rejects_reserved_before_scan_headers_and_http() {
|
||||
assert_remote_rejects(
|
||||
CreateTableMode::Create,
|
||||
"b4g_remote_create",
|
||||
&valid_reserved_payload(),
|
||||
"remote Create reserved admission",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remote_overwrite_rejects_reserved_before_scan_headers_and_http() {
|
||||
assert_remote_rejects(
|
||||
CreateTableMode::Overwrite,
|
||||
"b4g_remote_overwrite",
|
||||
&malformed_reserved_payload(),
|
||||
"remote Overwrite reserved admission",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remote_exist_ok_rejects_reserved_before_scan_headers_and_http() {
|
||||
assert_remote_rejects(
|
||||
CreateTableMode::exist_ok(|req| req),
|
||||
"b4g_remote_exist_ok",
|
||||
&valid_reserved_payload(),
|
||||
"remote ExistOk reserved admission",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remote_ordinary_field_metadata_is_transmitted_unchanged() {
|
||||
let conn = Connection::new_with_handler(|request| {
|
||||
assert_eq!(request.method(), &reqwest::Method::POST);
|
||||
assert_eq!(
|
||||
request.url().path(),
|
||||
"/v1/table/b4g_remote_ordinary/create/"
|
||||
);
|
||||
let body = request
|
||||
.body()
|
||||
.and_then(|b| b.as_bytes())
|
||||
.expect("ordinary create must send an Arrow IPC body");
|
||||
let reader = StreamReader::try_new(Cursor::new(body), None).unwrap();
|
||||
let schema = reader.schema();
|
||||
let md = schema.field_with_name(GEN_OUT).unwrap().metadata();
|
||||
assert_eq!(
|
||||
md.get(ORDINARY_META_KEY).map(String::as_str),
|
||||
Some(ORDINARY_META_VALUE),
|
||||
"ordinary field metadata must be transmitted unchanged"
|
||||
);
|
||||
assert!(!md.contains_key(GENERATED_COLUMN_METADATA_KEY));
|
||||
// Consume stream to completion for a well-formed IPC body.
|
||||
for batch in reader {
|
||||
batch.unwrap();
|
||||
}
|
||||
http::Response::builder()
|
||||
.status(200)
|
||||
.body(String::new())
|
||||
.unwrap()
|
||||
});
|
||||
|
||||
conn.create_table("b4g_remote_ordinary", ordinary_metadata_batch())
|
||||
.mode(CreateTableMode::Create)
|
||||
.execute()
|
||||
.await
|
||||
.expect("ordinary field metadata must remain accepted on remote create");
|
||||
}
|
||||
}
|
||||
@@ -23,6 +23,7 @@ use crate::connection::ConnectRequest;
|
||||
use crate::database::ReadConsistency;
|
||||
use crate::database::namespace::LanceNamespaceDatabase;
|
||||
use crate::error::{CreateDirSnafu, Error, Result};
|
||||
use crate::function::schema_admission::reject_caller_authored_generated_column_schema;
|
||||
use crate::io::object_store::MirroringObjectStoreWrapper;
|
||||
use crate::table::NativeTable;
|
||||
use crate::utils::validate_table_name;
|
||||
@@ -1032,13 +1033,16 @@ impl Database for ListingDatabase {
|
||||
};
|
||||
|
||||
Ok(ListTablesResponse {
|
||||
context: None,
|
||||
tables: f,
|
||||
page_token: next_page_token,
|
||||
})
|
||||
}
|
||||
|
||||
async fn create_table(&self, request: CreateTableRequest) -> Result<Arc<dyn BaseTable>> {
|
||||
// Admit schema before namespace forwarding, URI/config work, or NativeTable::create.
|
||||
// Scannable::schema is free; must not call scan_as_stream yet.
|
||||
reject_caller_authored_generated_column_schema(request.data.schema().as_ref())?;
|
||||
|
||||
if !request.namespace_path.is_empty() {
|
||||
return self.namespace_database().create_table(request).await;
|
||||
}
|
||||
@@ -1292,21 +1296,16 @@ impl Database for ListingDatabase {
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::Table;
|
||||
use crate::arrow::{SendableRecordBatchStream, SimpleRecordBatchStream};
|
||||
use crate::connection::ConnectRequest;
|
||||
use crate::data::scannable::Scannable;
|
||||
use crate::database::{CreateTableMode, CreateTableRequest};
|
||||
use crate::query::QueryRequest;
|
||||
use crate::table::{AnyQuery, WriteOptions};
|
||||
use arrow_array::{Int32Array, RecordBatch, StringArray};
|
||||
use arrow_schema::{DataType, Field, Schema, SchemaRef};
|
||||
use futures::{TryStreamExt, stream::once};
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
use futures::TryStreamExt;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use tempfile::tempdir;
|
||||
use tokio::sync::Barrier;
|
||||
use tokio::time::timeout;
|
||||
|
||||
async fn setup_database() -> (tempfile::TempDir, ListingDatabase) {
|
||||
let tempdir = tempdir().unwrap();
|
||||
@@ -1330,114 +1329,6 @@ mod tests {
|
||||
(tempdir, db)
|
||||
}
|
||||
|
||||
struct BarrierScannable {
|
||||
batch: RecordBatch,
|
||||
barrier: Arc<Barrier>,
|
||||
}
|
||||
|
||||
impl Scannable for BarrierScannable {
|
||||
fn schema(&self) -> SchemaRef {
|
||||
self.batch.schema()
|
||||
}
|
||||
|
||||
fn scan_as_stream(&mut self) -> SendableRecordBatchStream {
|
||||
let batch = self.batch.clone();
|
||||
let schema = batch.schema();
|
||||
let barrier = self.barrier.clone();
|
||||
Box::pin(SimpleRecordBatchStream {
|
||||
schema,
|
||||
stream: once(async move {
|
||||
barrier.wait().await;
|
||||
Ok(batch)
|
||||
}),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn create_request(name: &str, data: Box<dyn Scannable>) -> CreateTableRequest {
|
||||
CreateTableRequest {
|
||||
name: name.to_string(),
|
||||
namespace_path: vec![],
|
||||
data,
|
||||
mode: CreateTableMode::Create,
|
||||
write_options: Default::default(),
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_create_ignores_uncommitted_storage_without_manifest() {
|
||||
let (tmp_dir, db) = setup_database().await;
|
||||
let data_dir = tmp_dir.path().join("test.lance/data");
|
||||
std::fs::create_dir_all(&data_dir).unwrap();
|
||||
std::fs::write(data_dir.join("orphan.lance"), b"uncommitted").unwrap();
|
||||
|
||||
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]));
|
||||
let batch =
|
||||
RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(vec![1]))]).unwrap();
|
||||
|
||||
let table = db
|
||||
.create_table(create_request("test", Box::new(batch)))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(table.count_rows(None).await.unwrap(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_concurrent_create_is_arbitrated_by_manifest_commit() {
|
||||
let uri = format!("memory:///concurrent-create-{}", uuid::Uuid::new_v4());
|
||||
let db = crate::connect(&uri).execute().await.unwrap();
|
||||
let store: Arc<dyn object_store::ObjectStore> =
|
||||
Arc::new(object_store::memory::InMemory::new());
|
||||
let table_url = url::Url::parse("memory:///database/test.lance").unwrap();
|
||||
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]));
|
||||
let batch =
|
||||
RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(vec![1]))]).unwrap();
|
||||
let barrier = Arc::new(Barrier::new(2));
|
||||
|
||||
#[allow(deprecated)]
|
||||
let request = |batch, barrier| {
|
||||
let mut request = create_request("test", Box::new(BarrierScannable { batch, barrier }));
|
||||
request.write_options = WriteOptions {
|
||||
lance_write_params: Some(lance::dataset::WriteParams {
|
||||
store_params: Some(ObjectStoreParams {
|
||||
object_store: Some((store.clone(), table_url.clone())),
|
||||
..Default::default()
|
||||
}),
|
||||
commit_handler: Some(Arc::new(
|
||||
lance_table::io::commit::ConditionalPutCommitHandler,
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
};
|
||||
request
|
||||
};
|
||||
|
||||
let left = db
|
||||
.database()
|
||||
.create_table(request(batch.clone(), barrier.clone()));
|
||||
let right = db.database().create_table(request(batch, barrier));
|
||||
let (left, right) = timeout(Duration::from_secs(30), async { tokio::join!(left, right) })
|
||||
.await
|
||||
.expect("concurrent creates deadlocked");
|
||||
|
||||
let results = [left, right];
|
||||
assert_eq!(
|
||||
results.iter().filter(|result| result.is_ok()).count(),
|
||||
1,
|
||||
"expected one successful create, got {results:?}"
|
||||
);
|
||||
assert_eq!(
|
||||
results
|
||||
.iter()
|
||||
.filter(|result| matches!(result, Err(Error::TableAlreadyExists { .. })))
|
||||
.count(),
|
||||
1,
|
||||
"expected one manifest conflict, got {results:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_listing_database_root_ops_do_not_create_manifest() {
|
||||
let tempdir = tempdir().unwrap();
|
||||
|
||||
@@ -34,6 +34,7 @@ use crate::database::read_freshness::{
|
||||
FreshnessBaselines, ReadFreshnessContextProvider, TableFreshness,
|
||||
};
|
||||
use crate::error::{Error, Result};
|
||||
use crate::function::schema_admission::reject_caller_authored_generated_column_schema;
|
||||
use crate::table::{NativeTable, map_namespace_lance_error};
|
||||
use lance::dataset::WriteMode;
|
||||
|
||||
@@ -349,6 +350,10 @@ impl Database for LanceNamespaceDatabase {
|
||||
}
|
||||
|
||||
async fn create_table(&self, request: DbCreateTableRequest) -> Result<Arc<dyn BaseTable>> {
|
||||
// Admit schema before any mode branch, describe, declare, or storage work.
|
||||
// Scannable::schema is free; must not call scan_as_stream yet.
|
||||
reject_caller_authored_generated_column_schema(request.data.schema().as_ref())?;
|
||||
|
||||
let mut table_id = request.namespace_path.clone();
|
||||
table_id.push(request.name.clone());
|
||||
let mut existing_table = None;
|
||||
|
||||
+104
-8
@@ -6,10 +6,91 @@ use std::sync::{Arc, PoisonError};
|
||||
|
||||
use arrow_schema::ArrowError;
|
||||
use datafusion_common::DataFusionError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
use snafu::Snafu;
|
||||
|
||||
pub(crate) type BoxError = Box<dyn std::error::Error + Send + Sync>;
|
||||
|
||||
/// Stable Function error category (FF-006).
|
||||
///
|
||||
/// The known variants serialize to fixed JSON strings. Any other wire string
|
||||
/// decodes as [`Self::Unrecognized`] with the exact value preserved, and
|
||||
/// re-serializes unchanged. Category judgment is structural: do not infer a
|
||||
/// code from diagnostic message text, HTTP status, job phase, or retryability.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
#[non_exhaustive]
|
||||
pub enum FunctionErrorCode {
|
||||
/// Function definition failed validation.
|
||||
DefinitionValidationFailure,
|
||||
/// A named Function or Function reference was not found.
|
||||
NameOrFunctionNotFound,
|
||||
/// A Function name conflicts with an existing name.
|
||||
NameConflict,
|
||||
/// The requested runtime or capability is not supported.
|
||||
UnsupportedRuntimeOrCapability,
|
||||
/// The Function has been revoked and cannot be used.
|
||||
RevokedFunction,
|
||||
/// User-defined Function execution failed.
|
||||
UdfExecutionFailure,
|
||||
/// A generated column was not fully materialized.
|
||||
GeneratedColumnIncomplete,
|
||||
/// Input was stale or conflicted with the current state.
|
||||
StaleOrConflictingInput,
|
||||
/// A wire string this client version does not recognize.
|
||||
///
|
||||
/// The inner value is preserved exactly for forward compatibility.
|
||||
Unrecognized(String),
|
||||
}
|
||||
|
||||
impl FunctionErrorCode {
|
||||
/// The stable JSON / wire string for this code.
|
||||
pub fn as_str(&self) -> &str {
|
||||
match self {
|
||||
Self::DefinitionValidationFailure => "definition_validation_failure",
|
||||
Self::NameOrFunctionNotFound => "name_or_function_not_found",
|
||||
Self::NameConflict => "name_conflict",
|
||||
Self::UnsupportedRuntimeOrCapability => "unsupported_runtime_or_capability",
|
||||
Self::RevokedFunction => "revoked_function",
|
||||
Self::UdfExecutionFailure => "udf_execution_failure",
|
||||
Self::GeneratedColumnIncomplete => "generated_column_incomplete",
|
||||
Self::StaleOrConflictingInput => "stale_or_conflicting_input",
|
||||
Self::Unrecognized(raw) => raw.as_str(),
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(value: &str) -> Self {
|
||||
match value {
|
||||
"definition_validation_failure" => Self::DefinitionValidationFailure,
|
||||
"name_or_function_not_found" => Self::NameOrFunctionNotFound,
|
||||
"name_conflict" => Self::NameConflict,
|
||||
"unsupported_runtime_or_capability" => Self::UnsupportedRuntimeOrCapability,
|
||||
"revoked_function" => Self::RevokedFunction,
|
||||
"udf_execution_failure" => Self::UdfExecutionFailure,
|
||||
"generated_column_incomplete" => Self::GeneratedColumnIncomplete,
|
||||
"stale_or_conflicting_input" => Self::StaleOrConflictingInput,
|
||||
other => Self::Unrecognized(other.to_string()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Display for FunctionErrorCode {
|
||||
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
|
||||
f.write_str(self.as_str())
|
||||
}
|
||||
}
|
||||
|
||||
impl Serialize for FunctionErrorCode {
|
||||
fn serialize<S: Serializer>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error> {
|
||||
serializer.serialize_str(self.as_str())
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for FunctionErrorCode {
|
||||
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> std::result::Result<Self, D::Error> {
|
||||
Ok(Self::from_wire(String::deserialize(deserializer)?.as_str()))
|
||||
}
|
||||
}
|
||||
|
||||
/// Why a job failed, to whatever precision the backend provides.
|
||||
///
|
||||
/// A job run in this process carries the error it failed with in [`Self::source`].
|
||||
@@ -18,6 +99,12 @@ pub(crate) type BoxError = Box<dyn std::error::Error + Send + Sync>;
|
||||
/// backend does not supply it.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct JobFailure {
|
||||
/// Stable Function error category, when the backend supplied one.
|
||||
///
|
||||
/// Present only when copied from [`Error::Function`] or decoded from a
|
||||
/// remote `error_code` field. Never inferred from message, phase,
|
||||
/// retryable, HTTP status, or other diagnostics.
|
||||
pub error_code: Option<FunctionErrorCode>,
|
||||
/// The stage the job was in, when known.
|
||||
pub phase: Option<String>,
|
||||
/// A human-readable reason, when known.
|
||||
@@ -30,8 +117,16 @@ pub struct JobFailure {
|
||||
|
||||
impl JobFailure {
|
||||
/// A failure whose only known detail is the error that caused it.
|
||||
///
|
||||
/// When `source` is [`Error::Function`], [`Self::error_code`] is copied
|
||||
/// from that error. Other error kinds leave `error_code` as [`None`].
|
||||
pub(crate) fn from_source(source: Arc<Error>) -> Self {
|
||||
let error_code = match source.as_ref() {
|
||||
Error::Function { code, .. } => Some(code.clone()),
|
||||
_ => None,
|
||||
};
|
||||
Self {
|
||||
error_code,
|
||||
message: Some(source.to_string()),
|
||||
source: Some(source),
|
||||
..Default::default()
|
||||
@@ -71,14 +166,6 @@ pub enum Error {
|
||||
IndexNotFound { name: String },
|
||||
#[snafu(display("Embedding function '{name}' was not found. : {reason}"))]
|
||||
EmbeddingFunctionNotFound { name: String, reason: String },
|
||||
#[snafu(display("Column '{name}' was not found"))]
|
||||
ColumnNotFound { name: String },
|
||||
#[snafu(display("Column '{name}' already exists"))]
|
||||
ColumnAlreadyExists { name: String },
|
||||
#[snafu(display("Column '{name}' is not a computed column"))]
|
||||
NotAComputedColumn { name: String },
|
||||
#[snafu(display("Invalid expression for column '{column}': {message}"))]
|
||||
InvalidExpression { column: String, message: String },
|
||||
|
||||
#[snafu(display("Table '{name}' already exists"))]
|
||||
TableAlreadyExists { name: String },
|
||||
@@ -100,6 +187,15 @@ pub enum Error {
|
||||
},
|
||||
#[snafu(display("Job{} was cancelled", job_id.as_ref().map(|id| format!(" {id}")).unwrap_or_default()))]
|
||||
JobCancelled { job_id: Option<String> },
|
||||
/// A first-class Function operation failed with a stable category.
|
||||
///
|
||||
/// [`Self::Function::code`] is the semantic category. [`Self::Function::message`]
|
||||
/// is diagnostic only and must not be used to recover or override the code.
|
||||
#[snafu(display("Function error ({code}): {message}"))]
|
||||
Function {
|
||||
code: FunctionErrorCode,
|
||||
message: String,
|
||||
},
|
||||
|
||||
// 3rd party / external errors
|
||||
#[snafu(display("object_store error: {source}"))]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,776 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Atomic generated-column binding snapshot projection (FF-029).
|
||||
//!
|
||||
//! This is an implementation projection for table call binding. It is not a
|
||||
//! catalog resource, Job, persistent model, wire payload, or table-version
|
||||
//! replacement.
|
||||
|
||||
use std::collections::HashSet;
|
||||
|
||||
use arrow_schema::FieldRef;
|
||||
|
||||
use super::{
|
||||
FunctionCall, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition, invalid_input,
|
||||
};
|
||||
use crate::Result;
|
||||
|
||||
/// One top-level field identity from a single table snapshot.
|
||||
///
|
||||
/// Pairs a non-negative Lance stable field ID with the exact Arrow field from
|
||||
/// that same snapshot. IDs are carried only here; they are never injected into
|
||||
/// Arrow field metadata.
|
||||
#[doc(hidden)]
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct GeneratedColumnBindingEntry {
|
||||
field_id: i32,
|
||||
field: FieldRef,
|
||||
}
|
||||
|
||||
impl GeneratedColumnBindingEntry {
|
||||
/// Stable Lance field ID for this top-level entry.
|
||||
pub fn field_id(&self) -> i32 {
|
||||
self.field_id
|
||||
}
|
||||
|
||||
/// Exact Arrow field from the same snapshot.
|
||||
pub fn field(&self) -> &FieldRef {
|
||||
&self.field
|
||||
}
|
||||
|
||||
/// Strict generated-column definition from this entry's Arrow metadata.
|
||||
///
|
||||
/// Reads only [`GENERATED_COLUMN_METADATA_KEY`] on the exact snapshot field
|
||||
/// and decodes through
|
||||
/// [`GeneratedColumnDefinition::from_metadata_json`] with
|
||||
/// [`Self::field_id`] as the expected output identity. The same-snapshot
|
||||
/// stable field ID is mandatory so decode rejects metadata whose embedded
|
||||
/// `output_field_id` does not match this entry; name/ordinal/hash fallbacks
|
||||
/// are not used.
|
||||
///
|
||||
/// Returns [`Ok`]`(`[`None`]`)` when the key is absent. Present but invalid
|
||||
/// metadata fails closed as [`crate::Error::InvalidInput`] with a short
|
||||
/// field-ID diagnostic that does not echo the raw metadata payload.
|
||||
pub(crate) fn generated_column_definition(&self) -> Result<Option<GeneratedColumnDefinition>> {
|
||||
let Some(raw) = self.field.metadata().get(GENERATED_COLUMN_METADATA_KEY) else {
|
||||
return Ok(None);
|
||||
};
|
||||
match GeneratedColumnDefinition::from_metadata_json(raw, self.field_id) {
|
||||
Ok(definition) => Ok(Some(definition)),
|
||||
Err(_) => Err(invalid_input(format!(
|
||||
"invalid generated-column metadata for field id {}",
|
||||
self.field_id
|
||||
))),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Atomic table snapshot projection for generated-column call binding.
|
||||
///
|
||||
/// Contains one table version and immutable top-level field entries in schema
|
||||
/// order. Construction validates field/ID count equality, non-negative unique
|
||||
/// IDs, and unique top-level names.
|
||||
#[doc(hidden)]
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct GeneratedColumnBindingSnapshot {
|
||||
version: u64,
|
||||
entries: Vec<GeneratedColumnBindingEntry>,
|
||||
}
|
||||
|
||||
impl GeneratedColumnBindingSnapshot {
|
||||
/// Build a binding snapshot from one version and ordered field/ID pairs.
|
||||
///
|
||||
/// `fields` and `field_ids` must have the same length. Every ID must be
|
||||
/// non-negative and unique. Top-level field names must be unique. Order is
|
||||
/// preserved exactly as provided.
|
||||
pub fn try_new(
|
||||
version: u64,
|
||||
fields: impl IntoIterator<Item = FieldRef>,
|
||||
field_ids: impl IntoIterator<Item = i32>,
|
||||
) -> Result<Self> {
|
||||
let fields: Vec<FieldRef> = fields.into_iter().collect();
|
||||
let field_ids: Vec<i32> = field_ids.into_iter().collect();
|
||||
if fields.len() != field_ids.len() {
|
||||
return Err(invalid_input(
|
||||
"generated-column binding snapshot field count must equal field_ids count",
|
||||
));
|
||||
}
|
||||
|
||||
let mut seen_ids = HashSet::with_capacity(field_ids.len());
|
||||
let mut seen_names = HashSet::with_capacity(fields.len());
|
||||
let mut entries = Vec::with_capacity(fields.len());
|
||||
|
||||
for (field, field_id) in fields.into_iter().zip(field_ids) {
|
||||
if field_id < 0 {
|
||||
return Err(invalid_input(
|
||||
"generated-column binding snapshot field IDs must be non-negative",
|
||||
));
|
||||
}
|
||||
if !seen_ids.insert(field_id) {
|
||||
return Err(invalid_input(
|
||||
"generated-column binding snapshot field IDs must be unique",
|
||||
));
|
||||
}
|
||||
if !seen_names.insert(field.name().clone()) {
|
||||
return Err(invalid_input(
|
||||
"generated-column binding snapshot top-level field names must be unique",
|
||||
));
|
||||
}
|
||||
entries.push(GeneratedColumnBindingEntry { field_id, field });
|
||||
}
|
||||
|
||||
Ok(Self { version, entries })
|
||||
}
|
||||
|
||||
/// Table version for this snapshot.
|
||||
pub fn version(&self) -> u64 {
|
||||
self.version
|
||||
}
|
||||
|
||||
/// Top-level entries in schema order.
|
||||
pub fn entries(&self) -> &[GeneratedColumnBindingEntry] {
|
||||
&self.entries
|
||||
}
|
||||
|
||||
/// Exact case-sensitive top-level field name lookup.
|
||||
///
|
||||
/// A name containing `.` is a literal top-level field name, not a nested
|
||||
/// path. Lookup does not fold case or interpret dotted selectors.
|
||||
pub fn field(&self, name: &str) -> Option<&GeneratedColumnBindingEntry> {
|
||||
self.entries
|
||||
.iter()
|
||||
.find(|entry| entry.field().name() == name)
|
||||
}
|
||||
|
||||
/// Strict generated-column definition for one top-level column name.
|
||||
///
|
||||
/// Looks up the exact case-sensitive top-level name (`.` is literal, not a
|
||||
/// nested path), decodes through
|
||||
/// [`GeneratedColumnBindingEntry::generated_column_definition`] (preserving
|
||||
/// output stable-ID checking and raw-metadata redaction), then validates
|
||||
/// stored field arguments against this same snapshot via
|
||||
/// [`Self::validate_field_arguments`]. Returns the complete or incomplete
|
||||
/// definition unchanged. Does not perform table, catalog, network, or Job
|
||||
/// work and does not resolve a Function.
|
||||
///
|
||||
/// Returns [`crate::Error::InvalidInput`] for an empty name, a missing
|
||||
/// top-level field, an ordinary field without a valid generated-column
|
||||
/// definition, invalid metadata, or a field-argument identity/type
|
||||
/// mismatch against this snapshot.
|
||||
#[doc(hidden)]
|
||||
pub fn generated_column_definition(
|
||||
&self,
|
||||
column_name: impl AsRef<str>,
|
||||
) -> Result<GeneratedColumnDefinition> {
|
||||
let column_name = column_name.as_ref();
|
||||
if column_name.is_empty() {
|
||||
return Err(invalid_input("generated column name must not be empty"));
|
||||
}
|
||||
let Some(entry) = self.field(column_name) else {
|
||||
return Err(invalid_input(format!(
|
||||
"generated column '{column_name}' was not found in the table schema"
|
||||
)));
|
||||
};
|
||||
let Some(definition) = entry.generated_column_definition()? else {
|
||||
return Err(invalid_input(format!(
|
||||
"column '{column_name}' is not a generated column"
|
||||
)));
|
||||
};
|
||||
self.validate_field_arguments(definition.function_call())?;
|
||||
Ok(definition)
|
||||
}
|
||||
|
||||
/// Validate table-dependent field arguments of an already canonical call.
|
||||
///
|
||||
/// For every field argument, finds the snapshot entry by stable Lance field
|
||||
/// ID and requires exact Arrow [`arrow_schema::DataType`] equality. Literal
|
||||
/// arguments are table-independent and ignored. Missing field ID or type
|
||||
/// mismatch returns [`crate::Error::InvalidInput`] without modifying `call`
|
||||
/// or this snapshot.
|
||||
///
|
||||
/// This check is orthogonal to [`FunctionCall::validate_against`]: it does
|
||||
/// not perform catalog lookup, Function identity/signature validation, or
|
||||
/// table mutation.
|
||||
pub fn validate_field_arguments(&self, call: &FunctionCall) -> Result<()> {
|
||||
for (_parameter, argument) in call.arguments() {
|
||||
let Some(field_id) = argument.field_id() else {
|
||||
continue;
|
||||
};
|
||||
let Some(entry) = self.entry_by_field_id(field_id) else {
|
||||
return Err(invalid_input(format!(
|
||||
"generated-column binding snapshot missing field id {field_id}"
|
||||
)));
|
||||
};
|
||||
let expected = argument.data_type();
|
||||
let current = entry.field().data_type();
|
||||
if current != expected {
|
||||
return Err(invalid_input(format!(
|
||||
"generated-column binding snapshot field id {field_id} type mismatch: \
|
||||
expected {expected}, found {current}"
|
||||
)));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn entry_by_field_id(&self, field_id: i32) -> Option<&GeneratedColumnBindingEntry> {
|
||||
self.entries
|
||||
.iter()
|
||||
.find(|entry| entry.field_id() == field_id)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use arrow_schema::{DataType, Field};
|
||||
|
||||
use super::*;
|
||||
use crate::Error;
|
||||
|
||||
fn fields() -> Vec<FieldRef> {
|
||||
vec![
|
||||
Arc::new(Field::new("text", DataType::Utf8, true)),
|
||||
Arc::new(Field::new("Score", DataType::Int32, false)),
|
||||
Arc::new(Field::new("a.b", DataType::Utf8, true)),
|
||||
]
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn try_new_preserves_version_order_and_entry_data() {
|
||||
let snapshot =
|
||||
GeneratedColumnBindingSnapshot::try_new(11, fields(), vec![2, 4, 8]).unwrap();
|
||||
assert_eq!(snapshot.version(), 11);
|
||||
assert_eq!(snapshot.entries().len(), 3);
|
||||
assert_eq!(snapshot.entries()[0].field_id(), 2);
|
||||
assert_eq!(snapshot.entries()[0].field().name(), "text");
|
||||
assert_eq!(snapshot.entries()[0].field().data_type(), &DataType::Utf8);
|
||||
assert_eq!(snapshot.entries()[1].field_id(), 4);
|
||||
assert_eq!(snapshot.entries()[1].field().name(), "Score");
|
||||
assert_eq!(snapshot.entries()[2].field_id(), 8);
|
||||
assert_eq!(snapshot.entries()[2].field().name(), "a.b");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn lookup_is_exact_case_sensitive_and_treats_dot_literally() {
|
||||
let snapshot = GeneratedColumnBindingSnapshot::try_new(1, fields(), vec![2, 4, 8]).unwrap();
|
||||
assert_eq!(snapshot.field("Score").unwrap().field_id(), 4);
|
||||
assert!(snapshot.field("score").is_none());
|
||||
assert!(snapshot.field("TEXT").is_none());
|
||||
assert!(snapshot.field("a").is_none());
|
||||
assert!(snapshot.field("b").is_none());
|
||||
assert_eq!(snapshot.field("a.b").unwrap().field_id(), 8);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn try_new_rejects_count_mismatch_negative_duplicate_ids_and_names() {
|
||||
let base = fields();
|
||||
assert!(matches!(
|
||||
GeneratedColumnBindingSnapshot::try_new(1, base.clone(), vec![1, 2]),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
assert!(matches!(
|
||||
GeneratedColumnBindingSnapshot::try_new(1, base.clone(), vec![1, 2, -3]),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
assert!(matches!(
|
||||
GeneratedColumnBindingSnapshot::try_new(1, base.clone(), vec![1, 2, 1]),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
let duplicate_names = vec![
|
||||
Arc::new(Field::new("text", DataType::Utf8, true)),
|
||||
Arc::new(Field::new("text", DataType::Int32, false)),
|
||||
];
|
||||
assert!(matches!(
|
||||
GeneratedColumnBindingSnapshot::try_new(1, duplicate_names, vec![1, 2]),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
}
|
||||
|
||||
fn sample_function() -> crate::function::Function {
|
||||
use crate::function::{
|
||||
Function, FunctionId, FunctionOutput, FunctionParameter, FunctionSignature,
|
||||
};
|
||||
let id = FunctionId::try_new("fn.exact.snapshot.lib").unwrap();
|
||||
let signature = FunctionSignature::try_new(
|
||||
vec![
|
||||
FunctionParameter::new("payload_arg", DataType::Utf8),
|
||||
FunctionParameter::new("metric_arg", DataType::Int32),
|
||||
],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap();
|
||||
Function::new(id, signature)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_field_arguments_value_cases() {
|
||||
use crate::function::{FunctionArgument, FunctionCall};
|
||||
use arrow_array::{ArrayRef, Int32Array};
|
||||
|
||||
let snapshot = GeneratedColumnBindingSnapshot::try_new(2, fields(), vec![2, 4, 8]).unwrap();
|
||||
let function = sample_function();
|
||||
|
||||
let valid = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![
|
||||
(
|
||||
"payload_arg".to_string(),
|
||||
FunctionArgument::try_field(2, DataType::Utf8).unwrap(),
|
||||
),
|
||||
(
|
||||
"metric_arg".to_string(),
|
||||
FunctionArgument::try_field(4, DataType::Int32).unwrap(),
|
||||
),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
snapshot.validate_field_arguments(&valid).unwrap();
|
||||
|
||||
let missing = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![
|
||||
(
|
||||
"payload_arg".to_string(),
|
||||
FunctionArgument::try_field(99, DataType::Utf8).unwrap(),
|
||||
),
|
||||
(
|
||||
"metric_arg".to_string(),
|
||||
FunctionArgument::try_field(4, DataType::Int32).unwrap(),
|
||||
),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
snapshot.validate_field_arguments(&missing),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
|
||||
// Same stable ID, different Arrow type: exact-type equality must reject.
|
||||
let type_mismatch = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![
|
||||
(
|
||||
"payload_arg".to_string(),
|
||||
// ID 4 is Int32 in the snapshot.
|
||||
FunctionArgument::try_field(4, DataType::Utf8).unwrap(),
|
||||
),
|
||||
(
|
||||
"metric_arg".to_string(),
|
||||
FunctionArgument::try_field(4, DataType::Int32).unwrap(),
|
||||
),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
let err = snapshot
|
||||
.validate_field_arguments(&type_mismatch)
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, Error::InvalidInput { .. }));
|
||||
let message = err.to_string();
|
||||
assert!(message.contains('4'));
|
||||
assert!(message.contains("Utf8") && message.contains("Int32"));
|
||||
assert!(!message.contains("Score") && !message.contains("text"));
|
||||
|
||||
let mixed = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![
|
||||
(
|
||||
"payload_arg".to_string(),
|
||||
FunctionArgument::try_field(2, DataType::Utf8).unwrap(),
|
||||
),
|
||||
(
|
||||
"metric_arg".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(Int32Array::from(vec![Some(1)])) as ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
snapshot.validate_field_arguments(&mixed).unwrap();
|
||||
|
||||
let literal_only_fn = {
|
||||
use crate::function::{
|
||||
Function, FunctionId, FunctionOutput, FunctionParameter, FunctionSignature,
|
||||
};
|
||||
Function::new(
|
||||
FunctionId::try_new("fn.exact.snapshot.literal").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("constant_arg", DataType::Int32)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
};
|
||||
let literal_only =
|
||||
FunctionCall::try_new(
|
||||
&literal_only_fn,
|
||||
vec![(
|
||||
"constant_arg".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(Int32Array::from(vec![Some(9)])) as ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
// Empty snapshot still accepts literal-only calls.
|
||||
let empty =
|
||||
GeneratedColumnBindingSnapshot::try_new(1, Vec::<FieldRef>::new(), vec![]).unwrap();
|
||||
empty.validate_field_arguments(&literal_only).unwrap();
|
||||
}
|
||||
|
||||
fn status_sample_function() -> crate::function::Function {
|
||||
use crate::function::{
|
||||
Function, FunctionId, FunctionOutput, FunctionParameter, FunctionSignature,
|
||||
};
|
||||
Function::new(
|
||||
FunctionId::try_new("fn.exact.status.binding").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("label", DataType::Utf8)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn status_sample_call() -> crate::function::FunctionCall {
|
||||
use crate::function::{FunctionArgument, FunctionCall};
|
||||
use arrow_array::{ArrayRef, StringArray};
|
||||
let function = status_sample_function();
|
||||
FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"label".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(StringArray::from(vec![Some("ok")])) as ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn definition_json(
|
||||
output_field_id: i32,
|
||||
dependency_epoch: u64,
|
||||
materialized_epoch: u64,
|
||||
) -> String {
|
||||
use crate::function::GeneratedColumnDefinition;
|
||||
GeneratedColumnDefinition::try_new(
|
||||
output_field_id,
|
||||
status_sample_call(),
|
||||
dependency_epoch,
|
||||
materialized_epoch,
|
||||
)
|
||||
.unwrap()
|
||||
.to_metadata_json()
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn entry_with_metadata(
|
||||
name: &str,
|
||||
field_id: i32,
|
||||
metadata_json: Option<&str>,
|
||||
) -> GeneratedColumnBindingEntry {
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
let field = if let Some(json) = metadata_json {
|
||||
Field::new(name, DataType::Int32, true).with_metadata(
|
||||
[(GENERATED_COLUMN_METADATA_KEY.to_string(), json.to_string())].into(),
|
||||
)
|
||||
} else {
|
||||
Field::new(name, DataType::Int32, true)
|
||||
};
|
||||
let snapshot =
|
||||
GeneratedColumnBindingSnapshot::try_new(1, vec![Arc::new(field)], vec![field_id])
|
||||
.unwrap();
|
||||
snapshot.entries()[0].clone()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_column_definition_absent_returns_none() {
|
||||
let entry = entry_with_metadata("ordinary", 3, None);
|
||||
let got = entry.generated_column_definition().unwrap();
|
||||
assert!(got.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_column_definition_decodes_complete_and_incomplete() {
|
||||
use crate::function::{GeneratedColumnDefinition, GeneratedColumnStatus};
|
||||
|
||||
let complete_json = definition_json(5, 3, 3);
|
||||
let complete_entry = entry_with_metadata("gen_complete", 5, Some(&complete_json));
|
||||
let complete = complete_entry
|
||||
.generated_column_definition()
|
||||
.unwrap()
|
||||
.expect("complete metadata present");
|
||||
assert_eq!(complete.output_field_id(), 5);
|
||||
assert_eq!(complete.dependency_epoch(), 3);
|
||||
assert_eq!(complete.materialized_epoch(), 3);
|
||||
assert_eq!(complete.status(), GeneratedColumnStatus::Complete);
|
||||
assert_eq!(
|
||||
complete,
|
||||
GeneratedColumnDefinition::from_metadata_json(&complete_json, 5).unwrap()
|
||||
);
|
||||
|
||||
let incomplete_json = definition_json(7, 4, 2);
|
||||
let incomplete_entry = entry_with_metadata("gen_incomplete", 7, Some(&incomplete_json));
|
||||
let incomplete = incomplete_entry
|
||||
.generated_column_definition()
|
||||
.unwrap()
|
||||
.expect("incomplete metadata present");
|
||||
assert_eq!(incomplete.output_field_id(), 7);
|
||||
assert_eq!(incomplete.dependency_epoch(), 4);
|
||||
assert_eq!(incomplete.materialized_epoch(), 2);
|
||||
assert_eq!(incomplete.status(), GeneratedColumnStatus::Incomplete);
|
||||
assert_eq!(
|
||||
incomplete,
|
||||
GeneratedColumnDefinition::from_metadata_json(&incomplete_json, 7).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_column_definition_fail_closed_for_invalid_metadata() {
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
|
||||
let field_id = 9i32;
|
||||
let valid = definition_json(field_id, 2, 2);
|
||||
let mut mismatched: serde_json::Value = serde_json::from_str(&valid).unwrap();
|
||||
mismatched["output_field_id"] = serde_json::json!(field_id + 1);
|
||||
|
||||
let mut unsupported: serde_json::Value = serde_json::from_str(&valid).unwrap();
|
||||
unsupported["format_version"] = serde_json::json!(2);
|
||||
|
||||
let mut reversed: serde_json::Value = serde_json::from_str(&valid).unwrap();
|
||||
reversed["dependency_epoch"] = serde_json::json!(1);
|
||||
reversed["materialized_epoch"] = serde_json::json!(2);
|
||||
|
||||
let malformed_json = "{not-json";
|
||||
let mut malformed_call: serde_json::Value = serde_json::from_str(&valid).unwrap();
|
||||
malformed_call["function_call"] = serde_json::json!("not-an-object");
|
||||
|
||||
for raw in [
|
||||
mismatched.to_string(),
|
||||
unsupported.to_string(),
|
||||
reversed.to_string(),
|
||||
malformed_json.to_string(),
|
||||
malformed_call.to_string(),
|
||||
] {
|
||||
let entry = entry_with_metadata("gen_bad", field_id, Some(&raw));
|
||||
assert!(
|
||||
entry
|
||||
.field()
|
||||
.metadata()
|
||||
.contains_key(GENERATED_COLUMN_METADATA_KEY),
|
||||
"fixture must carry generated-column metadata"
|
||||
);
|
||||
let err = entry.generated_column_definition().unwrap_err();
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"expected InvalidInput, got {err:?}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_column_definition_errors_omit_raw_metadata_marker() {
|
||||
const MARKER: &str = "SENSITIVE_STATUS_METADATA_MARKER_b3d1_9f2e";
|
||||
let raw = format!(
|
||||
r#"{{"format_version":1,"output_field_id":3,"function_call":{MARKER},"dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
);
|
||||
assert!(raw.contains(MARKER));
|
||||
let entry = entry_with_metadata("gen_redact", 3, Some(&raw));
|
||||
let err = entry.generated_column_definition().unwrap_err();
|
||||
assert!(matches!(err, Error::InvalidInput { .. }));
|
||||
let text = format!("{err}\n{err:?}");
|
||||
assert!(
|
||||
!text.contains(MARKER),
|
||||
"status definition diagnostics must not echo raw metadata marker: {text}"
|
||||
);
|
||||
assert!(
|
||||
!text.contains(&raw),
|
||||
"status definition diagnostics must not echo raw metadata payload: {text}"
|
||||
);
|
||||
}
|
||||
|
||||
/// Build a definition whose stored field argument matches `input_type`.
|
||||
/// Construction succeeds even when the snapshot field at `input_field_id`
|
||||
/// later has a different Arrow type; same-snapshot validation catches that.
|
||||
fn field_arg_definition(
|
||||
output_field_id: i32,
|
||||
input_field_id: i32,
|
||||
input_type: DataType,
|
||||
dependency_epoch: u64,
|
||||
materialized_epoch: u64,
|
||||
) -> GeneratedColumnDefinition {
|
||||
use crate::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput,
|
||||
FunctionParameter, FunctionSignature,
|
||||
};
|
||||
let function = Function::new(
|
||||
FunctionId::try_new("fn.exact.snapshot.field_arg").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("payload", input_type.clone())],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"payload".to_string(),
|
||||
FunctionArgument::try_field(input_field_id, input_type).unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
GeneratedColumnDefinition::try_new(
|
||||
output_field_id,
|
||||
call,
|
||||
dependency_epoch,
|
||||
materialized_epoch,
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn snapshot_with_definition(
|
||||
version: u64,
|
||||
ordinary_name: &str,
|
||||
ordinary_id: i32,
|
||||
ordinary_type: DataType,
|
||||
gen_name: &str,
|
||||
gen_id: i32,
|
||||
definition: &GeneratedColumnDefinition,
|
||||
) -> GeneratedColumnBindingSnapshot {
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
let gen_field = Field::new(gen_name, DataType::Int32, true).with_metadata(
|
||||
[(
|
||||
GENERATED_COLUMN_METADATA_KEY.to_string(),
|
||||
definition.to_metadata_json().unwrap(),
|
||||
)]
|
||||
.into(),
|
||||
);
|
||||
GeneratedColumnBindingSnapshot::try_new(
|
||||
version,
|
||||
vec![
|
||||
Arc::new(Field::new(ordinary_name, ordinary_type, true)),
|
||||
Arc::new(gen_field),
|
||||
],
|
||||
vec![ordinary_id, gen_id],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn assert_snapshot_definition_invalid_input(err: &Error, label: &str) {
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"{label}: expected InvalidInput, got {err:?}"
|
||||
);
|
||||
let rendered = format!("{err}\n{err:?}");
|
||||
assert!(
|
||||
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
|
||||
"{label}: diagnostic leaked metadata wire key: {rendered}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn snapshot_generated_column_definition_returns_complete_and_incomplete() {
|
||||
use crate::function::GeneratedColumnStatus;
|
||||
|
||||
let complete = field_arg_definition(11, 3, DataType::Utf8, 4, 4);
|
||||
let snapshot =
|
||||
snapshot_with_definition(9, "text", 3, DataType::Utf8, "gen_out", 11, &complete);
|
||||
// High-level seam: name lookup + decode + same-snapshot field-arg check.
|
||||
// Callers keep using snapshot.version() for the FF-011 source pin.
|
||||
let got = snapshot.generated_column_definition("gen_out").unwrap();
|
||||
assert_eq!(got, complete);
|
||||
assert_eq!(got.status(), GeneratedColumnStatus::Complete);
|
||||
assert_eq!(snapshot.version(), 9);
|
||||
|
||||
let incomplete = field_arg_definition(11, 3, DataType::Utf8, 5, 2);
|
||||
let snapshot =
|
||||
snapshot_with_definition(10, "text", 3, DataType::Utf8, "gen_out", 11, &incomplete);
|
||||
let got = snapshot.generated_column_definition("gen_out").unwrap();
|
||||
assert_eq!(got, incomplete);
|
||||
assert_eq!(got.status(), GeneratedColumnStatus::Incomplete);
|
||||
|
||||
// Literal-only definitions remain valid (no field args to re-check).
|
||||
let literal = GeneratedColumnDefinition::try_new(13, status_sample_call(), 2, 2).unwrap();
|
||||
let snapshot = snapshot_with_definition(1, "text", 3, DataType::Utf8, "a.b", 13, &literal);
|
||||
assert_eq!(
|
||||
snapshot.generated_column_definition("a.b").unwrap(),
|
||||
literal
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn snapshot_generated_column_definition_rejects_empty_missing_ordinary_and_case() {
|
||||
let definition = field_arg_definition(11, 3, DataType::Utf8, 1, 1);
|
||||
let snapshot =
|
||||
snapshot_with_definition(1, "ordinary", 3, DataType::Utf8, "gen_out", 11, &definition);
|
||||
|
||||
for name in ["", "missing", "Gen_Out", "GEN_OUT", "ordinary", "gen.out"] {
|
||||
let err = snapshot.generated_column_definition(name).unwrap_err();
|
||||
assert_snapshot_definition_invalid_input(&err, name);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn snapshot_generated_column_definition_fail_closed_for_invalid_metadata() {
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
|
||||
let field_id = 11i32;
|
||||
let valid = definition_json(field_id, 2, 2);
|
||||
let mut mismatched: serde_json::Value = serde_json::from_str(&valid).unwrap();
|
||||
mismatched["output_field_id"] = serde_json::json!(field_id + 1);
|
||||
|
||||
const MARKER: &str = "SENSITIVE_SNAPSHOT_DEF_MARKER_c8e4_1a90";
|
||||
let malformed = format!(
|
||||
r#"{{"format_version":1,"output_field_id":{field_id},"function_call":{MARKER},"dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
);
|
||||
assert!(malformed.contains(MARKER));
|
||||
|
||||
for (label, raw) in [
|
||||
("output_field_id mismatch", mismatched.to_string()),
|
||||
("malformed function_call", malformed.clone()),
|
||||
] {
|
||||
let field = Field::new("gen_out", DataType::Int32, true)
|
||||
.with_metadata([(GENERATED_COLUMN_METADATA_KEY.to_string(), raw.clone())].into());
|
||||
let snapshot =
|
||||
GeneratedColumnBindingSnapshot::try_new(1, vec![Arc::new(field)], vec![field_id])
|
||||
.unwrap();
|
||||
let err = snapshot.generated_column_definition("gen_out").unwrap_err();
|
||||
assert_snapshot_definition_invalid_input(&err, label);
|
||||
let rendered = format!("{err}\n{err:?}");
|
||||
assert!(
|
||||
!rendered.contains(MARKER) && !rendered.contains(&raw),
|
||||
"{label}: must not echo raw metadata: {rendered}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn snapshot_generated_column_definition_validates_field_args_against_same_snapshot() {
|
||||
// Missing stable input identity: fixture constructs cleanly; projection fails.
|
||||
let missing = field_arg_definition(11, 99_999, DataType::Utf8, 3, 3);
|
||||
let snapshot =
|
||||
snapshot_with_definition(2, "text", 3, DataType::Utf8, "gen_out", 11, &missing);
|
||||
let err = snapshot.generated_column_definition("gen_out").unwrap_err();
|
||||
assert_snapshot_definition_invalid_input(&err, "missing stored input field id");
|
||||
|
||||
// Type drift: stored argument type matches FunctionCall construction, not
|
||||
// the snapshot field at that id.
|
||||
let mistyped = field_arg_definition(11, 3, DataType::Int32, 4, 4);
|
||||
let snapshot =
|
||||
snapshot_with_definition(3, "text", 3, DataType::Utf8, "gen_out", 11, &mistyped);
|
||||
assert_eq!(
|
||||
snapshot.field("text").unwrap().field().data_type(),
|
||||
&DataType::Utf8
|
||||
);
|
||||
let err = snapshot.generated_column_definition("gen_out").unwrap_err();
|
||||
assert_snapshot_definition_invalid_input(&err, "stored input Arrow type mismatch");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,168 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Immutable ChangeGeneratedColumnJobSpec change-generated-column Job
|
||||
//! operation input (FF-011).
|
||||
//!
|
||||
//! This type is Job operation input only. It does not look up catalogs or
|
||||
//! tables, execute Jobs, stage artifacts, call Lance, derive candidate
|
||||
//! definitions, or mutate epochs.
|
||||
|
||||
use std::fmt;
|
||||
|
||||
use serde::de::Error as DeError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
|
||||
use super::{Function, FunctionCall, GeneratedColumnDefinition, invalid_input};
|
||||
use crate::Result;
|
||||
|
||||
const FORMAT_VERSION_V1: u32 = 1;
|
||||
|
||||
/// Immutable Job operation input for changing a generated column (format
|
||||
/// version 1).
|
||||
///
|
||||
/// Semantic fields are exactly the expected [`GeneratedColumnDefinition`] CAS
|
||||
/// precondition and the new [`FunctionCall`]. Wire keys are exactly
|
||||
/// `format_version`, `expected_generated_column_definition`, and
|
||||
/// `new_function_call`.
|
||||
///
|
||||
/// Construction via [`Self::try_new`] validates only the new call against the
|
||||
/// new catalog [`Function`]. The expected definition is an opaque exact CAS
|
||||
/// precondition and is not validated against an old Function handle.
|
||||
/// Structural deserialize does not validate the new call either; execution
|
||||
/// consumers must call [`Self::validate_against`].
|
||||
///
|
||||
/// Both complete and incomplete expected definitions are accepted. Same-call
|
||||
/// change and new Functions whose output type or nullability differs from the
|
||||
/// old Function are valid. Status and output-type equality are not constructor
|
||||
/// or wire restrictions.
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct ChangeGeneratedColumnJobSpec {
|
||||
expected_generated_column_definition: GeneratedColumnDefinition,
|
||||
new_function_call: FunctionCall,
|
||||
}
|
||||
|
||||
impl ChangeGeneratedColumnJobSpec {
|
||||
/// Create a change-generated-column Job operation input.
|
||||
///
|
||||
/// Requires [`FunctionCall::validate_against`] to succeed for
|
||||
/// `new_function_call` and `new_function` before returning (exact Function
|
||||
/// ID, parameter name/order, argument count, and Arrow type equality).
|
||||
///
|
||||
/// The `expected_definition` is stored as an opaque exact CAS
|
||||
/// precondition. Its nested call is not validated against any Function.
|
||||
pub fn try_new(
|
||||
expected_definition: GeneratedColumnDefinition,
|
||||
new_function: &Function,
|
||||
new_function_call: FunctionCall,
|
||||
) -> Result<Self> {
|
||||
new_function_call.validate_against(new_function)?;
|
||||
Ok(Self {
|
||||
expected_generated_column_definition: expected_definition,
|
||||
new_function_call,
|
||||
})
|
||||
}
|
||||
|
||||
/// Wire format version (always 1 for this type).
|
||||
pub fn format_version(&self) -> u32 {
|
||||
FORMAT_VERSION_V1
|
||||
}
|
||||
|
||||
/// Expected generated-column definition used as an exact CAS precondition.
|
||||
pub fn expected_generated_column_definition(&self) -> &GeneratedColumnDefinition {
|
||||
&self.expected_generated_column_definition
|
||||
}
|
||||
|
||||
/// New function call to apply.
|
||||
pub fn new_function_call(&self) -> &FunctionCall {
|
||||
&self.new_function_call
|
||||
}
|
||||
|
||||
/// Validate the new call against a catalog [`Function`].
|
||||
///
|
||||
/// Structural decode does not perform this check. Execution consumers must
|
||||
/// call this before using the new call. The expected definition remains an
|
||||
/// opaque CAS precondition and is not validated here.
|
||||
pub fn validate_against(&self, new_function: &Function) -> Result<()> {
|
||||
self.new_function_call.validate_against(new_function)
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> ChangeGeneratedColumnJobSpecWire {
|
||||
ChangeGeneratedColumnJobSpecWire {
|
||||
format_version: FORMAT_VERSION_V1,
|
||||
expected_generated_column_definition: self.expected_generated_column_definition.clone(),
|
||||
new_function_call: self.new_function_call.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(wire: ChangeGeneratedColumnJobSpecWire) -> Result<Self> {
|
||||
if wire.format_version != FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported ChangeGeneratedColumnJobSpec format_version {}",
|
||||
wire.format_version
|
||||
)));
|
||||
}
|
||||
Ok(Self {
|
||||
expected_generated_column_definition: wire.expected_generated_column_definition,
|
||||
new_function_call: wire.new_function_call,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for ChangeGeneratedColumnJobSpec {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
let expected = &self.expected_generated_column_definition;
|
||||
let old_call = expected.function_call();
|
||||
let new_call = &self.new_function_call;
|
||||
let old_field_ids: Vec<_> = old_call
|
||||
.arguments()
|
||||
.iter()
|
||||
.filter_map(|(_, argument)| argument.field_id())
|
||||
.collect();
|
||||
let new_field_ids: Vec<_> = new_call
|
||||
.arguments()
|
||||
.iter()
|
||||
.filter_map(|(_, argument)| argument.field_id())
|
||||
.collect();
|
||||
f.debug_struct("ChangeGeneratedColumnJobSpec")
|
||||
.field("output_field_id", &expected.output_field_id())
|
||||
.field("old_function_id", &old_call.function_id().as_str())
|
||||
.field("new_function_id", &new_call.function_id().as_str())
|
||||
.field("dependency_epoch", &expected.dependency_epoch())
|
||||
.field("materialized_epoch", &expected.materialized_epoch())
|
||||
.field("old_argument_count", &old_call.arguments().len())
|
||||
.field("new_argument_count", &new_call.arguments().len())
|
||||
.field("old_field_ids", &old_field_ids)
|
||||
.field("new_field_ids", &new_field_ids)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
// Do not derive Debug: nested GeneratedColumnDefinition / FunctionCall may
|
||||
// carry typed literal payloads on the trusted change-generated-column wire.
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct ChangeGeneratedColumnJobSpecWire {
|
||||
format_version: u32,
|
||||
expected_generated_column_definition: GeneratedColumnDefinition,
|
||||
new_function_call: FunctionCall,
|
||||
}
|
||||
|
||||
impl Serialize for ChangeGeneratedColumnJobSpec {
|
||||
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
self.to_wire().serialize(serializer)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for ChangeGeneratedColumnJobSpec {
|
||||
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let wire = ChangeGeneratedColumnJobSpecWire::deserialize(deserializer)?;
|
||||
Self::from_wire(wire).map_err(D::Error::custom)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Immutable CreateGeneratedColumnJobSpec create-generated-column Job operation
|
||||
//! input (FF-009).
|
||||
//!
|
||||
//! This type is Job operation input only. It does not allocate output fields,
|
||||
//! construct [`super::GeneratedColumnDefinition`], mutate tables, or execute
|
||||
//! Jobs.
|
||||
|
||||
use std::fmt;
|
||||
|
||||
use serde::de::Error as DeError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
|
||||
use super::{Function, FunctionCall, invalid_input};
|
||||
use crate::Result;
|
||||
|
||||
const FORMAT_VERSION_V1: u32 = 1;
|
||||
|
||||
/// Immutable Job operation input for creating a generated column (format
|
||||
/// version 1).
|
||||
///
|
||||
/// Semantic fields are exactly `column_name` and [`FunctionCall`]. Wire keys
|
||||
/// are exactly `format_version`, `column_name`, and `function_call`.
|
||||
///
|
||||
/// Construction via [`Self::try_new`] validates the call against a catalog
|
||||
/// [`Function`]. Structural deserialize does not; execution consumers must
|
||||
/// call [`Self::validate_against`].
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct CreateGeneratedColumnJobSpec {
|
||||
column_name: String,
|
||||
function_call: FunctionCall,
|
||||
}
|
||||
|
||||
impl CreateGeneratedColumnJobSpec {
|
||||
/// Create a create-generated-column Job operation input.
|
||||
///
|
||||
/// Rejects an empty `column_name`. Requires
|
||||
/// [`FunctionCall::validate_against`] to succeed for `function` before
|
||||
/// returning (exact Function ID, parameter name/order, argument count, and
|
||||
/// Arrow type equality).
|
||||
pub fn try_new(
|
||||
column_name: impl Into<String>,
|
||||
function: &Function,
|
||||
call: FunctionCall,
|
||||
) -> Result<Self> {
|
||||
let column_name = column_name.into();
|
||||
if column_name.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"CreateGeneratedColumnJobSpec column_name must be non-empty",
|
||||
));
|
||||
}
|
||||
call.validate_against(function)?;
|
||||
Ok(Self {
|
||||
column_name,
|
||||
function_call: call,
|
||||
})
|
||||
}
|
||||
|
||||
/// Wire format version (always 1 for this type).
|
||||
pub fn format_version(&self) -> u32 {
|
||||
FORMAT_VERSION_V1
|
||||
}
|
||||
|
||||
/// Target generated column name.
|
||||
pub fn column_name(&self) -> &str {
|
||||
&self.column_name
|
||||
}
|
||||
|
||||
/// Embedded function call.
|
||||
pub fn function_call(&self) -> &FunctionCall {
|
||||
&self.function_call
|
||||
}
|
||||
|
||||
/// Validate the embedded call against a catalog [`Function`].
|
||||
///
|
||||
/// Structural decode does not perform this check. Execution consumers must
|
||||
/// call this before using the call.
|
||||
pub fn validate_against(&self, function: &Function) -> Result<()> {
|
||||
self.function_call.validate_against(function)
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> CreateGeneratedColumnJobSpecWire {
|
||||
CreateGeneratedColumnJobSpecWire {
|
||||
format_version: FORMAT_VERSION_V1,
|
||||
column_name: self.column_name.clone(),
|
||||
function_call: self.function_call.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(wire: CreateGeneratedColumnJobSpecWire) -> Result<Self> {
|
||||
if wire.format_version != FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported CreateGeneratedColumnJobSpec format_version {}",
|
||||
wire.format_version
|
||||
)));
|
||||
}
|
||||
if wire.column_name.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"CreateGeneratedColumnJobSpec column_name must be non-empty",
|
||||
));
|
||||
}
|
||||
Ok(Self {
|
||||
column_name: wire.column_name,
|
||||
function_call: wire.function_call,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for CreateGeneratedColumnJobSpec {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
let field_ids: Vec<_> = self
|
||||
.function_call
|
||||
.arguments()
|
||||
.iter()
|
||||
.filter_map(|(_, argument)| argument.field_id())
|
||||
.collect();
|
||||
f.debug_struct("CreateGeneratedColumnJobSpec")
|
||||
.field("column_name", &self.column_name)
|
||||
.field("function_id", &self.function_call.function_id().as_str())
|
||||
.field("argument_count", &self.function_call.arguments().len())
|
||||
.field("field_ids", &field_ids)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
// Do not derive Debug: nested FunctionCall may carry typed literal payloads on
|
||||
// the trusted create-generated-column wire.
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct CreateGeneratedColumnJobSpecWire {
|
||||
format_version: u32,
|
||||
column_name: String,
|
||||
function_call: FunctionCall,
|
||||
}
|
||||
|
||||
impl Serialize for CreateGeneratedColumnJobSpec {
|
||||
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
self.to_wire().serialize(serializer)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for CreateGeneratedColumnJobSpec {
|
||||
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let wire = CreateGeneratedColumnJobSpecWire::deserialize(deserializer)?;
|
||||
Self::from_wire(wire).map_err(D::Error::custom)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,427 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Immutable FunctionDefinition registration input (B1c / FF-007).
|
||||
//!
|
||||
//! These types are authoring/transport values only. They do not mint identity,
|
||||
//! store digests/artifacts, or execute Python.
|
||||
|
||||
use std::collections::HashSet;
|
||||
use std::fmt;
|
||||
|
||||
use serde::de::Error as DeError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
|
||||
use super::{FunctionSignature, SignatureWire, invalid_input};
|
||||
use crate::Result;
|
||||
|
||||
const FORMAT_VERSION_V1: u32 = 1;
|
||||
|
||||
/// Immutable Python implementation description for a [`FunctionDefinition`].
|
||||
///
|
||||
/// The source body is carried on the trusted registration wire but is omitted
|
||||
/// from [`Debug`] output.
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct PythonFunctionDefinition {
|
||||
module: String,
|
||||
callable: String,
|
||||
source: String,
|
||||
python: String,
|
||||
packages: Vec<String>,
|
||||
}
|
||||
|
||||
impl PythonFunctionDefinition {
|
||||
/// Create a Python implementation description.
|
||||
///
|
||||
/// Rejects empty `module`, `callable`, `source`, `python`, or any empty
|
||||
/// package requirement, and rejects duplicate package requirement strings.
|
||||
pub fn try_new(
|
||||
module: impl Into<String>,
|
||||
callable: impl Into<String>,
|
||||
source: impl Into<String>,
|
||||
python: impl Into<String>,
|
||||
packages: Vec<String>,
|
||||
) -> Result<Self> {
|
||||
let module = module.into();
|
||||
let callable = callable.into();
|
||||
let source = source.into();
|
||||
let python = python.into();
|
||||
validate_python_fields(&module, &callable, &source, &python, &packages)?;
|
||||
Ok(Self {
|
||||
module,
|
||||
callable,
|
||||
source,
|
||||
python,
|
||||
packages,
|
||||
})
|
||||
}
|
||||
|
||||
/// Python module name.
|
||||
pub fn module(&self) -> &str {
|
||||
&self.module
|
||||
}
|
||||
|
||||
/// Callable name within the module.
|
||||
pub fn callable(&self) -> &str {
|
||||
&self.callable
|
||||
}
|
||||
|
||||
/// Source body submitted at the trusted registration boundary.
|
||||
pub fn source(&self) -> &str {
|
||||
&self.source
|
||||
}
|
||||
|
||||
/// Requested Python runtime version string.
|
||||
pub fn python(&self) -> &str {
|
||||
&self.python
|
||||
}
|
||||
|
||||
/// Ordered package requirements.
|
||||
pub fn packages(&self) -> &[String] {
|
||||
&self.packages
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for PythonFunctionDefinition {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("PythonFunctionDefinition")
|
||||
.field("module", &self.module)
|
||||
.field("callable", &self.callable)
|
||||
.field("source", &"<redacted>")
|
||||
.field("python", &self.python)
|
||||
.field("packages", &self.packages)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_python_fields(
|
||||
module: &str,
|
||||
callable: &str,
|
||||
source: &str,
|
||||
python: &str,
|
||||
packages: &[String],
|
||||
) -> Result<()> {
|
||||
if module.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"PythonFunctionDefinition module must be non-empty",
|
||||
));
|
||||
}
|
||||
if callable.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"PythonFunctionDefinition callable must be non-empty",
|
||||
));
|
||||
}
|
||||
if source.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"PythonFunctionDefinition source must be non-empty",
|
||||
));
|
||||
}
|
||||
if python.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"PythonFunctionDefinition python must be non-empty",
|
||||
));
|
||||
}
|
||||
let mut seen = HashSet::with_capacity(packages.len());
|
||||
for package in packages {
|
||||
if package.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"PythonFunctionDefinition package must be non-empty",
|
||||
));
|
||||
}
|
||||
if !seen.insert(package.as_str()) {
|
||||
return Err(invalid_input(
|
||||
"PythonFunctionDefinition packages must not contain duplicates",
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Explicit capability grant attached to a [`FunctionDefinition`].
|
||||
///
|
||||
/// Secret references are carried on the trusted registration wire but are
|
||||
/// omitted from [`Debug`] output. Plaintext secret values are never part of
|
||||
/// this type.
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct FunctionCapability {
|
||||
kind: FunctionCapabilityKind,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
enum FunctionCapabilityKind {
|
||||
Network {
|
||||
origin: String,
|
||||
},
|
||||
Secret {
|
||||
reference: String,
|
||||
environment_variable: String,
|
||||
},
|
||||
}
|
||||
|
||||
impl FunctionCapability {
|
||||
/// Create a network capability for a non-empty origin.
|
||||
pub fn try_network(origin: impl Into<String>) -> Result<Self> {
|
||||
let origin = origin.into();
|
||||
if origin.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"FunctionCapability network origin must be non-empty",
|
||||
));
|
||||
}
|
||||
Ok(Self {
|
||||
kind: FunctionCapabilityKind::Network { origin },
|
||||
})
|
||||
}
|
||||
|
||||
/// Create a secret capability for a non-empty reference and environment variable.
|
||||
///
|
||||
/// Errors name the fields and never echo the reference value.
|
||||
pub fn try_secret(
|
||||
reference: impl Into<String>,
|
||||
environment_variable: impl Into<String>,
|
||||
) -> Result<Self> {
|
||||
let reference = reference.into();
|
||||
let environment_variable = environment_variable.into();
|
||||
if reference.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"FunctionCapability secret reference must be non-empty",
|
||||
));
|
||||
}
|
||||
if environment_variable.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"FunctionCapability secret environment_variable must be non-empty",
|
||||
));
|
||||
}
|
||||
Ok(Self {
|
||||
kind: FunctionCapabilityKind::Secret {
|
||||
reference,
|
||||
environment_variable,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
/// Network origin when this capability is a network grant.
|
||||
pub fn origin(&self) -> Option<&str> {
|
||||
match &self.kind {
|
||||
FunctionCapabilityKind::Network { origin } => Some(origin.as_str()),
|
||||
FunctionCapabilityKind::Secret { .. } => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Secret reference when this capability is a secret grant.
|
||||
pub fn reference(&self) -> Option<&str> {
|
||||
match &self.kind {
|
||||
FunctionCapabilityKind::Secret { reference, .. } => Some(reference.as_str()),
|
||||
FunctionCapabilityKind::Network { .. } => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Environment variable name when this capability is a secret grant.
|
||||
pub fn environment_variable(&self) -> Option<&str> {
|
||||
match &self.kind {
|
||||
FunctionCapabilityKind::Secret {
|
||||
environment_variable,
|
||||
..
|
||||
} => Some(environment_variable.as_str()),
|
||||
FunctionCapabilityKind::Network { .. } => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> CapabilityWire {
|
||||
match &self.kind {
|
||||
FunctionCapabilityKind::Network { origin } => CapabilityWire::Network {
|
||||
origin: origin.clone(),
|
||||
},
|
||||
FunctionCapabilityKind::Secret {
|
||||
reference,
|
||||
environment_variable,
|
||||
} => CapabilityWire::Secret {
|
||||
reference: reference.clone(),
|
||||
environment_variable: environment_variable.clone(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(wire: CapabilityWire) -> Result<Self> {
|
||||
match wire {
|
||||
CapabilityWire::Network { origin } => Self::try_network(origin),
|
||||
CapabilityWire::Secret {
|
||||
reference,
|
||||
environment_variable,
|
||||
} => Self::try_secret(reference, environment_variable),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for FunctionCapability {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match &self.kind {
|
||||
FunctionCapabilityKind::Network { origin } => f
|
||||
.debug_struct("FunctionCapability")
|
||||
.field("kind", &"network")
|
||||
.field("origin", origin)
|
||||
.finish(),
|
||||
FunctionCapabilityKind::Secret {
|
||||
environment_variable,
|
||||
..
|
||||
} => f
|
||||
.debug_struct("FunctionCapability")
|
||||
.field("kind", &"secret")
|
||||
.field("reference", &"<redacted>")
|
||||
.field("environment_variable", environment_variable)
|
||||
.finish(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Immutable registration input for a first-class Function (format version 1).
|
||||
///
|
||||
/// This value has no catalog identity. Source bodies and secret references are
|
||||
/// present on the trusted serde wire but omitted from [`Debug`].
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct FunctionDefinition {
|
||||
signature: FunctionSignature,
|
||||
python_definition: PythonFunctionDefinition,
|
||||
capabilities: Vec<FunctionCapability>,
|
||||
}
|
||||
|
||||
impl FunctionDefinition {
|
||||
/// Create a definition from a signature, Python implementation, and capabilities.
|
||||
///
|
||||
/// Emptiness and package uniqueness are enforced by the child constructors.
|
||||
pub fn try_new(
|
||||
signature: FunctionSignature,
|
||||
python_definition: PythonFunctionDefinition,
|
||||
capabilities: Vec<FunctionCapability>,
|
||||
) -> Result<Self> {
|
||||
Ok(Self {
|
||||
signature,
|
||||
python_definition,
|
||||
capabilities,
|
||||
})
|
||||
}
|
||||
|
||||
/// Function signature.
|
||||
pub fn signature(&self) -> &FunctionSignature {
|
||||
&self.signature
|
||||
}
|
||||
|
||||
/// Python implementation description.
|
||||
pub fn python_definition(&self) -> &PythonFunctionDefinition {
|
||||
&self.python_definition
|
||||
}
|
||||
|
||||
/// Ordered capability grants.
|
||||
pub fn capabilities(&self) -> &[FunctionCapability] {
|
||||
&self.capabilities
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> Result<FunctionDefinitionWire> {
|
||||
Ok(FunctionDefinitionWire {
|
||||
format_version: FORMAT_VERSION_V1,
|
||||
signature: self.signature.to_wire()?,
|
||||
implementation: ImplementationWire::Python {
|
||||
module: self.python_definition.module.clone(),
|
||||
callable: self.python_definition.callable.clone(),
|
||||
source: self.python_definition.source.clone(),
|
||||
python: self.python_definition.python.clone(),
|
||||
packages: self.python_definition.packages.clone(),
|
||||
},
|
||||
capabilities: self
|
||||
.capabilities
|
||||
.iter()
|
||||
.map(FunctionCapability::to_wire)
|
||||
.collect(),
|
||||
})
|
||||
}
|
||||
|
||||
fn from_wire(wire: FunctionDefinitionWire) -> Result<Self> {
|
||||
if wire.format_version != FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported FunctionDefinition format_version {}",
|
||||
wire.format_version
|
||||
)));
|
||||
}
|
||||
let signature = FunctionSignature::from_wire(wire.signature)?;
|
||||
let python_definition = match wire.implementation {
|
||||
ImplementationWire::Python {
|
||||
module,
|
||||
callable,
|
||||
source,
|
||||
python,
|
||||
packages,
|
||||
} => PythonFunctionDefinition::try_new(module, callable, source, python, packages)?,
|
||||
};
|
||||
let capabilities = wire
|
||||
.capabilities
|
||||
.into_iter()
|
||||
.map(FunctionCapability::from_wire)
|
||||
.collect::<Result<Vec<_>>>()?;
|
||||
Self::try_new(signature, python_definition, capabilities)
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for FunctionDefinition {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("FunctionDefinition")
|
||||
.field("signature", &self.signature)
|
||||
.field("python_definition", &self.python_definition)
|
||||
.field("capabilities", &self.capabilities)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
// Do not derive Debug: wire payloads carry Python source and secret references.
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct FunctionDefinitionWire {
|
||||
format_version: u32,
|
||||
signature: SignatureWire,
|
||||
implementation: ImplementationWire,
|
||||
capabilities: Vec<CapabilityWire>,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", deny_unknown_fields)]
|
||||
enum ImplementationWire {
|
||||
#[serde(rename = "python")]
|
||||
Python {
|
||||
module: String,
|
||||
callable: String,
|
||||
source: String,
|
||||
python: String,
|
||||
packages: Vec<String>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", deny_unknown_fields)]
|
||||
enum CapabilityWire {
|
||||
#[serde(rename = "network")]
|
||||
Network { origin: String },
|
||||
#[serde(rename = "secret")]
|
||||
Secret {
|
||||
reference: String,
|
||||
environment_variable: String,
|
||||
},
|
||||
}
|
||||
|
||||
impl Serialize for FunctionDefinition {
|
||||
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
self.to_wire()
|
||||
.map_err(serde::ser::Error::custom)?
|
||||
.serialize(serializer)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for FunctionDefinition {
|
||||
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let wire = FunctionDefinitionWire::deserialize(deserializer)?;
|
||||
Self::from_wire(wire).map_err(D::Error::custom)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,315 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Pure crate-private generated-column invalidation planner (B4a).
|
||||
//!
|
||||
//! Plans column-wide dependency-epoch advances from a binding snapshot and a
|
||||
//! mutation impact. This module does not mutate tables, write metadata, or
|
||||
//! execute append/update/delete/merge paths. Native append and update consume
|
||||
//! the plan through the B4b / B4c runtime wiring.
|
||||
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use super::{GeneratedColumnBindingSnapshot, GeneratedColumnDefinition};
|
||||
use crate::Result;
|
||||
|
||||
/// Mutation impact considered by the crate-private invalidation planner.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum GeneratedColumnMutationImpact {
|
||||
/// Append or delete: whole-column coverage / row membership changed.
|
||||
RowSetChanged,
|
||||
/// Update of the listed stable field IDs (direct and transitive dependents).
|
||||
///
|
||||
/// Native update (B4c) constructs this impact. Native append (B4b) only
|
||||
/// constructs [`Self::RowSetChanged`].
|
||||
UpdatedFields(BTreeSet<i32>),
|
||||
}
|
||||
|
||||
/// One planned field-metadata replacement produced by the pure planner.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct PlannedGeneratedColumnMetadataUpdate {
|
||||
output_field_id: i32,
|
||||
metadata_json: String,
|
||||
}
|
||||
|
||||
impl PlannedGeneratedColumnMetadataUpdate {
|
||||
/// Stable output field ID whose metadata should be replaced.
|
||||
pub fn output_field_id(&self) -> i32 {
|
||||
self.output_field_id
|
||||
}
|
||||
|
||||
/// Canonical [`GeneratedColumnDefinition::to_metadata_json`] bytes.
|
||||
pub fn metadata_json(&self) -> &str {
|
||||
&self.metadata_json
|
||||
}
|
||||
}
|
||||
|
||||
/// Plan generated-column metadata replacements for `impact`.
|
||||
///
|
||||
/// Planning is pure: `snapshot` is never mutated. Every present
|
||||
/// `lancedb::generated_column` value is decoded and every decoded call's field
|
||||
/// arguments are validated against `snapshot` before impact is calculated.
|
||||
/// Decode, missing-field, type-mismatch, serialization, or overflow errors
|
||||
/// return no plan.
|
||||
///
|
||||
/// Impacted definitions advance `dependency_epoch` exactly once (checked
|
||||
/// arithmetic) while preserving `materialized_epoch`, output identity, and the
|
||||
/// embedded [`super::FunctionCall`]. Replacements are returned in snapshot
|
||||
/// schema order.
|
||||
pub fn plan_generated_column_invalidation(
|
||||
snapshot: &GeneratedColumnBindingSnapshot,
|
||||
impact: &GeneratedColumnMutationImpact,
|
||||
) -> Result<Vec<PlannedGeneratedColumnMetadataUpdate>> {
|
||||
let definitions = decode_and_validate_generated_columns(snapshot)?;
|
||||
let impacted = compute_impacted_output_ids(&definitions, impact);
|
||||
|
||||
let mut plan = Vec::new();
|
||||
for (output_field_id, definition) in &definitions {
|
||||
if !impacted.contains(output_field_id) {
|
||||
continue;
|
||||
}
|
||||
let mut next = definition.clone();
|
||||
next.invalidate()?;
|
||||
let metadata_json = next.to_metadata_json()?;
|
||||
plan.push(PlannedGeneratedColumnMetadataUpdate {
|
||||
output_field_id: *output_field_id,
|
||||
metadata_json,
|
||||
});
|
||||
}
|
||||
Ok(plan)
|
||||
}
|
||||
|
||||
/// Decode every present generated-column definition in schema order and
|
||||
/// validate field arguments against the same snapshot.
|
||||
fn decode_and_validate_generated_columns(
|
||||
snapshot: &GeneratedColumnBindingSnapshot,
|
||||
) -> Result<Vec<(i32, GeneratedColumnDefinition)>> {
|
||||
let mut definitions = Vec::new();
|
||||
for entry in snapshot.entries() {
|
||||
let Some(definition) = entry.generated_column_definition()? else {
|
||||
continue;
|
||||
};
|
||||
snapshot.validate_field_arguments(definition.function_call())?;
|
||||
definitions.push((entry.field_id(), definition));
|
||||
}
|
||||
Ok(definitions)
|
||||
}
|
||||
|
||||
/// Compute the set of impacted generated-column output field IDs.
|
||||
///
|
||||
/// `RowSetChanged` impacts every generated column. `UpdatedFields` computes a
|
||||
/// deterministic fixed point over generated output IDs: a definition is
|
||||
/// impacted when any field argument references a dirty ID, and each generated
|
||||
/// definition is added at most once so cycles terminate.
|
||||
fn compute_impacted_output_ids(
|
||||
definitions: &[(i32, GeneratedColumnDefinition)],
|
||||
impact: &GeneratedColumnMutationImpact,
|
||||
) -> BTreeSet<i32> {
|
||||
match impact {
|
||||
GeneratedColumnMutationImpact::RowSetChanged => {
|
||||
definitions.iter().map(|(id, _)| *id).collect()
|
||||
}
|
||||
GeneratedColumnMutationImpact::UpdatedFields(updated) => {
|
||||
let mut dirty = updated.clone();
|
||||
let mut impacted = BTreeSet::new();
|
||||
let mut progressed = true;
|
||||
while progressed {
|
||||
progressed = false;
|
||||
for (output_field_id, definition) in definitions {
|
||||
if impacted.contains(output_field_id) {
|
||||
continue;
|
||||
}
|
||||
let depends_on_dirty =
|
||||
definition
|
||||
.function_call()
|
||||
.arguments()
|
||||
.iter()
|
||||
.any(|(_, argument)| {
|
||||
argument
|
||||
.field_id()
|
||||
.is_some_and(|field_id| dirty.contains(&field_id))
|
||||
});
|
||||
if depends_on_dirty {
|
||||
impacted.insert(*output_field_id);
|
||||
dirty.insert(*output_field_id);
|
||||
progressed = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
impacted
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use arrow_array::{ArrayRef, Int32Array};
|
||||
use arrow_schema::{DataType, Field, FieldRef};
|
||||
|
||||
use super::*;
|
||||
use crate::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
|
||||
FunctionSignature, GENERATED_COLUMN_METADATA_KEY,
|
||||
};
|
||||
|
||||
fn int_field_function(id: &str) -> Function {
|
||||
Function::new(
|
||||
FunctionId::try_new(id).unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("upstream", DataType::Int32)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn int_field_bound_call(function: &Function, input_field_id: i32) -> FunctionCall {
|
||||
FunctionCall::try_new(
|
||||
function,
|
||||
vec![(
|
||||
"upstream".to_string(),
|
||||
FunctionArgument::try_field(input_field_id, DataType::Int32).unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn definition(
|
||||
output_field_id: i32,
|
||||
call: FunctionCall,
|
||||
dependency_epoch: u64,
|
||||
materialized_epoch: u64,
|
||||
) -> GeneratedColumnDefinition {
|
||||
GeneratedColumnDefinition::try_new(
|
||||
output_field_id,
|
||||
call,
|
||||
dependency_epoch,
|
||||
materialized_epoch,
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn generated_field(name: &str, def: &GeneratedColumnDefinition) -> FieldRef {
|
||||
let json = def.to_metadata_json().unwrap();
|
||||
Arc::new(
|
||||
Field::new(name, DataType::Int32, true)
|
||||
.with_metadata([(GENERATED_COLUMN_METADATA_KEY.to_string(), json)].into()),
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cyclic_dependency_fixed_point_impacts_each_definition_at_most_once() {
|
||||
// A <-> B cycle. Seeding either side must terminate and advance each
|
||||
// impacted definition exactly once. This proves planner termination; it
|
||||
// is not a public cyclic-dependency creation guarantee.
|
||||
let a_id = 60;
|
||||
let b_id = 70;
|
||||
let fn_a = int_field_function("fn.exact.b4a.cycle.a");
|
||||
let fn_b = int_field_function("fn.exact.b4a.cycle.b");
|
||||
let a = definition(a_id, int_field_bound_call(&fn_a, b_id), 1, 1);
|
||||
let b = definition(b_id, int_field_bound_call(&fn_b, a_id), 2, 2);
|
||||
let snap = GeneratedColumnBindingSnapshot::try_new(
|
||||
11,
|
||||
vec![generated_field("gen_a", &a), generated_field("gen_b", &b)],
|
||||
vec![a_id, b_id],
|
||||
)
|
||||
.unwrap();
|
||||
let before = snap.clone();
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([a_id])),
|
||||
)
|
||||
.expect("cyclic fixed point must terminate");
|
||||
assert_eq!(snap, before);
|
||||
assert_eq!(plan.len(), 2);
|
||||
assert_eq!(plan[0].output_field_id(), a_id);
|
||||
assert_eq!(plan[1].output_field_id(), b_id);
|
||||
|
||||
let decoded_a =
|
||||
GeneratedColumnDefinition::from_metadata_json(plan[0].metadata_json(), a_id).unwrap();
|
||||
let decoded_b =
|
||||
GeneratedColumnDefinition::from_metadata_json(plan[1].metadata_json(), b_id).unwrap();
|
||||
assert_eq!(decoded_a.dependency_epoch(), 2);
|
||||
assert_eq!(decoded_a.materialized_epoch(), 1);
|
||||
assert_eq!(decoded_b.dependency_epoch(), 3);
|
||||
assert_eq!(decoded_b.materialized_epoch(), 2);
|
||||
assert_eq!(decoded_a.function_call(), a.function_call());
|
||||
assert_eq!(decoded_b.function_call(), b.function_call());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn row_set_change_with_cycle_still_invalidates_each_column_once() {
|
||||
let a_id = 61;
|
||||
let b_id = 71;
|
||||
let fn_a = int_field_function("fn.exact.b4a.cycle.row.a");
|
||||
let fn_b = int_field_function("fn.exact.b4a.cycle.row.b");
|
||||
let a = definition(a_id, int_field_bound_call(&fn_a, b_id), 5, 5);
|
||||
let b = definition(b_id, int_field_bound_call(&fn_b, a_id), 8, 8);
|
||||
let snap = GeneratedColumnBindingSnapshot::try_new(
|
||||
12,
|
||||
vec![generated_field("gen_b", &b), generated_field("gen_a", &a)],
|
||||
vec![b_id, a_id],
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::RowSetChanged,
|
||||
)
|
||||
.expect("row-set change over a cycle must plan once per column");
|
||||
assert_eq!(plan.len(), 2);
|
||||
assert_eq!(plan[0].output_field_id(), b_id);
|
||||
assert_eq!(plan[1].output_field_id(), a_id);
|
||||
let decoded_b =
|
||||
GeneratedColumnDefinition::from_metadata_json(plan[0].metadata_json(), b_id).unwrap();
|
||||
let decoded_a =
|
||||
GeneratedColumnDefinition::from_metadata_json(plan[1].metadata_json(), a_id).unwrap();
|
||||
assert_eq!(decoded_b.dependency_epoch(), 9);
|
||||
assert_eq!(decoded_a.dependency_epoch(), 6);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn literal_only_is_ignored_by_updated_fields_even_with_empty_seed() {
|
||||
let literal_fn = Function::new(
|
||||
FunctionId::try_new("fn.exact.b4a.cycle.literal").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("constant", DataType::Int32)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let literal_id = 80;
|
||||
let literal = definition(
|
||||
literal_id,
|
||||
FunctionCall::try_new(
|
||||
&literal_fn,
|
||||
vec![(
|
||||
"constant".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(Int32Array::from(vec![Some(1)])) as ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap(),
|
||||
3,
|
||||
3,
|
||||
);
|
||||
let snap = GeneratedColumnBindingSnapshot::try_new(
|
||||
13,
|
||||
vec![generated_field("gen_literal", &literal)],
|
||||
vec![literal_id],
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::new()),
|
||||
)
|
||||
.expect("empty UpdatedFields must succeed");
|
||||
assert!(plan.is_empty());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,483 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Contract tests for the crate-private generated-column invalidation planner (B4a).
|
||||
//!
|
||||
//! These tests pin the pure planning surface implemented by
|
||||
//! [`super::plan_generated_column_invalidation`]. No runtime append/update/delete
|
||||
//! path is exercised.
|
||||
|
||||
use std::collections::BTreeSet;
|
||||
use std::sync::Arc;
|
||||
|
||||
use arrow_array::{ArrayRef, Int32Array};
|
||||
use arrow_schema::{DataType, Field, FieldRef};
|
||||
|
||||
use super::plan_generated_column_invalidation::{
|
||||
GeneratedColumnMutationImpact, PlannedGeneratedColumnMetadataUpdate,
|
||||
plan_generated_column_invalidation,
|
||||
};
|
||||
use super::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
|
||||
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnBindingSnapshot,
|
||||
GeneratedColumnDefinition,
|
||||
};
|
||||
use crate::Error;
|
||||
|
||||
fn utf8_field_function() -> Function {
|
||||
Function::new(
|
||||
FunctionId::try_new("fn.exact.b4a.utf8").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("payload", DataType::Utf8)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn literal_only_function() -> Function {
|
||||
Function::new(
|
||||
FunctionId::try_new("fn.exact.b4a.literal").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("constant", DataType::Int32)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn int_field_function() -> Function {
|
||||
Function::new(
|
||||
FunctionId::try_new("fn.exact.b4a.int").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("upstream", DataType::Int32)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn field_bound_call(input_field_id: i32) -> FunctionCall {
|
||||
FunctionCall::try_new(
|
||||
&utf8_field_function(),
|
||||
vec![(
|
||||
"payload".to_string(),
|
||||
FunctionArgument::try_field(input_field_id, DataType::Utf8).unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn literal_only_call() -> FunctionCall {
|
||||
FunctionCall::try_new(
|
||||
&literal_only_function(),
|
||||
vec![(
|
||||
"constant".to_string(),
|
||||
FunctionArgument::try_literal(Arc::new(Int32Array::from(vec![Some(7)])) as ArrayRef)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn int_field_bound_call(input_field_id: i32) -> FunctionCall {
|
||||
FunctionCall::try_new(
|
||||
&int_field_function(),
|
||||
vec![(
|
||||
"upstream".to_string(),
|
||||
FunctionArgument::try_field(input_field_id, DataType::Int32).unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn definition(
|
||||
output_field_id: i32,
|
||||
call: FunctionCall,
|
||||
dependency_epoch: u64,
|
||||
materialized_epoch: u64,
|
||||
) -> GeneratedColumnDefinition {
|
||||
GeneratedColumnDefinition::try_new(output_field_id, call, dependency_epoch, materialized_epoch)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn ordinary_field(name: &str, data_type: DataType) -> FieldRef {
|
||||
Arc::new(Field::new(name, data_type, true))
|
||||
}
|
||||
|
||||
fn generated_field(name: &str, def: &GeneratedColumnDefinition) -> FieldRef {
|
||||
let json = def.to_metadata_json().unwrap();
|
||||
Arc::new(
|
||||
Field::new(name, DataType::Int32, true)
|
||||
.with_metadata([(GENERATED_COLUMN_METADATA_KEY.to_string(), json)].into()),
|
||||
)
|
||||
}
|
||||
|
||||
fn generated_field_with_raw_metadata(name: &str, raw: &str) -> FieldRef {
|
||||
Arc::new(
|
||||
Field::new(name, DataType::Int32, true)
|
||||
.with_metadata([(GENERATED_COLUMN_METADATA_KEY.to_string(), raw.to_string())].into()),
|
||||
)
|
||||
}
|
||||
|
||||
fn snapshot(
|
||||
version: u64,
|
||||
fields: Vec<FieldRef>,
|
||||
field_ids: Vec<i32>,
|
||||
) -> GeneratedColumnBindingSnapshot {
|
||||
GeneratedColumnBindingSnapshot::try_new(version, fields, field_ids).unwrap()
|
||||
}
|
||||
|
||||
fn expected_invalidated(def: &GeneratedColumnDefinition) -> GeneratedColumnDefinition {
|
||||
let mut next = def.clone();
|
||||
next.invalidate().unwrap();
|
||||
next
|
||||
}
|
||||
|
||||
fn assert_planned_definition(
|
||||
update: &PlannedGeneratedColumnMetadataUpdate,
|
||||
expected: &GeneratedColumnDefinition,
|
||||
) {
|
||||
assert_eq!(update.output_field_id(), expected.output_field_id());
|
||||
let decoded = GeneratedColumnDefinition::from_metadata_json(
|
||||
update.metadata_json(),
|
||||
expected.output_field_id(),
|
||||
)
|
||||
.expect("planned metadata must decode");
|
||||
assert_eq!(&decoded, expected);
|
||||
assert_eq!(
|
||||
update.metadata_json(),
|
||||
expected.to_metadata_json().unwrap(),
|
||||
"planned metadata JSON must be canonical"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn no_generated_columns_returns_empty_plan() {
|
||||
let snap = snapshot(
|
||||
1,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
ordinary_field("score", DataType::Int32),
|
||||
],
|
||||
vec![1, 2],
|
||||
);
|
||||
let before = snap.clone();
|
||||
|
||||
let plan =
|
||||
plan_generated_column_invalidation(&snap, &GeneratedColumnMutationImpact::RowSetChanged)
|
||||
.expect("planner must succeed when no generated columns are present");
|
||||
assert!(plan.is_empty());
|
||||
assert_eq!(snap, before, "planner must not mutate the binding snapshot");
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([1])),
|
||||
)
|
||||
.expect("field update with no generated columns must succeed");
|
||||
assert!(plan.is_empty());
|
||||
assert_eq!(snap, before, "planner must not mutate the binding snapshot");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn row_set_change_invalidates_field_bound_and_literal_only_exactly_once() {
|
||||
let text_id = 10;
|
||||
let field_bound_id = 20;
|
||||
let literal_id = 30;
|
||||
let field_bound = definition(field_bound_id, field_bound_call(text_id), 3, 3);
|
||||
let literal_only = definition(literal_id, literal_only_call(), 4, 4);
|
||||
let snap = snapshot(
|
||||
2,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
generated_field("gen_field", &field_bound),
|
||||
generated_field("gen_literal", &literal_only),
|
||||
],
|
||||
vec![text_id, field_bound_id, literal_id],
|
||||
);
|
||||
let before = snap.clone();
|
||||
|
||||
let plan =
|
||||
plan_generated_column_invalidation(&snap, &GeneratedColumnMutationImpact::RowSetChanged)
|
||||
.expect("row-set change must plan invalidation");
|
||||
assert_eq!(snap, before, "planner must not mutate the binding snapshot");
|
||||
assert_eq!(
|
||||
plan.len(),
|
||||
2,
|
||||
"each generated column invalidates exactly once"
|
||||
);
|
||||
assert_eq!(plan[0].output_field_id(), field_bound_id);
|
||||
assert_eq!(plan[1].output_field_id(), literal_id);
|
||||
assert_planned_definition(&plan[0], &expected_invalidated(&field_bound));
|
||||
assert_planned_definition(&plan[1], &expected_invalidated(&literal_only));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn already_incomplete_advances_dependency_epoch_and_preserves_materialized_epoch() {
|
||||
let text_id = 11;
|
||||
let gen_id = 21;
|
||||
let incomplete = definition(gen_id, field_bound_call(text_id), 9, 2);
|
||||
assert_eq!(incomplete.dependency_epoch(), 9);
|
||||
assert_eq!(incomplete.materialized_epoch(), 2);
|
||||
let snap = snapshot(
|
||||
3,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
generated_field("gen_incomplete", &incomplete),
|
||||
],
|
||||
vec![text_id, gen_id],
|
||||
);
|
||||
|
||||
let plan =
|
||||
plan_generated_column_invalidation(&snap, &GeneratedColumnMutationImpact::RowSetChanged)
|
||||
.expect("incomplete definition must still advance");
|
||||
assert_eq!(plan.len(), 1);
|
||||
let decoded =
|
||||
GeneratedColumnDefinition::from_metadata_json(plan[0].metadata_json(), gen_id).unwrap();
|
||||
assert_eq!(decoded.dependency_epoch(), 10);
|
||||
assert_eq!(decoded.materialized_epoch(), 2);
|
||||
assert_eq!(
|
||||
decoded.function_call(),
|
||||
incomplete.function_call(),
|
||||
"invalidation must preserve the embedded function call"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn direct_field_update_invalidates_only_dependent_generated_column() {
|
||||
let text_id = 12;
|
||||
let score_id = 13;
|
||||
let dependent_id = 22;
|
||||
let unrelated_gen_id = 23;
|
||||
let dependent = definition(dependent_id, field_bound_call(text_id), 5, 5);
|
||||
let unrelated_gen = definition(unrelated_gen_id, literal_only_call(), 6, 6);
|
||||
let snap = snapshot(
|
||||
4,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
ordinary_field("score", DataType::Int32),
|
||||
generated_field("gen_dependent", &dependent),
|
||||
generated_field("gen_unrelated", &unrelated_gen),
|
||||
],
|
||||
vec![text_id, score_id, dependent_id, unrelated_gen_id],
|
||||
);
|
||||
let before = snap.clone();
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([text_id])),
|
||||
)
|
||||
.expect("dependent update must plan a single invalidation");
|
||||
assert_eq!(snap, before);
|
||||
assert_eq!(plan.len(), 1);
|
||||
assert_planned_definition(&plan[0], &expected_invalidated(&dependent));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unrelated_field_update_returns_empty_plan() {
|
||||
let text_id = 14;
|
||||
let score_id = 15;
|
||||
let gen_id = 24;
|
||||
let dependent = definition(gen_id, field_bound_call(text_id), 2, 2);
|
||||
let snap = snapshot(
|
||||
5,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
ordinary_field("score", DataType::Int32),
|
||||
generated_field("gen_text", &dependent),
|
||||
],
|
||||
vec![text_id, score_id, gen_id],
|
||||
);
|
||||
let before = snap.clone();
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([score_id])),
|
||||
)
|
||||
.expect("unrelated update must not invent invalidation");
|
||||
assert!(plan.is_empty());
|
||||
assert_eq!(snap, before);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transitive_dependency_propagation_follows_snapshot_order() {
|
||||
// A (ordinary) -> B (generated) -> C (generated). Update A invalidates B and C.
|
||||
let a_id = 30;
|
||||
let b_id = 40;
|
||||
let c_id = 50;
|
||||
let b = definition(b_id, field_bound_call(a_id), 1, 1);
|
||||
let c = definition(c_id, int_field_bound_call(b_id), 1, 1);
|
||||
// Schema order places C before B so the plan must follow snapshot order, not
|
||||
// dependency discovery order.
|
||||
let snap = snapshot(
|
||||
6,
|
||||
vec![
|
||||
ordinary_field("a", DataType::Utf8),
|
||||
generated_field("gen_c", &c),
|
||||
generated_field("gen_b", &b),
|
||||
],
|
||||
vec![a_id, c_id, b_id],
|
||||
);
|
||||
let before = snap.clone();
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([a_id])),
|
||||
)
|
||||
.expect("transitive dependents must invalidate");
|
||||
assert_eq!(snap, before);
|
||||
assert_eq!(plan.len(), 2);
|
||||
assert_eq!(plan[0].output_field_id(), c_id);
|
||||
assert_eq!(plan[1].output_field_id(), b_id);
|
||||
assert_planned_definition(&plan[0], &expected_invalidated(&c));
|
||||
assert_planned_definition(&plan[1], &expected_invalidated(&b));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn malformed_metadata_fails_closed_for_unrelated_update_without_echoing_payload() {
|
||||
const MARKER: &str = "SENSITIVE_B4A_METADATA_MARKER_7c91_e2aa";
|
||||
let text_id = 16;
|
||||
let score_id = 17;
|
||||
let bad_id = 25;
|
||||
let raw = format!(
|
||||
r#"{{"format_version":1,"output_field_id":{bad_id},"function_call":{MARKER},"dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
);
|
||||
assert!(raw.contains(MARKER));
|
||||
let snap = snapshot(
|
||||
7,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
ordinary_field("score", DataType::Int32),
|
||||
generated_field_with_raw_metadata("gen_bad", &raw),
|
||||
],
|
||||
vec![text_id, score_id, bad_id],
|
||||
);
|
||||
|
||||
let err = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([score_id])),
|
||||
)
|
||||
.expect_err("malformed metadata must fail closed even for an unrelated update");
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"expected InvalidInput, got {err:?}"
|
||||
);
|
||||
let text = format!("{err}\n{err:?}");
|
||||
assert!(
|
||||
!text.contains(MARKER),
|
||||
"diagnostics must not echo raw metadata marker: {text}"
|
||||
);
|
||||
assert!(
|
||||
!text.contains(&raw),
|
||||
"diagnostics must not echo raw metadata payload: {text}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_input_field_id_fails_closed() {
|
||||
let missing_input_id = 99;
|
||||
let gen_id = 26;
|
||||
let orphan = definition(gen_id, field_bound_call(missing_input_id), 1, 1);
|
||||
let snap = snapshot(
|
||||
8,
|
||||
vec![
|
||||
ordinary_field("score", DataType::Int32),
|
||||
generated_field("gen_orphan", &orphan),
|
||||
],
|
||||
vec![18, gen_id],
|
||||
);
|
||||
|
||||
let err = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([18])),
|
||||
)
|
||||
.expect_err("missing stable input field id must fail closed");
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"expected InvalidInput, got {err:?}"
|
||||
);
|
||||
let message = err.to_string();
|
||||
assert!(
|
||||
message.contains("99") || message.contains("missing"),
|
||||
"diagnostic should identify the missing field id: {message}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn field_type_mismatch_fails_closed() {
|
||||
let text_id = 19;
|
||||
let gen_id = 27;
|
||||
// Definition claims Utf8 for field 19, but the snapshot entry is Int32.
|
||||
let mismatched = definition(gen_id, field_bound_call(text_id), 1, 1);
|
||||
let snap = snapshot(
|
||||
9,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Int32),
|
||||
generated_field("gen_mismatch", &mismatched),
|
||||
],
|
||||
vec![text_id, gen_id],
|
||||
);
|
||||
|
||||
let err = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([text_id])),
|
||||
)
|
||||
.expect_err("field type mismatch must fail closed");
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"expected InvalidInput, got {err:?}"
|
||||
);
|
||||
let message = err.to_string();
|
||||
assert!(
|
||||
message.contains("mismatch")
|
||||
|| (message.contains("Utf8") && message.contains("Int32"))
|
||||
|| message.contains(&text_id.to_string()),
|
||||
"diagnostic should identify the type mismatch: {message}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn epoch_overflow_fails_atomically_with_stable_sanitized_diagnostic() {
|
||||
let text_id = 31;
|
||||
let overflow_id = 41;
|
||||
let other_id = 42;
|
||||
let at_max = definition(overflow_id, field_bound_call(text_id), u64::MAX, u64::MAX);
|
||||
let other = definition(other_id, literal_only_call(), 1, 1);
|
||||
let snap = snapshot(
|
||||
10,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
generated_field("gen_max", &at_max),
|
||||
generated_field("gen_other", &other),
|
||||
],
|
||||
vec![text_id, overflow_id, other_id],
|
||||
);
|
||||
|
||||
// Row-set change impacts every generated column, including the overflowed one.
|
||||
let err =
|
||||
plan_generated_column_invalidation(&snap, &GeneratedColumnMutationImpact::RowSetChanged)
|
||||
.expect_err("dependency_epoch overflow must fail closed");
|
||||
match err {
|
||||
Error::InvalidInput { message } => {
|
||||
assert_eq!(
|
||||
message, "dependency_epoch overflow",
|
||||
"overflow must use the existing sanitized InvalidInput diagnostic"
|
||||
);
|
||||
}
|
||||
other => panic!("expected InvalidInput overflow, got {other:?}"),
|
||||
}
|
||||
|
||||
// Direct update that impacts only the overflowed definition must also fail
|
||||
// atomically and must not return a partial plan for sibling columns.
|
||||
let err = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([text_id])),
|
||||
)
|
||||
.expect_err("impacted overflow must fail with no partial plan");
|
||||
match err {
|
||||
Error::InvalidInput { message } => {
|
||||
assert_eq!(message, "dependency_epoch overflow");
|
||||
}
|
||||
other => panic!("expected InvalidInput overflow, got {other:?}"),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,142 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Immutable RefreshGeneratedColumnJobSpec refresh-generated-column Job
|
||||
//! operation input (FF-010).
|
||||
//!
|
||||
//! This type is Job operation input only. It does not look up catalogs or
|
||||
//! tables, execute Jobs, stage artifacts, call Lance, or mutate epochs.
|
||||
|
||||
use std::fmt;
|
||||
|
||||
use serde::de::Error as DeError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
|
||||
use super::{Function, GeneratedColumnDefinition, invalid_input};
|
||||
use crate::Result;
|
||||
|
||||
const FORMAT_VERSION_V1: u32 = 1;
|
||||
|
||||
/// Immutable Job operation input for refreshing a generated column (format
|
||||
/// version 1).
|
||||
///
|
||||
/// Semantic field is exactly the nested [`GeneratedColumnDefinition`]. Wire
|
||||
/// keys are exactly `format_version` and `generated_column_definition`.
|
||||
///
|
||||
/// Construction via [`Self::try_new`] validates the nested call against a
|
||||
/// catalog [`Function`]. Structural deserialize does not; execution consumers
|
||||
/// must call [`Self::validate_against`], and later compare the full nested
|
||||
/// definition to current field metadata in the pinned snapshot.
|
||||
///
|
||||
/// Both complete and incomplete definitions are accepted. Status is not a
|
||||
/// constructor or wire restriction.
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct RefreshGeneratedColumnJobSpec {
|
||||
generated_column_definition: GeneratedColumnDefinition,
|
||||
}
|
||||
|
||||
impl RefreshGeneratedColumnJobSpec {
|
||||
/// Create a refresh-generated-column Job operation input.
|
||||
///
|
||||
/// Requires [`crate::function::FunctionCall::validate_against`] to succeed
|
||||
/// for the nested call and `function` before returning (exact Function ID,
|
||||
/// parameter name/order, argument count, and Arrow type equality).
|
||||
pub fn try_new(
|
||||
function: &Function,
|
||||
generated_column_definition: GeneratedColumnDefinition,
|
||||
) -> Result<Self> {
|
||||
generated_column_definition
|
||||
.function_call()
|
||||
.validate_against(function)?;
|
||||
Ok(Self {
|
||||
generated_column_definition,
|
||||
})
|
||||
}
|
||||
|
||||
/// Wire format version (always 1 for this type).
|
||||
pub fn format_version(&self) -> u32 {
|
||||
FORMAT_VERSION_V1
|
||||
}
|
||||
|
||||
/// Nested generated-column definition to refresh.
|
||||
pub fn generated_column_definition(&self) -> &GeneratedColumnDefinition {
|
||||
&self.generated_column_definition
|
||||
}
|
||||
|
||||
/// Validate the nested call against a catalog [`Function`].
|
||||
///
|
||||
/// Structural decode does not perform this check. Execution consumers must
|
||||
/// call this before using the call.
|
||||
pub fn validate_against(&self, function: &Function) -> Result<()> {
|
||||
self.generated_column_definition
|
||||
.function_call()
|
||||
.validate_against(function)
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> RefreshGeneratedColumnJobSpecWire {
|
||||
RefreshGeneratedColumnJobSpecWire {
|
||||
format_version: FORMAT_VERSION_V1,
|
||||
generated_column_definition: self.generated_column_definition.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(wire: RefreshGeneratedColumnJobSpecWire) -> Result<Self> {
|
||||
if wire.format_version != FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported RefreshGeneratedColumnJobSpec format_version {}",
|
||||
wire.format_version
|
||||
)));
|
||||
}
|
||||
Ok(Self {
|
||||
generated_column_definition: wire.generated_column_definition,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for RefreshGeneratedColumnJobSpec {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
let definition = &self.generated_column_definition;
|
||||
let call = definition.function_call();
|
||||
let field_ids: Vec<_> = call
|
||||
.arguments()
|
||||
.iter()
|
||||
.filter_map(|(_, argument)| argument.field_id())
|
||||
.collect();
|
||||
f.debug_struct("RefreshGeneratedColumnJobSpec")
|
||||
.field("output_field_id", &definition.output_field_id())
|
||||
.field("function_id", &call.function_id().as_str())
|
||||
.field("dependency_epoch", &definition.dependency_epoch())
|
||||
.field("materialized_epoch", &definition.materialized_epoch())
|
||||
.field("argument_count", &call.arguments().len())
|
||||
.field("field_ids", &field_ids)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
// Do not derive Debug: nested GeneratedColumnDefinition / FunctionCall may
|
||||
// carry typed literal payloads on the trusted refresh wire.
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct RefreshGeneratedColumnJobSpecWire {
|
||||
format_version: u32,
|
||||
generated_column_definition: GeneratedColumnDefinition,
|
||||
}
|
||||
|
||||
impl Serialize for RefreshGeneratedColumnJobSpec {
|
||||
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
self.to_wire().serialize(serializer)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for RefreshGeneratedColumnJobSpec {
|
||||
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let wire = RefreshGeneratedColumnJobSpecWire::deserialize(deserializer)?;
|
||||
Self::from_wire(wire).map_err(D::Error::custom)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,154 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Immutable RegisterFunctionJobSpec registration Job operation input (B1d / FF-008).
|
||||
//!
|
||||
//! This type is Job operation input only. It does not execute registration,
|
||||
//! upsert into a catalog, mint identity, or manage Job lifecycle.
|
||||
|
||||
use std::fmt;
|
||||
|
||||
use serde::de::Error as DeError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
|
||||
use super::{FunctionDefinition, FunctionId, invalid_input};
|
||||
use crate::Result;
|
||||
|
||||
const FORMAT_VERSION_V1: u32 = 1;
|
||||
|
||||
/// Immutable Job operation input for registering a first-class Function
|
||||
/// (format version 1).
|
||||
///
|
||||
/// `expected_current_function_id` is a precondition only:
|
||||
/// - [`None`] means create-if-absent (no current Function is expected).
|
||||
/// - [`Some`] with an exact opaque [`FunctionId`] means conditional replace of
|
||||
/// that current Function.
|
||||
///
|
||||
/// This type does not perform catalog execution or upsert.
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct RegisterFunctionJobSpec {
|
||||
name: String,
|
||||
definition: FunctionDefinition,
|
||||
expected_current_function_id: Option<FunctionId>,
|
||||
}
|
||||
|
||||
impl RegisterFunctionJobSpec {
|
||||
/// Create a registration Job operation input.
|
||||
///
|
||||
/// Rejects an empty `name`. Nested definition validation is enforced by
|
||||
/// [`FunctionDefinition`]. When `expected_current_function_id` is
|
||||
/// [`Some`], emptiness is enforced by [`FunctionId::try_new`].
|
||||
///
|
||||
/// - `expected_current_function_id = None`: create-if-absent.
|
||||
/// - `expected_current_function_id = Some(id)`: conditional replace of the
|
||||
/// Function with that exact opaque id.
|
||||
pub fn try_new(
|
||||
name: impl Into<String>,
|
||||
definition: FunctionDefinition,
|
||||
expected_current_function_id: Option<FunctionId>,
|
||||
) -> Result<Self> {
|
||||
let name = name.into();
|
||||
if name.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"RegisterFunctionJobSpec name must be non-empty",
|
||||
));
|
||||
}
|
||||
Ok(Self {
|
||||
name,
|
||||
definition,
|
||||
expected_current_function_id,
|
||||
})
|
||||
}
|
||||
|
||||
/// Wire format version (always 1 for this type).
|
||||
pub fn format_version(&self) -> u32 {
|
||||
FORMAT_VERSION_V1
|
||||
}
|
||||
|
||||
/// Catalog Function name to register.
|
||||
pub fn name(&self) -> &str {
|
||||
&self.name
|
||||
}
|
||||
|
||||
/// Nested registration definition (exact FF-007 [`FunctionDefinition`]).
|
||||
pub fn definition(&self) -> &FunctionDefinition {
|
||||
&self.definition
|
||||
}
|
||||
|
||||
/// Precondition on the current Function id.
|
||||
///
|
||||
/// [`None`] is create-if-absent. [`Some`] is conditional replace of that
|
||||
/// exact opaque id.
|
||||
pub fn expected_current_function_id(&self) -> Option<&FunctionId> {
|
||||
self.expected_current_function_id.as_ref()
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> RegisterFunctionJobSpecWire {
|
||||
RegisterFunctionJobSpecWire {
|
||||
format_version: FORMAT_VERSION_V1,
|
||||
name: self.name.clone(),
|
||||
definition: self.definition.clone(),
|
||||
expected_current_function_id: self
|
||||
.expected_current_function_id
|
||||
.as_ref()
|
||||
.map(|id| id.as_str().to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(wire: RegisterFunctionJobSpecWire) -> Result<Self> {
|
||||
if wire.format_version != FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported RegisterFunctionJobSpec format_version {}",
|
||||
wire.format_version
|
||||
)));
|
||||
}
|
||||
let expected_current_function_id = match wire.expected_current_function_id {
|
||||
None => None,
|
||||
Some(id) => Some(FunctionId::try_new(id)?),
|
||||
};
|
||||
Self::try_new(wire.name, wire.definition, expected_current_function_id)
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for RegisterFunctionJobSpec {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("RegisterFunctionJobSpec")
|
||||
.field("name", &self.name)
|
||||
.field("definition", &self.definition)
|
||||
.field(
|
||||
"expected_current_function_id",
|
||||
&self.expected_current_function_id,
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
// Do not derive Debug: nested FunctionDefinition carries Python source and
|
||||
// secret references on the trusted registration wire.
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct RegisterFunctionJobSpecWire {
|
||||
format_version: u32,
|
||||
name: String,
|
||||
definition: FunctionDefinition,
|
||||
expected_current_function_id: Option<String>,
|
||||
}
|
||||
|
||||
impl Serialize for RegisterFunctionJobSpec {
|
||||
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
self.to_wire().serialize(serializer)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for RegisterFunctionJobSpec {
|
||||
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let wire = RegisterFunctionJobSpecWire::deserialize(deserializer)?;
|
||||
Self::from_wire(wire).map_err(D::Error::custom)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Schema admission for caller-authored generated-column definition ingress.
|
||||
//!
|
||||
//! General-purpose table-schema inputs (for example `Database::create_table`
|
||||
//! and Native `add_columns` schema-bearing transforms) must not invent or
|
||||
//! mutate Job-owned `lancedb::generated_column` top-level field metadata. Only
|
||||
//! generated-column create/change/refresh Job publication may create or change
|
||||
//! that reserved key.
|
||||
//!
|
||||
//! This helper checks raw key presence on top-level fields only. It does not
|
||||
//! recurse into nested children, inspect schema-level metadata, decode the
|
||||
//! payload, look up a Function, or validate epochs.
|
||||
|
||||
use arrow_schema::Schema;
|
||||
|
||||
use super::GENERATED_COLUMN_METADATA_KEY;
|
||||
use crate::{Error, Result};
|
||||
|
||||
/// Reject a caller-authored Arrow schema that carries reserved generated-column
|
||||
/// definition metadata on any top-level field.
|
||||
///
|
||||
/// Safe to call at the start of create-table and Native add-columns paths
|
||||
/// before source consumption, namespace mutation, or HTTP.
|
||||
pub fn reject_caller_authored_generated_column_schema(schema: &Schema) -> Result<()> {
|
||||
for field in schema.fields() {
|
||||
if field.metadata().contains_key(GENERATED_COLUMN_METADATA_KEY) {
|
||||
return Err(Error::NotSupported {
|
||||
message: "generated column definitions are owned by create/change/refresh Jobs \
|
||||
and cannot be supplied through general-purpose table schema input"
|
||||
.into(),
|
||||
});
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Conditionally admit an input schema for append vs overwrite.
|
||||
///
|
||||
/// Overwrite is schema replacement and must reject reserved top-level field
|
||||
/// metadata. Append is not schema replacement: caller field metadata is
|
||||
/// discarded by cast-to-table-schema, so reserved input keys remain accepted.
|
||||
pub fn reject_caller_authored_generated_column_schema_on_overwrite(
|
||||
schema: &Schema,
|
||||
is_overwrite: bool,
|
||||
) -> Result<()> {
|
||||
if is_overwrite {
|
||||
reject_caller_authored_generated_column_schema(schema)
|
||||
} else {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
+292
-12
@@ -6,10 +6,124 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde::de::Error as DeError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
use tokio::sync::watch;
|
||||
use tokio::task::{AbortHandle, JoinHandle};
|
||||
|
||||
use crate::error::{Error, JobFailure, Result};
|
||||
use crate::function::Function;
|
||||
|
||||
const JOB_RESULT_FORMAT_VERSION_V1: u32 = 1;
|
||||
|
||||
fn invalid_input(message: impl Into<String>) -> Error {
|
||||
Error::InvalidInput {
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Result value produced by a completed Job (format version 1).
|
||||
///
|
||||
/// This is a non-resource transport value. It is not a Job handle, does not
|
||||
/// observe lifecycle, and does not preserve unknown wire shapes.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
#[non_exhaustive]
|
||||
pub enum JobResult {
|
||||
/// The Job completed without a Function result.
|
||||
None,
|
||||
/// The Job completed with a [`Function`] value.
|
||||
Function(Function),
|
||||
}
|
||||
|
||||
impl JobResult {
|
||||
/// Wire format version (always 1 for this type).
|
||||
pub fn format_version(&self) -> u32 {
|
||||
JOB_RESULT_FORMAT_VERSION_V1
|
||||
}
|
||||
|
||||
/// Borrow the nested [`Function`] when this is [`JobResult::Function`].
|
||||
pub fn function(&self) -> Option<&Function> {
|
||||
match self {
|
||||
Self::None => None,
|
||||
Self::Function(function) => Some(function),
|
||||
}
|
||||
}
|
||||
|
||||
/// Consume this value and return the nested [`Function`] when present.
|
||||
pub fn into_function(self) -> Option<Function> {
|
||||
match self {
|
||||
Self::None => None,
|
||||
Self::Function(function) => Some(function),
|
||||
}
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> JobResultWire {
|
||||
match self {
|
||||
Self::None => JobResultWire::None {
|
||||
format_version: JOB_RESULT_FORMAT_VERSION_V1,
|
||||
},
|
||||
Self::Function(function) => JobResultWire::Function {
|
||||
format_version: JOB_RESULT_FORMAT_VERSION_V1,
|
||||
function: function.clone(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(wire: JobResultWire) -> Result<Self> {
|
||||
match wire {
|
||||
JobResultWire::None { format_version } => {
|
||||
if format_version != JOB_RESULT_FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported JobResult format_version {format_version}"
|
||||
)));
|
||||
}
|
||||
Ok(Self::None)
|
||||
}
|
||||
JobResultWire::Function {
|
||||
format_version,
|
||||
function,
|
||||
} => {
|
||||
if format_version != JOB_RESULT_FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported JobResult format_version {format_version}"
|
||||
)));
|
||||
}
|
||||
Ok(Self::Function(function))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", deny_unknown_fields)]
|
||||
enum JobResultWire {
|
||||
#[serde(rename = "none")]
|
||||
None { format_version: u32 },
|
||||
#[serde(rename = "function")]
|
||||
Function {
|
||||
format_version: u32,
|
||||
function: Function,
|
||||
},
|
||||
}
|
||||
|
||||
impl Serialize for JobResult {
|
||||
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
self.to_wire().serialize(serializer)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for JobResult {
|
||||
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let wire = JobResultWire::deserialize(deserializer)?;
|
||||
Self::from_wire(wire).map_err(D::Error::custom)
|
||||
}
|
||||
}
|
||||
|
||||
/// Backend-specific tracking for an asynchronous operation.
|
||||
#[async_trait]
|
||||
@@ -19,7 +133,7 @@ pub(crate) trait JobHandle: Send + Sync {
|
||||
None
|
||||
}
|
||||
async fn status(&self) -> Result<String>;
|
||||
async fn wait(&self) -> Result<()>;
|
||||
async fn wait(&self) -> Result<JobResult>;
|
||||
async fn cancel(&self) -> Result<()>;
|
||||
}
|
||||
|
||||
@@ -52,7 +166,7 @@ impl Job {
|
||||
}
|
||||
|
||||
/// A job running as a task in this process.
|
||||
pub(crate) fn spawned(task: JoinHandle<Result<()>>) -> Self {
|
||||
pub(crate) fn spawned(task: JoinHandle<Result<JobResult>>) -> Self {
|
||||
Self::new(Box::new(SpawnedJob::new(task)))
|
||||
}
|
||||
|
||||
@@ -81,11 +195,14 @@ impl Job {
|
||||
|
||||
/// Waits until the operation reaches a terminal state.
|
||||
///
|
||||
/// On success, returns the job's [`JobResult`]. Operations that produce no
|
||||
/// resource result yield [`JobResult::None`].
|
||||
///
|
||||
/// Returns [`crate::Error::JobFailed`] if the operation failed and
|
||||
/// [`crate::Error::JobCancelled`] if it was cancelled.
|
||||
pub async fn wait(&self) -> Result<()> {
|
||||
pub async fn wait(&self) -> Result<JobResult> {
|
||||
match &self.handle {
|
||||
None => Ok(()),
|
||||
None => Ok(JobResult::None),
|
||||
Some(handle) => handle.wait().await,
|
||||
}
|
||||
}
|
||||
@@ -105,15 +222,15 @@ impl Job {
|
||||
/// the outcome; [`Error`] is not, so failures share one behind an [`Arc`].
|
||||
#[derive(Clone)]
|
||||
enum Outcome {
|
||||
Succeeded,
|
||||
Succeeded(JobResult),
|
||||
Failed(Arc<Error>),
|
||||
Cancelled,
|
||||
}
|
||||
|
||||
impl Outcome {
|
||||
fn into_result(self) -> Result<()> {
|
||||
fn into_result(self) -> Result<JobResult> {
|
||||
match self {
|
||||
Self::Succeeded => Ok(()),
|
||||
Self::Succeeded(result) => Ok(result),
|
||||
Self::Failed(source) => Err(Error::JobFailed {
|
||||
job_id: None,
|
||||
failure: JobFailure::from_source(source),
|
||||
@@ -132,12 +249,12 @@ struct SpawnedJob {
|
||||
}
|
||||
|
||||
impl SpawnedJob {
|
||||
fn new(task: JoinHandle<Result<()>>) -> Self {
|
||||
fn new(task: JoinHandle<Result<JobResult>>) -> Self {
|
||||
let abort = task.abort_handle();
|
||||
let (tx, outcome) = watch::channel(None);
|
||||
tokio::spawn(async move {
|
||||
let outcome = match task.await {
|
||||
Ok(Ok(())) => Outcome::Succeeded,
|
||||
Ok(Ok(result)) => Outcome::Succeeded(result),
|
||||
Ok(Err(err)) => Outcome::Failed(Arc::new(err)),
|
||||
Err(err) if err.is_cancelled() => Outcome::Cancelled,
|
||||
Err(err) => Outcome::Failed(Arc::new(Error::Runtime {
|
||||
@@ -155,20 +272,20 @@ impl JobHandle for SpawnedJob {
|
||||
async fn status(&self) -> Result<String> {
|
||||
let label = match &*self.outcome.borrow() {
|
||||
None => "running",
|
||||
Some(Outcome::Succeeded) => "finished",
|
||||
Some(Outcome::Succeeded(_)) => "finished",
|
||||
Some(Outcome::Failed(_)) => "failed",
|
||||
Some(Outcome::Cancelled) => "cancelled",
|
||||
};
|
||||
Ok(label.to_string())
|
||||
}
|
||||
|
||||
async fn wait(&self) -> Result<()> {
|
||||
async fn wait(&self) -> Result<JobResult> {
|
||||
let mut outcome = self.outcome.clone();
|
||||
let settled = outcome
|
||||
.wait_for(|outcome| outcome.is_some())
|
||||
.await
|
||||
.map_err(|_| Error::Runtime {
|
||||
message: "index job outcome was dropped before it completed".to_string(),
|
||||
message: "job outcome was dropped before it completed".to_string(),
|
||||
})?
|
||||
.clone()
|
||||
.expect("wait_for returns once an outcome is set");
|
||||
@@ -180,3 +297,166 @@ impl JobHandle for SpawnedJob {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::future::Future;
|
||||
use std::pin::pin;
|
||||
use std::task::{Context, Poll, Waker};
|
||||
|
||||
use arrow_schema::DataType;
|
||||
use tokio::sync::oneshot;
|
||||
|
||||
use super::*;
|
||||
use crate::error::FunctionErrorCode;
|
||||
use crate::function::{
|
||||
Function, FunctionId, FunctionOutput, FunctionParameter, FunctionSignature,
|
||||
};
|
||||
|
||||
fn sample_success_function() -> Function {
|
||||
let id = FunctionId::try_new("fn.exact.local-job-result").expect("valid FunctionId");
|
||||
let signature = FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("x", DataType::Int32)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.expect("valid FunctionSignature");
|
||||
Function::new(id, signature)
|
||||
}
|
||||
|
||||
fn assert_exact_function(actual: &Function, expected: &Function) {
|
||||
assert_eq!(actual.id(), expected.id());
|
||||
assert_eq!(actual.signature(), expected.signature());
|
||||
}
|
||||
|
||||
/// A completed-before-handle local job projects success as None.
|
||||
#[tokio::test]
|
||||
async fn local_job_result_new_done_wait_returns_none() {
|
||||
let job = Job::new_done();
|
||||
let result = job.wait().await.expect("new_done must succeed");
|
||||
assert_eq!(result, JobResult::None);
|
||||
}
|
||||
|
||||
/// A local spawned unit / no-resource success projects as None.
|
||||
#[tokio::test]
|
||||
async fn local_job_result_spawned_unit_success_projects_none() {
|
||||
let job = Job::spawned(tokio::spawn(async { Ok(JobResult::None) }));
|
||||
let result = job
|
||||
.wait()
|
||||
.await
|
||||
.expect("unit success must finish without error");
|
||||
assert_eq!(result, JobResult::None);
|
||||
}
|
||||
|
||||
/// Function success is cloneable and shared by concurrent + late waiters.
|
||||
///
|
||||
/// Wait futures are pinned and polled once to Pending while success is still
|
||||
/// gated, proving they observed the running state before publication.
|
||||
#[tokio::test]
|
||||
async fn local_job_result_spawned_function_shared_by_waiters() {
|
||||
let expected = sample_success_function();
|
||||
let (release_tx, release_rx) = oneshot::channel();
|
||||
|
||||
let job = Job::spawned(tokio::spawn({
|
||||
let function = expected.clone();
|
||||
async move {
|
||||
release_rx
|
||||
.await
|
||||
.expect("success task must be released by the test");
|
||||
Ok(JobResult::Function(function))
|
||||
}
|
||||
}));
|
||||
|
||||
let mut wait_a = pin!(job.wait());
|
||||
let mut wait_b = pin!(job.wait());
|
||||
let waker = Waker::noop();
|
||||
let mut cx = Context::from_waker(waker);
|
||||
|
||||
assert!(
|
||||
matches!(wait_a.as_mut().poll(&mut cx), Poll::Pending),
|
||||
"waiter A must poll Pending before success publication"
|
||||
);
|
||||
assert!(
|
||||
matches!(wait_b.as_mut().poll(&mut cx), Poll::Pending),
|
||||
"waiter B must poll Pending before success publication"
|
||||
);
|
||||
|
||||
release_tx
|
||||
.send(())
|
||||
.expect("success task must still be waiting on the gate");
|
||||
|
||||
let result_a = wait_a
|
||||
.await
|
||||
.expect("concurrent waiter A must observe success");
|
||||
let result_b = wait_b
|
||||
.await
|
||||
.expect("concurrent waiter B must observe success");
|
||||
let result_late = job
|
||||
.wait()
|
||||
.await
|
||||
.expect("late waiter must observe the same success");
|
||||
|
||||
for result in [&result_a, &result_b, &result_late] {
|
||||
match result {
|
||||
JobResult::Function(function) => assert_exact_function(function, &expected),
|
||||
JobResult::None => panic!("Function success must not project as JobResult::None"),
|
||||
}
|
||||
}
|
||||
assert_eq!(result_a, result_b);
|
||||
assert_eq!(result_a, result_late);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn spawned_job_function_failure_returns_job_failed_with_same_code() {
|
||||
let job = Job::spawned(tokio::spawn(async {
|
||||
Err(Error::Function {
|
||||
code: FunctionErrorCode::UdfExecutionFailure,
|
||||
// Message names a different category on purpose; code is structural.
|
||||
message: "looks like name_conflict to a string parser".to_string(),
|
||||
})
|
||||
}));
|
||||
|
||||
let err = job
|
||||
.wait()
|
||||
.await
|
||||
.expect_err("Function failure must fail the job");
|
||||
match err {
|
||||
Error::JobFailed { failure, .. } => match &failure.error_code {
|
||||
Some(code) => {
|
||||
assert_eq!(code, &FunctionErrorCode::UdfExecutionFailure);
|
||||
assert_ne!(code, &FunctionErrorCode::NameConflict);
|
||||
}
|
||||
None => panic!("local Function failure must project error_code onto JobFailure"),
|
||||
},
|
||||
other => panic!("expected Error::JobFailed, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn spawned_job_preserves_unrecognized_function_error_code() {
|
||||
let raw = "enterprise_future_category_xyz";
|
||||
let job = Job::spawned(tokio::spawn({
|
||||
let raw = raw.to_string();
|
||||
async move {
|
||||
Err(Error::Function {
|
||||
code: FunctionErrorCode::Unrecognized(raw),
|
||||
message: "future server category".to_string(),
|
||||
})
|
||||
}
|
||||
}));
|
||||
|
||||
let err = job
|
||||
.wait()
|
||||
.await
|
||||
.expect_err("Function failure must fail the job");
|
||||
match err {
|
||||
Error::JobFailed { failure, .. } => match &failure.error_code {
|
||||
Some(FunctionErrorCode::Unrecognized(preserved)) => {
|
||||
assert_eq!(preserved, raw);
|
||||
}
|
||||
Some(other) => panic!("unrecognized code must not become known: {other:?}"),
|
||||
None => panic!("unrecognized Function code must be preserved on JobFailure"),
|
||||
},
|
||||
other => panic!("expected Error::JobFailed, got {other:?}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -181,6 +181,7 @@ pub mod dataloader;
|
||||
pub mod embeddings;
|
||||
pub mod error;
|
||||
pub mod expr;
|
||||
pub mod function;
|
||||
pub mod index;
|
||||
pub mod io;
|
||||
pub mod ipc;
|
||||
@@ -205,7 +206,7 @@ use serde::{Deserialize, Serialize};
|
||||
pub use blob::{BlobRangeRequest, blob, is_blob};
|
||||
pub use connection::{ConnectNamespaceBuilder, Connection};
|
||||
pub use error::{Error, JobFailure, Result};
|
||||
pub use job::Job;
|
||||
pub use job::{Job, JobResult};
|
||||
use lance_index::vector::ApproxMode as LanceApproxMode;
|
||||
use lance_linalg::distance::DistanceType as LanceDistanceType;
|
||||
/// Re-export of the [`metrics`](https://docs.rs/metrics) crate facade. Enable
|
||||
@@ -214,7 +215,7 @@ use lance_linalg::distance::DistanceType as LanceDistanceType;
|
||||
/// a built-in pull-based adapter.
|
||||
#[cfg(feature = "metrics")]
|
||||
pub use metrics;
|
||||
pub use table::{FtsToken, Table, TableBase};
|
||||
pub use table::{FtsToken, Table};
|
||||
|
||||
/// Tokenize a full-text search query using an explicit FTS tokenizer configuration.
|
||||
///
|
||||
|
||||
+1345
-3
File diff suppressed because it is too large
Load Diff
@@ -8,10 +8,12 @@
|
||||
|
||||
pub(crate) mod client;
|
||||
pub(crate) mod db;
|
||||
pub(crate) mod function;
|
||||
pub(crate) mod job;
|
||||
pub mod oauth;
|
||||
mod retry;
|
||||
pub(crate) mod table;
|
||||
mod transport;
|
||||
pub(crate) mod util;
|
||||
|
||||
const ARROW_STREAM_CONTENT_TYPE: &str = "application/vnd.apache.arrow.stream";
|
||||
@@ -19,15 +21,6 @@ const ARROW_FILE_CONTENT_TYPE: &str = "application/vnd.apache.arrow.file";
|
||||
#[cfg(test)]
|
||||
const JSON_CONTENT_TYPE: &str = "application/json";
|
||||
|
||||
fn extract_job_id(body: &str) -> Option<String> {
|
||||
serde_json::from_str::<serde_json::Value>(body)
|
||||
.ok()?
|
||||
.get("job_id")?
|
||||
.as_str()
|
||||
.filter(|job_id| !job_id.is_empty())
|
||||
.map(str::to_string)
|
||||
}
|
||||
|
||||
pub use client::{ClientConfig, HeaderProvider, RetryConfig, TimeoutConfig, TlsConfig};
|
||||
pub use db::{RemoteDatabaseOptions, RemoteDatabaseOptionsBuilder};
|
||||
pub use oauth::{OAuthConfig, OAuthFlow, OAuthHeaderProvider};
|
||||
|
||||
@@ -15,6 +15,51 @@ use crate::remote::retry::{ResolvedRetryConfig, RetryCounter};
|
||||
|
||||
const REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-request-id");
|
||||
|
||||
/// Privacy mode for request logging and non-success response handling.
|
||||
///
|
||||
/// [`RequestPrivacy::Standard`] preserves the existing harmless JSON body
|
||||
/// visibility. [`RequestPrivacy::Sensitive`] never includes request bodies or
|
||||
/// headers in logs, and never folds response bodies into error chains.
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
enum RequestPrivacy {
|
||||
Standard,
|
||||
Sensitive,
|
||||
}
|
||||
|
||||
/// Format a request for debug logging according to [`RequestPrivacy`].
|
||||
fn format_request_log(request: &Request, request_id: &str, privacy: RequestPrivacy) -> String {
|
||||
match privacy {
|
||||
RequestPrivacy::Standard => {
|
||||
let content_type = request
|
||||
.headers()
|
||||
.get("content-type")
|
||||
.and_then(|v| v.to_str().ok());
|
||||
if content_type == Some("application/json") {
|
||||
let body = request
|
||||
.body()
|
||||
.and_then(|b| b.as_bytes())
|
||||
.map(|b| String::from_utf8_lossy(b).into_owned())
|
||||
.unwrap_or_default();
|
||||
format!(
|
||||
"Sending request_id={}: {:?} with body {}",
|
||||
request_id, request, body
|
||||
)
|
||||
} else {
|
||||
format!("Sending request_id={}: {:?}", request_id, request)
|
||||
}
|
||||
}
|
||||
RequestPrivacy::Sensitive => {
|
||||
// Safe context only: request id, method, and URL. Never body or headers.
|
||||
format!(
|
||||
"Sending request_id={}: {} {}",
|
||||
request_id,
|
||||
request.method(),
|
||||
request.url()
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Configuration for TLS/mTLS settings.
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct TlsConfig {
|
||||
@@ -746,6 +791,41 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
Ok((request_id, response))
|
||||
}
|
||||
|
||||
/// Send one attempt with a caller-owned request id.
|
||||
///
|
||||
/// Shared by explicit-`error_code` classifiers for Function catalog and
|
||||
/// Remote table query routes. Keeps the caller-owned request ID, uses
|
||||
/// sensitive logging, applies dynamic headers, sends one uninterpreted
|
||||
/// attempt, and leaves status/body classification to the caller.
|
||||
pub(crate) async fn send_attempt_with_request_id(
|
||||
&self,
|
||||
req_builder: RequestBuilder,
|
||||
request_id: &str,
|
||||
) -> Result<Response> {
|
||||
let (client, request) = req_builder.build_split();
|
||||
let mut request = request.map_err(|e| Error::Runtime {
|
||||
message: format!("Failed to build request: {}", e),
|
||||
})?;
|
||||
self.set_request_id(&mut request, request_id);
|
||||
request = self.apply_dynamic_headers(request).await?;
|
||||
if log::log_enabled!(log::Level::Debug) {
|
||||
debug!(
|
||||
"{}",
|
||||
format_request_log(&request, request_id, RequestPrivacy::Sensitive)
|
||||
);
|
||||
}
|
||||
let response = self
|
||||
.sender
|
||||
.send(&client, request)
|
||||
.await
|
||||
.err_to_http(request_id.to_string())?;
|
||||
debug!(
|
||||
"Received response for request_id={}: {:?}",
|
||||
request_id, response
|
||||
);
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
/// Send the request using retries configured in the RetryConfig.
|
||||
/// If retry_5xx is false, 5xx requests will not be retried regardless of the statuses configured
|
||||
/// in the RetryConfig.
|
||||
@@ -753,9 +833,37 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
pub async fn send_with_retry(
|
||||
&self,
|
||||
req_builder: RequestBuilder,
|
||||
mut make_body: Option<Box<dyn FnMut() -> Result<Body> + Send + 'static>>,
|
||||
make_body: Option<Box<dyn FnMut() -> Result<Body> + Send + 'static>>,
|
||||
retry_5xx: bool,
|
||||
) -> Result<(String, Response)> {
|
||||
self.send_with_retry_inner(req_builder, make_body, retry_5xx, RequestPrivacy::Standard)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Like [`Self::send_with_retry`], but never logs request bodies/headers and
|
||||
/// never folds non-success response bodies into retry or HTTP error chains.
|
||||
///
|
||||
/// Privacy affects only logging and error-body exposure; retry budgets are
|
||||
/// identical to [`Self::send_with_retry`] for the same [`RetryConfig`].
|
||||
pub(crate) async fn send_sensitive_with_retry(
|
||||
&self,
|
||||
req_builder: RequestBuilder,
|
||||
make_body: Option<Box<dyn FnMut() -> Result<Body> + Send + 'static>>,
|
||||
retry_5xx: bool,
|
||||
) -> Result<(String, Response)> {
|
||||
self.send_with_retry_inner(req_builder, make_body, retry_5xx, RequestPrivacy::Sensitive)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn send_with_retry_inner(
|
||||
&self,
|
||||
req_builder: RequestBuilder,
|
||||
mut make_body: Option<Box<dyn FnMut() -> Result<Body> + Send + 'static>>,
|
||||
retry_5xx: bool,
|
||||
privacy: RequestPrivacy,
|
||||
) -> Result<(String, Response)> {
|
||||
// Privacy must never alter retry budgets: both Standard and Sensitive
|
||||
// share the same ResolvedRetryConfig / RetryCounter semantics.
|
||||
let retry_config = &self.retry_config;
|
||||
let non_5xx_statuses = retry_config
|
||||
.statuses
|
||||
@@ -772,6 +880,7 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
let mut r = r.map_err(|e| Error::Runtime {
|
||||
message: format!("Failed to build request: {}", e),
|
||||
})?;
|
||||
// One SDK-generated request id is reused across every retry attempt.
|
||||
let request_id = self.extract_request_id(&mut r);
|
||||
let mut retry_counter = RetryCounter::new(retry_config, request_id.clone());
|
||||
|
||||
@@ -790,12 +899,14 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
let mut request = request.map_err(|e| Error::Runtime {
|
||||
message: format!("Failed to build request: {}", e),
|
||||
})?;
|
||||
self.set_request_id(&mut request, &request_id.clone());
|
||||
self.set_request_id(&mut request, &request_id);
|
||||
|
||||
// Apply dynamic headers before each retry attempt
|
||||
request = self.apply_dynamic_headers(request).await?;
|
||||
|
||||
self.log_request(&request, &request_id);
|
||||
if log::log_enabled!(log::Level::Debug) {
|
||||
debug!("{}", format_request_log(&request, &request_id, privacy));
|
||||
}
|
||||
|
||||
let response = self.sender.send(&c, request).await.map(|r| (r.status(), r));
|
||||
|
||||
@@ -811,10 +922,16 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
if (retry_5xx && retry_config.statuses.contains(&status))
|
||||
|| non_5xx_statuses.contains(&status) =>
|
||||
{
|
||||
let source = self
|
||||
.check_response(&retry_counter.request_id, response)
|
||||
.await
|
||||
.unwrap_err();
|
||||
let source = match privacy {
|
||||
RequestPrivacy::Standard => self
|
||||
.check_response(&retry_counter.request_id, response)
|
||||
.await
|
||||
.unwrap_err(),
|
||||
RequestPrivacy::Sensitive => self
|
||||
.check_sensitive_response(&retry_counter.request_id, response)
|
||||
.await
|
||||
.unwrap_err(),
|
||||
};
|
||||
retry_counter.increment_request_failures(source)?;
|
||||
}
|
||||
Err(err) if err.is_connect() => {
|
||||
@@ -839,22 +956,12 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn log_request(&self, request: &Request, request_id: &String) {
|
||||
pub(crate) fn log_request(&self, request: &Request, request_id: &str) {
|
||||
if log::log_enabled!(log::Level::Debug) {
|
||||
let content_type = request
|
||||
.headers()
|
||||
.get("content-type")
|
||||
.map(|v| v.to_str().unwrap());
|
||||
if content_type == Some("application/json") {
|
||||
let body = request.body().as_ref().unwrap().as_bytes().unwrap();
|
||||
let body = String::from_utf8_lossy(body);
|
||||
debug!(
|
||||
"Sending request_id={}: {:?} with body {}",
|
||||
request_id, request, body
|
||||
);
|
||||
} else {
|
||||
debug!("Sending request_id={}: {:?}", request_id, request);
|
||||
}
|
||||
debug!(
|
||||
"{}",
|
||||
format_request_log(request, request_id, RequestPrivacy::Standard)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -898,6 +1005,27 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Like [`Self::check_response`], but discards the response body on failure
|
||||
/// so marker-bearing payloads never enter [`Error::Http`] chains.
|
||||
pub(crate) async fn check_sensitive_response(
|
||||
&self,
|
||||
request_id: &str,
|
||||
response: Response,
|
||||
) -> Result<Response> {
|
||||
let status = response.status();
|
||||
if status.is_success() {
|
||||
Ok(response)
|
||||
} else {
|
||||
// Discard the body entirely; never fold it into Error::Http.
|
||||
let _ = response.bytes().await;
|
||||
Err(Error::Http {
|
||||
source: status.to_string().into(),
|
||||
request_id: request_id.into(),
|
||||
status_code: Some(status),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub trait RequestResultExt {
|
||||
@@ -1066,6 +1194,7 @@ pub mod test_utils {
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serial_test::serial;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::time::Duration;
|
||||
|
||||
// Serializes the env-var-mutating tests below: cargo test runs tests in
|
||||
@@ -1664,4 +1793,253 @@ mod tests {
|
||||
}
|
||||
assert!(matches!(err, Error::InvalidInput { .. }), "got: {err:?}");
|
||||
}
|
||||
|
||||
// -------------------------------------------------------------------------
|
||||
// Sensitive-request privacy mode (generic transport; RED until helpers exist)
|
||||
// -------------------------------------------------------------------------
|
||||
|
||||
const PRIVACY_SOURCE_MARKER: &str = "SENSITIVE_PRIVACY_SOURCE_BODY_MARKER_client";
|
||||
const PRIVACY_SECRET_MARKER: &str = "secret://team/client-privacy-token";
|
||||
|
||||
fn privacy_json_request(url: &str, body: &str, request_id: &str) -> Request {
|
||||
reqwest::Client::new()
|
||||
.post(url)
|
||||
.header("content-type", "application/json")
|
||||
.header("x-request-id", request_id)
|
||||
.body(body.to_string())
|
||||
.build()
|
||||
.expect("build privacy fixture request")
|
||||
}
|
||||
|
||||
fn assert_markers_absent(text: &str) {
|
||||
assert!(
|
||||
!text.contains(PRIVACY_SOURCE_MARKER),
|
||||
"source marker must be absent: {text}"
|
||||
);
|
||||
assert!(
|
||||
!text.contains(PRIVACY_SECRET_MARKER),
|
||||
"secret marker must be absent: {text}"
|
||||
);
|
||||
}
|
||||
|
||||
fn error_chain_text(err: &Error) -> String {
|
||||
let mut text = format!("{err}\n{err:?}");
|
||||
let mut current: Option<&(dyn std::error::Error + 'static)> = Some(err);
|
||||
while let Some(e) = current {
|
||||
text.push('\n');
|
||||
text.push_str(&e.to_string());
|
||||
text.push('\n');
|
||||
text.push_str(&format!("{e:?}"));
|
||||
current = e.source();
|
||||
}
|
||||
text
|
||||
}
|
||||
|
||||
/// Standard JSON request logging keeps the current harmless body visibility.
|
||||
#[test]
|
||||
fn format_request_log_standard_retains_harmless_json_body() {
|
||||
let request_id = "req-privacy-standard";
|
||||
let body = r#"{"ok":true,"note":"harmless-visible-body"}"#;
|
||||
let request = privacy_json_request("http://localhost/v1/table/", body, request_id);
|
||||
|
||||
let log = format_request_log(&request, request_id, RequestPrivacy::Standard);
|
||||
|
||||
assert!(
|
||||
log.contains(request_id),
|
||||
"standard log must retain request id: {log}"
|
||||
);
|
||||
assert!(
|
||||
log.contains("POST"),
|
||||
"standard log must retain method: {log}"
|
||||
);
|
||||
assert!(
|
||||
log.contains("/v1/table/"),
|
||||
"standard log must retain URL path: {log}"
|
||||
);
|
||||
assert!(
|
||||
log.contains("harmless-visible-body"),
|
||||
"standard JSON logging must retain body visibility: {log}"
|
||||
);
|
||||
assert!(
|
||||
log.contains(body) || log.contains(r#""note":"harmless-visible-body""#),
|
||||
"standard JSON logging must include the harmless JSON body: {log}"
|
||||
);
|
||||
}
|
||||
|
||||
/// Sensitive JSON formatting redacts the entire body and keeps only safe context.
|
||||
#[test]
|
||||
fn format_request_log_sensitive_redacts_json_body_keeps_safe_context() {
|
||||
let request_id = "req-privacy-sensitive";
|
||||
let body =
|
||||
format!(r#"{{"source":"{PRIVACY_SOURCE_MARKER}","secret":"{PRIVACY_SECRET_MARKER}"}}"#);
|
||||
let request =
|
||||
privacy_json_request("http://localhost/v1/functions/register", &body, request_id);
|
||||
|
||||
let log = format_request_log(&request, request_id, RequestPrivacy::Sensitive);
|
||||
|
||||
assert!(
|
||||
log.contains(request_id),
|
||||
"sensitive log must retain request id: {log}"
|
||||
);
|
||||
assert!(
|
||||
log.contains("POST"),
|
||||
"sensitive log must retain method: {log}"
|
||||
);
|
||||
assert!(
|
||||
log.contains("/v1/functions/register"),
|
||||
"sensitive log must retain URL path: {log}"
|
||||
);
|
||||
assert_markers_absent(&log);
|
||||
assert!(
|
||||
!log.contains(&body),
|
||||
"sensitive JSON formatting must redact the entire body: {log}"
|
||||
);
|
||||
}
|
||||
|
||||
/// Sensitive non-success responses omit the response body from Error::Http text.
|
||||
#[tokio::test]
|
||||
async fn check_sensitive_response_omits_non_success_response_body() {
|
||||
let client = test_utils::client_with_handler(|_| {
|
||||
http::Response::builder().status(200).body("").unwrap()
|
||||
});
|
||||
let response: Response = http::Response::builder()
|
||||
.status(400)
|
||||
.body(format!(
|
||||
"client error echoed {PRIVACY_SOURCE_MARKER} and {PRIVACY_SECRET_MARKER}"
|
||||
))
|
||||
.unwrap()
|
||||
.into();
|
||||
|
||||
let err = client
|
||||
.check_sensitive_response("req-privacy-check", response)
|
||||
.await
|
||||
.expect_err("non-success sensitive response must fail closed");
|
||||
|
||||
assert!(
|
||||
matches!(err, Error::Http { .. }),
|
||||
"expected Error::Http, got {err:?}"
|
||||
);
|
||||
assert_markers_absent(&error_chain_text(&err));
|
||||
}
|
||||
|
||||
/// Sensitive send+retry must not leak request/response markers into retry errors.
|
||||
#[tokio::test]
|
||||
async fn send_sensitive_with_retry_omits_markers_from_exhausted_retry_errors() {
|
||||
let call_count = Arc::new(AtomicUsize::new(0));
|
||||
let counted = call_count.clone();
|
||||
let client = test_utils::client_with_handler_and_config(
|
||||
move |request| {
|
||||
counted.fetch_add(1, Ordering::SeqCst);
|
||||
let body = request.body().and_then(|b| b.as_bytes()).unwrap_or(b"");
|
||||
let body = std::str::from_utf8(body).unwrap_or("");
|
||||
assert!(
|
||||
body.contains(PRIVACY_SOURCE_MARKER) && body.contains(PRIVACY_SECRET_MARKER),
|
||||
"trusted wire body must still carry sensitive fields"
|
||||
);
|
||||
http::Response::builder()
|
||||
.status(500)
|
||||
.body(format!(
|
||||
"server echoed {PRIVACY_SOURCE_MARKER} and {PRIVACY_SECRET_MARKER}"
|
||||
))
|
||||
.unwrap()
|
||||
},
|
||||
ClientConfig {
|
||||
retry_config: RetryConfig {
|
||||
// RetryCounter treats `retries` as max request failures, so
|
||||
// retries=2 yields exactly two transport attempts before Error::Retry.
|
||||
retries: Some(2),
|
||||
backoff_factor: Some(0.0),
|
||||
backoff_jitter: Some(0.0),
|
||||
..Default::default()
|
||||
},
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
|
||||
let payload = serde_json::json!({
|
||||
"source": PRIVACY_SOURCE_MARKER,
|
||||
"secret": PRIVACY_SECRET_MARKER,
|
||||
});
|
||||
let req = client.post("/v1/functions/register").json(&payload);
|
||||
let err = client
|
||||
.send_sensitive_with_retry(req, None, true)
|
||||
.await
|
||||
.expect_err("exhausted sensitive 5xx retries must fail");
|
||||
|
||||
assert!(
|
||||
matches!(err, Error::Retry { .. }),
|
||||
"expected Error::Retry, got {err:?}"
|
||||
);
|
||||
assert_markers_absent(&error_chain_text(&err));
|
||||
assert_eq!(
|
||||
call_count.load(Ordering::SeqCst),
|
||||
2,
|
||||
"RetryCounter max request failures=2 must make exactly two transport attempts"
|
||||
);
|
||||
}
|
||||
|
||||
/// Standard and Sensitive share the same RetryCounter attempt budget.
|
||||
#[tokio::test]
|
||||
async fn send_with_retry_standard_and_sensitive_share_attempt_budget() {
|
||||
async fn exhausted_attempts(sensitive: bool) -> usize {
|
||||
let call_count = Arc::new(AtomicUsize::new(0));
|
||||
let counted = call_count.clone();
|
||||
let client = test_utils::client_with_handler_and_config(
|
||||
move |_| {
|
||||
counted.fetch_add(1, Ordering::SeqCst);
|
||||
http::Response::builder()
|
||||
.status(500)
|
||||
.body(format!(
|
||||
"server echoed {PRIVACY_SOURCE_MARKER} and {PRIVACY_SECRET_MARKER}"
|
||||
))
|
||||
.unwrap()
|
||||
},
|
||||
ClientConfig {
|
||||
retry_config: RetryConfig {
|
||||
retries: Some(2),
|
||||
backoff_factor: Some(0.0),
|
||||
backoff_jitter: Some(0.0),
|
||||
..Default::default()
|
||||
},
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
|
||||
let payload = serde_json::json!({
|
||||
"source": PRIVACY_SOURCE_MARKER,
|
||||
"secret": PRIVACY_SECRET_MARKER,
|
||||
});
|
||||
let req = client.post("/v1/functions/register").json(&payload);
|
||||
let err = if sensitive {
|
||||
client
|
||||
.send_sensitive_with_retry(req, None, true)
|
||||
.await
|
||||
.expect_err("exhausted sensitive 5xx retries must fail")
|
||||
} else {
|
||||
client
|
||||
.send_with_retry(req, None, true)
|
||||
.await
|
||||
.expect_err("exhausted standard 5xx retries must fail")
|
||||
};
|
||||
assert!(
|
||||
matches!(err, Error::Retry { .. }),
|
||||
"expected Error::Retry, got {err:?}"
|
||||
);
|
||||
if sensitive {
|
||||
assert_markers_absent(&error_chain_text(&err));
|
||||
}
|
||||
call_count.load(Ordering::SeqCst)
|
||||
}
|
||||
|
||||
let standard_attempts = exhausted_attempts(false).await;
|
||||
let sensitive_attempts = exhausted_attempts(true).await;
|
||||
assert_eq!(
|
||||
standard_attempts, sensitive_attempts,
|
||||
"Standard and Sensitive must share the same attempt budget for identical RetryConfig"
|
||||
);
|
||||
assert_eq!(
|
||||
standard_attempts, 2,
|
||||
"RetryCounter max request failures=2 must make exactly two transport attempts"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+3274
-156
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,270 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Remote first-class Function catalog wire helpers.
|
||||
//!
|
||||
//! POST `/v1/functions/lookup` resolves a database-scoped name or exact
|
||||
//! [`FunctionId`] to an immutable [`Function`] value. Name is lookup
|
||||
//! indirection only and never becomes part of the returned handle.
|
||||
//!
|
||||
//! POST `/v1/functions/remove` performs a direct synchronous catalog CAS that
|
||||
//! unbinds a name when the caller's observed [`Function`] id still matches.
|
||||
//! This is not a Job, not physical Function deletion, and not revocation.
|
||||
//!
|
||||
//! POST `/v1/functions/revoke` performs a direct synchronous administrator
|
||||
//! catalog set-bit for an exact [`Function`] id. This is not a Job, not name
|
||||
//! removal, not physical deletion, and not Function mutation.
|
||||
|
||||
use reqwest::{RequestBuilder, StatusCode};
|
||||
use serde::Deserialize;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::error::{Error, Result};
|
||||
use crate::function::{Function, FunctionId};
|
||||
|
||||
use super::client::{HttpSend, RestfulLanceDbClient};
|
||||
use super::transport::{
|
||||
BeforeBody, BodyAction, explicit_error_code, post_with_body_classification,
|
||||
};
|
||||
|
||||
const LOOKUP_PATH: &str = "/v1/functions/lookup";
|
||||
const REMOVE_PATH: &str = "/v1/functions/remove";
|
||||
const REVOKE_PATH: &str = "/v1/functions/revoke";
|
||||
|
||||
/// Fixed client diagnostic for [`Error::Function`]. Never carry server text,
|
||||
/// selector values, or response payload bytes.
|
||||
const LOOKUP_FUNCTION_ERROR_MESSAGE: &str = "function lookup failed";
|
||||
|
||||
/// Fixed client diagnostic for protocol / HTTP failures. Never include response
|
||||
/// payload bytes or selector values.
|
||||
const LOOKUP_HTTP_ERROR_MESSAGE: &str = "function lookup request failed";
|
||||
|
||||
/// Fixed client diagnostic for malformed success payloads.
|
||||
const LOOKUP_INVALID_SUCCESS_MESSAGE: &str = "function lookup response missing or invalid function";
|
||||
|
||||
/// Fixed client diagnostic for remove [`Error::Function`]. Never carry server
|
||||
/// text, catalog name, Function id, or response payload bytes.
|
||||
const REMOVE_FUNCTION_ERROR_MESSAGE: &str = "function name removal failed";
|
||||
|
||||
/// Fixed client diagnostic for remove protocol / HTTP failures.
|
||||
const REMOVE_HTTP_ERROR_MESSAGE: &str = "function name removal request failed";
|
||||
|
||||
/// Fixed client diagnostic for revoke [`Error::Function`]. Never carry server
|
||||
/// text, Function id, or response payload bytes.
|
||||
const REVOKE_FUNCTION_ERROR_MESSAGE: &str = "function revocation failed";
|
||||
|
||||
/// Fixed client diagnostic for revoke protocol / HTTP failures.
|
||||
const REVOKE_HTTP_ERROR_MESSAGE: &str = "function revocation request failed";
|
||||
|
||||
/// One exact lookup selector. Exactly one variant is serialized on the wire.
|
||||
pub enum FunctionLookupSelector {
|
||||
Name(String),
|
||||
FunctionId(String),
|
||||
}
|
||||
|
||||
impl FunctionLookupSelector {
|
||||
pub fn by_name(name: impl Into<String>) -> Result<Self> {
|
||||
let name = name.into();
|
||||
if name.is_empty() {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "function lookup name must be non-empty".into(),
|
||||
});
|
||||
}
|
||||
Ok(Self::Name(name))
|
||||
}
|
||||
|
||||
pub fn by_function_id(function_id: &FunctionId) -> Self {
|
||||
Self::FunctionId(function_id.as_str().to_string())
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> Value {
|
||||
match self {
|
||||
Self::Name(name) => serde_json::json!({ "name": name }),
|
||||
Self::FunctionId(function_id) => {
|
||||
serde_json::json!({ "function_id": function_id })
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct LookupSuccessResponse {
|
||||
function: Function,
|
||||
}
|
||||
|
||||
/// Resolve a Function via POST `/v1/functions/lookup`.
|
||||
///
|
||||
/// Transport classification matches [`RestfulLanceDbClient::send_with_retry`]:
|
||||
/// connect → connect_retries; timeout/body/decode (including response-byte
|
||||
/// reads) → read_retries; configured retryable statuses without an explicit
|
||||
/// `error_code` → request retries; all other transport/client errors return
|
||||
/// immediately. An explicit nonempty `error_code` is terminal and wins over
|
||||
/// HTTP status. Request/response payload bytes never enter error chains.
|
||||
pub async fn lookup_function<S: HttpSend>(
|
||||
client: &RestfulLanceDbClient<S>,
|
||||
selector: FunctionLookupSelector,
|
||||
) -> Result<Function> {
|
||||
let req_builder = client.post(LOOKUP_PATH).json(&selector.to_wire());
|
||||
post_with_body_classification(
|
||||
client,
|
||||
req_builder,
|
||||
LOOKUP_HTTP_ERROR_MESSAGE,
|
||||
|_status, _request_id| BeforeBody::ReadBody,
|
||||
|status, bytes, request_id| {
|
||||
if status.is_success() {
|
||||
return BodyAction::Done(decode_lookup_success(bytes, request_id));
|
||||
}
|
||||
if let Some(code) = explicit_error_code(bytes) {
|
||||
return BodyAction::Done(Err(Error::Function {
|
||||
code,
|
||||
message: LOOKUP_FUNCTION_ERROR_MESSAGE.to_string(),
|
||||
}));
|
||||
}
|
||||
if client.retry_config.statuses.contains(&status) {
|
||||
return BodyAction::RetryRequest;
|
||||
}
|
||||
BodyAction::Done(Err(Error::Http {
|
||||
source: LOOKUP_HTTP_ERROR_MESSAGE.into(),
|
||||
request_id,
|
||||
status_code: Some(status),
|
||||
}))
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Conditionally remove a database-scoped Function name via POST
|
||||
/// `/v1/functions/remove`.
|
||||
///
|
||||
/// Direct synchronous catalog CAS: the wire body is exactly
|
||||
/// `{"name","expected_current_function_id"}` using only `current.id`. Only
|
||||
/// HTTP 204 means the CAS completed; other 2xx are payload-free protocol
|
||||
/// [`Error::Http`]. Empty names are [`Error::InvalidInput`] before transport.
|
||||
///
|
||||
/// Retry budgets match lookup: stable internal request id and exact cloned
|
||||
/// body across attempts; response-byte failures consume read budget; configured
|
||||
/// retryable status without explicit `error_code` consumes request budget;
|
||||
/// header/client errors are immediate. Sophon deduplicates the internal request
|
||||
/// id; it is not a user-facing idempotency key.
|
||||
pub async fn remove_function_name<S: HttpSend>(
|
||||
client: &RestfulLanceDbClient<S>,
|
||||
name: &str,
|
||||
current: &Function,
|
||||
) -> Result<()> {
|
||||
if name.is_empty() {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "function name removal name must be non-empty".into(),
|
||||
});
|
||||
}
|
||||
|
||||
// Authority is the observed immutable Function id only; never send
|
||||
// signature, raw Function objects, Job fields, or user idempotency keys.
|
||||
let body = serde_json::json!({
|
||||
"name": name,
|
||||
"expected_current_function_id": current.id().as_str(),
|
||||
});
|
||||
let req_builder = client.post(REMOVE_PATH).json(&body);
|
||||
|
||||
catalog_mutation_with_retry(
|
||||
client,
|
||||
req_builder,
|
||||
REMOVE_HTTP_ERROR_MESSAGE,
|
||||
REMOVE_FUNCTION_ERROR_MESSAGE,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Revoke an exact Function via POST `/v1/functions/revoke`.
|
||||
///
|
||||
/// Direct synchronous administrator catalog set-bit: the wire body is exactly
|
||||
/// `{"function_id"}` from `function.id`. Only HTTP 204 means the set-bit
|
||||
/// completed; other 2xx are payload-free protocol [`Error::Http`] and are not
|
||||
/// retried or body-read. There is no empty-input validation because
|
||||
/// [`Function`] is already a validated exact handle.
|
||||
///
|
||||
/// Retry and explicit-code classification match remove. Sophon owns durable
|
||||
/// idempotent set-bit semantics; repeated logical calls that each receive 204
|
||||
/// succeed with no client already-revoked branch.
|
||||
pub async fn revoke_function<S: HttpSend>(
|
||||
client: &RestfulLanceDbClient<S>,
|
||||
function: &Function,
|
||||
) -> Result<()> {
|
||||
let body = serde_json::json!({
|
||||
"function_id": function.id().as_str(),
|
||||
});
|
||||
let req_builder = client.post(REVOKE_PATH).json(&body);
|
||||
|
||||
catalog_mutation_with_retry(
|
||||
client,
|
||||
req_builder,
|
||||
REVOKE_HTTP_ERROR_MESSAGE,
|
||||
REVOKE_FUNCTION_ERROR_MESSAGE,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Shared remove/revoke catalog-mutation response classification.
|
||||
///
|
||||
/// Exact HTTP 204 succeeds without reading the body. Other 2xx are immediate
|
||||
/// payload-free [`Error::Http`]. Non-success bodies use explicit nonempty
|
||||
/// `error_code` as terminal [`Error::Function`], else configured retryable
|
||||
/// status retry, else payload-free [`Error::Http`]. Each caller supplies its
|
||||
/// own fixed sanitized messages.
|
||||
async fn catalog_mutation_with_retry<S: HttpSend>(
|
||||
client: &RestfulLanceDbClient<S>,
|
||||
req_builder: RequestBuilder,
|
||||
http_error_message: &'static str,
|
||||
function_error_message: &'static str,
|
||||
) -> Result<()> {
|
||||
post_with_body_classification(
|
||||
client,
|
||||
req_builder,
|
||||
http_error_message,
|
||||
|status, request_id| {
|
||||
// Exact HTTP 204 completes the mutation; do not read or interpret any body.
|
||||
if status == StatusCode::NO_CONTENT {
|
||||
BeforeBody::Done(Ok(()))
|
||||
} else if status.is_success() {
|
||||
// Other 2xx are payload-free protocol failures from status alone.
|
||||
BeforeBody::Done(Err(Error::Http {
|
||||
source: http_error_message.into(),
|
||||
request_id: request_id.to_string(),
|
||||
status_code: Some(status),
|
||||
}))
|
||||
} else {
|
||||
BeforeBody::ReadBody
|
||||
}
|
||||
},
|
||||
|status, bytes, request_id| {
|
||||
// Explicit nonempty error_code wins over HTTP status and precludes retry.
|
||||
if let Some(code) = explicit_error_code(bytes) {
|
||||
return BodyAction::Done(Err(Error::Function {
|
||||
code,
|
||||
message: function_error_message.to_string(),
|
||||
}));
|
||||
}
|
||||
|
||||
if client.retry_config.statuses.contains(&status) {
|
||||
return BodyAction::RetryRequest;
|
||||
}
|
||||
|
||||
BodyAction::Done(Err(Error::Http {
|
||||
source: http_error_message.into(),
|
||||
request_id,
|
||||
status_code: Some(status),
|
||||
}))
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn decode_lookup_success(bytes: &[u8], request_id: String) -> Result<Function> {
|
||||
match serde_json::from_slice::<LookupSuccessResponse>(bytes) {
|
||||
Ok(body) => Ok(body.function),
|
||||
Err(_) => Err(Error::Http {
|
||||
source: LOOKUP_INVALID_SUCCESS_MESSAGE.into(),
|
||||
request_id,
|
||||
status_code: None,
|
||||
}),
|
||||
}
|
||||
}
|
||||
+962
-19
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user