Compare commits

...

65 Commits

Author SHA1 Message Date
Xuanwo 3a3ddfda01 test: verify function state survives server restart 2026-08-14 16:55:38 +08:00
Xuanwo c4371eb500 test: add enterprise function reliability e2e 2026-08-14 16:19:35 +08:00
Xuanwo 7a08580400 fix(remote): project generated column job results 2026-08-14 14:27:43 +08:00
Xuanwo 2fccab172f test: add enterprise first-class function e2e 2026-08-14 12:10:33 +08:00
Xuanwo e478b80985 feat: author generated column change jobs 2026-08-12 20:34:39 +08:00
Xuanwo 72767b17fa feat: project definitions from binding snapshots 2026-08-12 20:13:04 +08:00
Xuanwo d7d25cd5ef feat: author generated column refresh jobs 2026-08-12 20:01:09 +08:00
Xuanwo 4843445a7e feat: project generated column definitions 2026-08-12 19:42:36 +08:00
Xuanwo 713510375b feat: resolve generated column functions by id 2026-08-12 19:34:16 +08:00
Xuanwo a9617bf830 feat: submit generated column refresh jobs 2026-08-12 19:25:18 +08:00
Xuanwo 208787ae7b feat: submit generated column change jobs 2026-08-12 19:18:55 +08:00
Xuanwo 9b46e7a448 feat: guard generated metadata on overwrite 2026-08-12 19:08:13 +08:00
Xuanwo 7a46da2e67 feat: guard generated metadata on add columns 2026-08-12 18:52:21 +08:00
Xuanwo 705f7e7760 feat: guard generated metadata on table creation 2026-08-12 18:41:58 +08:00
Xuanwo d38a566282 feat: reserve generated column metadata updates 2026-08-12 18:25:31 +08:00
Xuanwo 5a27c71ab8 feat: reject merge insert on generated columns 2026-08-12 18:12:30 +08:00
Xuanwo 76f6487d92 feat: invalidate generated columns on native delete 2026-08-12 17:57:26 +08:00
Xuanwo 6f6c3c33e0 build: pin Lance zero-row delete attachment fix 2026-08-12 17:31:34 +08:00
Xuanwo 1aa3665d67 feat: invalidate generated columns on native update 2026-08-12 17:03:31 +08:00
Xuanwo 0194f2317a build: pin Lance no-op update attachment fix 2026-08-12 16:44:38 +08:00
Xuanwo c2a647189d feat: invalidate generated columns on native append 2026-08-12 16:14:22 +08:00
Xuanwo df67ee4028 chore: pin Lance A4 substrate 2026-08-12 15:49:25 +08:00
Xuanwo d9e41228c8 feat: plan generated column invalidation 2026-08-12 15:11:25 +08:00
Xuanwo 68597070d2 feat: project remote query function errors 2026-08-12 14:49:01 +08:00
Xuanwo c825737780 feat: guard native generated column queries 2026-08-12 14:14:41 +08:00
Xuanwo 0a43795996 feat: add generated column query guard analysis 2026-08-12 13:45:13 +08:00
Xuanwo 0fa2fa05ad feat(python): expose generated column status 2026-08-12 12:52:27 +08:00
Xuanwo 93ba442ac2 feat: expose generated column status 2026-08-12 12:17:02 +08:00
Xuanwo 7a94ab7d6c feat: add generated column Python API 2026-08-12 11:50:54 +08:00
Xuanwo 6ed1a25439 feat: submit generated column creation jobs 2026-08-12 10:58:04 +08:00
Xuanwo ca1d04db25 feat(python): bind function calls to table snapshots 2026-08-12 10:31:46 +08:00
Xuanwo efe3300404 feat: validate bound function call fields 2026-08-12 10:05:42 +08:00
Xuanwo ecf87f6371 feat: add atomic generated column binding snapshots 2026-08-12 09:49:51 +08:00
Xuanwo 47213e31f8 feat(python): add first-class function call authoring 2026-08-12 09:19:49 +08:00
Xuanwo f65bf89c98 feat(python): add exact function revocation 2026-08-12 08:22:49 +08:00
Xuanwo d902144605 feat(rust): add exact function revocation 2026-08-12 08:06:30 +08:00
Xuanwo a49dc5c71d feat(python): add conditional function name removal 2026-08-12 07:56:07 +08:00
Xuanwo 98fed41efa feat(rust): add conditional function name removal 2026-08-12 07:33:19 +08:00
Xuanwo 1524ee0669 feat(python): add conditional function replacement 2026-08-12 07:02:38 +08:00
Xuanwo 29be3e5509 feat(python): expose function job error codes 2026-08-12 06:44:03 +08:00
Xuanwo 8cedd50495 feat: expose function lookup in Python 2026-08-12 06:25:33 +08:00
Xuanwo b71ada0fae feat: add function catalog lookup 2026-08-12 06:06:18 +08:00
Xuanwo 206efd98ff feat: register functions from Python 2026-08-12 05:34:28 +08:00
Xuanwo 65c0968c0f feat: submit function registration jobs 2026-08-12 05:03:17 +08:00
Xuanwo 2b10f2a7ce feat: bridge Python UDF definitions to Rust 2026-08-12 04:34:58 +08:00
Xuanwo f8bb90405f feat: declare Python function capabilities 2026-08-12 03:57:52 +08:00
Xuanwo 76aac96749 feat: validate Python UDF source packages 2026-08-12 03:47:17 +08:00
Xuanwo 0093bc8179 feat: add Python UDF declarations 2026-08-12 03:28:05 +08:00
Xuanwo ac35a687f1 feat: expose Python function job results 2026-08-12 03:12:36 +08:00
Xuanwo 203f6536a6 feat: expose typed remote job results 2026-08-12 02:09:33 +08:00
Xuanwo 9d3d0d0640 feat: decode remote job results 2026-08-12 01:39:51 +08:00
Xuanwo a9ed8dba27 feat: return results from jobs 2026-08-12 01:18:41 +08:00
Xuanwo 04acf1d3b5 feat: add first-class function job result 2026-08-12 00:43:18 +08:00
Xuanwo 3746118374 feat: add generated column change job spec 2026-08-12 00:26:01 +08:00
Xuanwo d0b5cbe510 feat: add generated column refresh job spec 2026-08-12 00:09:57 +08:00
Xuanwo 7b195adc3a feat: add generated column create job spec 2026-08-11 23:44:27 +08:00
Xuanwo 818d6d1f59 feat: add function registration job spec 2026-08-11 23:26:05 +08:00
Xuanwo 9d589bea44 feat: add function definition contract 2026-08-11 23:11:01 +08:00
Xuanwo 1798ece362 feat: add stable function error codes 2026-08-11 22:43:39 +08:00
Xuanwo 82b82711ba feat: add first-class function value model 2026-08-11 22:25:09 +08:00
Sravan Avvaru a615306f39 feat(python): add on_transform_error fault tolerance to StreamingDataset (#3763)
Closes #3704

## Problem

Transforms can fail on bad data (e.g. nulls/NaNs from incomplete user
surveys). Today any transform exception aborts iteration, and there is
no way to skip invalid rows during loading.

## Solution

New `on_transform_error` parameter on `StreamingDataset`:

- `"raise"` (default, matches current behavior and the convention in
tf.data / WebDataset / Ray Data)
- `"skip"` — drop the failing rows and continue
- `"warn"` — like skip, plus a logged warning per failing batch
- a WebDataset-style callable `handler(exc) -> bool`, so users can skip
only expected error types

Key design points:

- **Row-granular skipping**: when a batch fails, the transform is re-run
on single-row slices so only the rows that actually fail are dropped
(avoids Ray-style whole-block loss). Skips are counted in a new
`rows_skipped` property.
- **No crash on uneven skips**: the round-robin loop now ends the epoch
at the last cycle where every split still has a row, instead of hitting
`IndexError` when a split runs dry early.
- **Exact resumability under skips**: checkpoints are now
position-based. `state_dict` gains `positions_consumed_per_split` (exact
for owned splits), and a new `merge_state_dicts` static method combines
per-rank states via elementwise max for elastic resume across topology
changes. Old checkpoints without the new key still load. Positions equal
sample counts when nothing is skipped, so existing behavior is
unchanged.
- **Guardrail**: transforms returning the wrong number of rows now raise
a clear `ValueError` instead of silently corrupting split accounting.

### Answers to the issue's open questions

- *Can we do this?* Yes — all transforms funnel through one guarded call
in the Stage 2 pipeline.
- *What do other libraries do?* tf.data `ignore_errors()`, WebDataset
`handler=`, Ray `max_errored_blocks`; MosaicML StreamingDataset offers
nothing (skipping conflicts with its determinism model). This design
follows the common conventions: raise by default, opt-in skipping,
count/log drops.
- *Error handling or pre-filtering?* Both: the existing `filter=`
remains the recommended tool for predictable bad data (splits are built
post-filter, so all guarantees hold — now documented);
`on_transform_error` covers failures not expressible as a predicate.
- *Impact on splits / elastic determinism?* Per-split sample sequences
stay deterministic (skips are data-dependent, not topology-dependent).
With unequal bad-row counts across splits the last few global steps of
an epoch can differ across topologies (bounded by the skew), which is
documented on the parameter. With equal counts per split, full
determinism is preserved — covered by a test.

## Testing

15 new tests in `test_elastic_dataloader.py` covering: default raise,
invalid values, uniform and uneven skips (including epoch-end
truncation), warn logging, selective callable handlers, wrong-row-count
guardrail, determinism across runs and across world sizes (1/2/3/4) with
skips, exact mid-epoch resume with skips on the same topology, elastic
resume via `merge_state_dicts` (ws=2 → ws=1), merge validation, and
backward-compat loading of old checkpoints.

Note: relying on CI for the test run — my local machine OOMs during the
final link of the native extension. The change itself is pure Python.

---------

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
2026-08-10 09:22:06 -07:00
Xuanwo 920fc0e455 fix(python): set native module metadata (#3913)
PyO3 defaults native extension classes to `builtins`, so
mkdocstrings/Griffe could not resolve the newly documented
`lancedb.Session` alias and `Deploy docs to Pages` failed on `main`.
Declare the extension module for the public native types referenced by
the Python API docs so Griffe resolves them through `lancedb._lancedb`
and Pages can build again.

Validated with the docs toolchain used by CI (`griffe==0.49.0`,
`mkdocstrings==0.25.2`, and `mkdocs==1.6.1`); `PYTHONPATH=. mkdocs
build` succeeds.
2026-08-10 21:40:31 +08:00
Xuanwo 5acce6782e ci(docs): report link checker failures through issues (#3909) 2026-08-10 15:08:36 +08:00
ForwardXu 12405a4077 chore: drop explicit goosefs-sdk pin in favor of opendal 0.58.1 transitive dep (#3910)
## Summary

`opendal 0.58.1` (the version pulled in transitively via Lance) already
ships
`goosefs-sdk 0.1.9`, which includes the upstream fix for the 0.1.6
compile
break. The explicit version pin that lancedb has been carrying since the
GooseFS feature was introduced is therefore no longer necessary and is
now
redundant work to maintain.

## Changes

- Remove the direct `goosefs-sdk` dependency from
`rust/lancedb/Cargo.toml`
(it was pinned to `=0.1.9` with a comment referencing the 0.1.6 compile
  break).
- Remove the `dep:goosefs-sdk` entry from the `goosefs` cargo feature,
since
  no source file in lancedb imports the crate directly.
- Refresh `Cargo.lock`; `goosefs-sdk 0.1.9` now resolves transitively
through
  `lance` → `opendal 0.58.1`.

## Verification

- `cargo fmt --all` — clean
- `cargo check --features remote,goosefs --tests --examples` — passes
- `Cargo.lock` confirms `goosefs-sdk 0.1.9` is still resolved (now
transitively), so the `goosefs` feature continues to enable the same set
of
  Lance/IOPaths as before.

## Backwards compatibility

No public API changes. The `goosefs` cargo feature still activates
`lance/goosefs`, `lance-io/goosefs`, and
`lance-namespace-impls/dir-goosefs`,
and the same `goosefs-sdk 0.1.9` version is selected by the resolver.
2026-08-10 12:16:21 +08:00
lancedb-gatefixer[bot] 36054be576 fix(node): preserve nested Arrow data across versions (#3900)
<!-- lance-gatekeeper-fix:v1 agent=613a074d606e626c5169d601373a32d8
generation=1 -->

## Root cause

When LanceDB accepted an Arrow table created by a different installed
Arrow package, its compatibility sanitizer rebuilt each Data node
without converting the foreign type or preserving nested children. It
also dropped the separate dictionary vector payload and did not preserve
identity shared by dictionary schema types, vector wrappers, or growing
dictionary chunks.

## Fix

Recursively sanitize nested Arrow data types and child data. Use one
table-scoped sanitization context to rebuild and memoize source type
objects, dictionary vectors, and Data nodes in the local Arrow realm,
preserving all identities required by Arrow IPC.

Add Arrow 15 through 18 regressions for list serialization, ordinary
dictionaries, dictionaries shared across fields and batches, growing
dictionaries, and IPC round trips.

## Validation

- pnpm test __test__/arrow.test.ts --runInBand (188 passed)
- pnpm lint
- pnpm build
- pnpm test --runInBand (706 passed, 5 skipped)
- pnpm run docs

Fixes #2256

---------

Co-authored-by: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com>
2026-08-09 03:34:39 +08:00
104 changed files with 43728 additions and 420 deletions
+70 -49
View File
@@ -36,7 +36,9 @@ jobs:
permissions:
contents: read
outputs:
checker_outcome: ${{ steps.lychee.outcome }}
exit_code: ${{ steps.lychee.outputs.exit_code }}
status: ${{ steps.validate.outputs.status }}
steps:
- name: Checkout
uses: actions/checkout@v6
@@ -50,6 +52,7 @@ jobs:
- name: Check links
id: lychee
continue-on-error: true
uses: lycheeverse/lychee-action@e7477775783ea5526144ba13e8db5eec57747ce8 # v2.9.0
with:
# Restricted to http(s) on purpose. Much of docs/src is generated
@@ -68,38 +71,50 @@ jobs:
format: json
output: ./lychee/out.json
jobSummary: false
# The report, not a red build, is the signal for broken links. The
# validation step below still fails the run if the check itself
# breaks.
# The report issue, not a red workflow run, is the signal for link
# findings and checker failures alike.
fail: false
- name: Validate report
id: validate
# lychee does not reserve exit code 2 for broken links: its CLI
# parser also exits 2 on an invalid option, before any link was
# checked or any report written. Only a parseable report whose
# counts agree with the exit code counts as a link verdict; anything
# else fails here, and the report job below is skipped entirely, so
# the tracking issue is never touched. Exit 2 covers timeouts as
# well as errors, and a timed-out host is exactly the transient
# unavailability this report exists to surface, so both count as
# findings. Requiring total > 0 also catches a glob that silently
# stopped matching any file.
if: steps.lychee.outputs.exit_code == 0 || steps.lychee.outputs.exit_code == 2
# counts agree with a completed exit code (0 or 2) counts as a link
# verdict. Everything else becomes a checker-error report instead of
# failing the workflow. Exit 2 covers timeouts as well as errors, and a
# timed-out host is exactly the transient unavailability this report
# exists to surface, so both count as findings. Requiring total > 0
# also catches a glob that silently stopped matching any file.
if: always()
env:
CHECKER_OUTCOME: ${{ steps.lychee.outcome }}
EXIT_CODE: ${{ steps.lychee.outputs.exit_code }}
run: |
jq -e --argjson code "$EXIT_CODE" '
(.total > 0) and
(if $code == 0
then .errors == 0 and .timeouts == 0
and (.error_map | length == 0) and (.timeout_map | length == 0)
else (.errors + .timeouts) > 0
and ((.error_map | length) + (.timeout_map | length)) > 0
end)
' ./lychee/out.json
status=checker-error
if [[ "$CHECKER_OUTCOME" == success ]] &&
[[ "$EXIT_CODE" == 0 || "$EXIT_CODE" == 2 ]] &&
jq -e --argjson code "$EXIT_CODE" '
(.total > 0) and
(if $code == 0
then .errors == 0 and .timeouts == 0
and (.error_map | length == 0) and (.timeout_map | length == 0)
else (.errors + .timeouts) > 0
and ((.error_map | length) + (.timeout_map | length)) > 0
end)
' ./lychee/out.json
then
if [[ "$EXIT_CODE" == 0 ]]; then
status=healthy
else
status=findings
fi
fi
echo "status=$status" >> "$GITHUB_OUTPUT"
echo "Validated link check as $status"
- name: Upload report
if: steps.lychee.outputs.exit_code == 2
if: steps.validate.outputs.status == 'findings'
uses: actions/upload-artifact@v7
with:
name: link-report
@@ -115,26 +130,11 @@ jobs:
permissions:
issues: write
env:
CHECKER_OUTCOME: ${{ needs.scan.outputs.checker_outcome }}
EXIT_CODE: ${{ needs.scan.outputs.exit_code }}
STATUS: ${{ needs.scan.outputs.status }}
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
steps:
- name: Classify checker result
# lychee exits 0 when every link resolves and 2 when links fail,
# both already cross-checked against the report by the scan job's
# validation step. Anything else (1 runtime, 3 bad config) means the
# check never produced a link verdict, which must surface as a failed
# run rather than be published as "broken documentation links".
run: |
case "$EXIT_CODE" in
0|2)
echo "lychee exit code $EXIT_CODE"
;;
*)
echo "::error::lychee exited with '$EXIT_CODE': the link check did not complete. Leaving the report issue untouched."
exit 1
;;
esac
- name: Find existing report issue
id: report
# Matched on title alone, and through search rather than a listing:
@@ -144,7 +144,7 @@ jobs:
# Closed issues are included because a healthy run closes the report:
# an open-only lookup would forget that identity and the next failing
# run would open a duplicate. The oldest match stays the canonical
# report and is reopened below when links break again.
# report and is reopened below when a problem recurs.
run: |
match=$(gh issue list --repo "$GITHUB_REPOSITORY" --state all \
--search "in:title \"$REPORT_TITLE\" author:app/github-actions" \
@@ -154,14 +154,14 @@ jobs:
echo "state=$(jq -r '.state // empty' <<<"$match")" >> "$GITHUB_OUTPUT"
- name: Download report
if: env.EXIT_CODE == 2
if: env.STATUS == 'findings'
uses: actions/download-artifact@v8
with:
name: link-report
path: ./lychee
- name: Compose report
if: env.EXIT_CODE == 2
if: env.STATUS == 'findings'
run: |
run_url="$GITHUB_SERVER_URL/$GITHUB_REPOSITORY/actions/runs/$GITHUB_RUN_ID"
{
@@ -185,22 +185,41 @@ jobs:
' ./lychee/out.json
} > ./lychee/issue.md
- name: Compose checker error report
if: env.STATUS == 'checker-error'
run: |
mkdir -p ./lychee
run_url="$GITHUB_SERVER_URL/$GITHUB_REPOSITORY/actions/runs/$GITHUB_RUN_ID"
{
echo "The documentation link check did not complete in [the latest run]($run_url)."
echo
echo "This issue is rewritten by every scheduled run and closed automatically once a trustworthy run finds that all links resolve."
echo
echo "The checker did not produce a trustworthy link verdict. Treat the previous result, if any, as stale until a later run completes."
echo
echo "* Action outcome: \`$CHECKER_OUTCOME\`"
echo "* Exit code: \`${EXIT_CODE:-not reported}\`"
echo "* Verdict validation: \`failed\`"
} > ./lychee/issue.md
- name: Reopen report issue
# A healthy run closes the report, and the issue action below only
# rewrites the body of whatever number it is given. Without an
# explicit reopen, the 2 -> 0 -> 2 sequence would keep rewriting a
# closed issue while links are broken. A CLOSED state implies the
# lookup found a canonical issue, so no separate emptiness check.
if: env.EXIT_CODE == 2 && steps.report.outputs.state == 'CLOSED'
# explicit reopen, a later finding or checker error would rewrite a
# closed issue. A CLOSED state implies the lookup found a canonical
# issue, so no separate emptiness check.
if: >-
env.STATUS != 'healthy' &&
steps.report.outputs.state == 'CLOSED'
env:
ISSUE_NUMBER: ${{ steps.report.outputs.number }}
run: |
run_url="$GITHUB_SERVER_URL/$GITHUB_REPOSITORY/actions/runs/$GITHUB_RUN_ID"
gh issue reopen "$ISSUE_NUMBER" --repo "$GITHUB_REPOSITORY" \
--comment "Broken documentation links found again in [the latest run]($run_url)."
--comment "The documentation link checker reported a problem again in [the latest run]($run_url)."
- name: Report broken links
if: env.EXIT_CODE == 2
- name: Report link-check problem
if: env.STATUS != 'healthy'
uses: peter-evans/create-issue-from-file@fca9117c27cdc29c6c4db3b86c48e4115a786710 # v6.0.0
with:
# Empty on the first failing run, which creates the issue; afterwards
@@ -213,7 +232,9 @@ jobs:
- name: Close report issue once links are healthy
# An OPEN state implies the lookup found a canonical issue; a report
# that is already closed needs nothing.
if: env.EXIT_CODE == 0 && steps.report.outputs.state == 'OPEN'
if: >-
env.STATUS == 'healthy' &&
steps.report.outputs.state == 'OPEN'
env:
ISSUE_NUMBER: ${{ steps.report.outputs.number }}
run: |
Generated
+43 -43
View File
@@ -3455,8 +3455,8 @@ checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c"
[[package]]
name = "fsst"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
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.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arc-swap",
"arrow",
@@ -4890,8 +4890,8 @@ dependencies = [
[[package]]
name = "lance-arrow"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -4913,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.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -4927,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.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-schema",
@@ -4936,8 +4936,8 @@ dependencies = [
[[package]]
name = "lance-bitpacking"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrayref",
"crunchy",
@@ -4947,8 +4947,8 @@ dependencies = [
[[package]]
name = "lance-core"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -4988,8 +4988,8 @@ dependencies = [
[[package]]
name = "lance-datafusion"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow",
"arrow-array",
@@ -5019,8 +5019,8 @@ dependencies = [
[[package]]
name = "lance-datagen"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow",
"arrow-array",
@@ -5037,8 +5037,8 @@ dependencies = [
[[package]]
name = "lance-derive"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"proc-macro2",
"quote",
@@ -5047,8 +5047,8 @@ dependencies = [
[[package]]
name = "lance-encoding"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-arith",
"arrow-array",
@@ -5082,8 +5082,8 @@ dependencies = [
[[package]]
name = "lance-file"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-arith",
"arrow-array",
@@ -5114,8 +5114,8 @@ dependencies = [
[[package]]
name = "lance-index"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arc-swap",
"arrow",
@@ -5182,8 +5182,8 @@ dependencies = [
[[package]]
name = "lance-index-core"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-schema",
@@ -5205,8 +5205,8 @@ dependencies = [
[[package]]
name = "lance-io"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow",
"arrow-array",
@@ -5242,8 +5242,8 @@ dependencies = [
[[package]]
name = "lance-linalg"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -5259,8 +5259,8 @@ dependencies = [
[[package]]
name = "lance-namespace"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow",
"async-trait",
@@ -5272,8 +5272,8 @@ dependencies = [
[[package]]
name = "lance-namespace-impls"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow",
"arrow-ipc",
@@ -5326,8 +5326,8 @@ dependencies = [
[[package]]
name = "lance-select"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -5342,8 +5342,8 @@ dependencies = [
[[package]]
name = "lance-table"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow",
"arrow-array",
@@ -5383,8 +5383,8 @@ dependencies = [
[[package]]
name = "lance-testing"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"arrow-array",
"arrow-schema",
@@ -5397,8 +5397,8 @@ dependencies = [
[[package]]
name = "lance-tokenizer"
version = "11.0.0-beta.3"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
version = "11.0.0-beta.4"
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
dependencies = [
"frostem",
"icu_segmenter",
@@ -5432,6 +5432,7 @@ dependencies = [
"aws-sdk-kms",
"aws-sdk-s3",
"aws-smithy-runtime",
"base64 0.22.1",
"bytes",
"candle-core",
"candle-nn",
@@ -5447,7 +5448,6 @@ dependencies = [
"datafusion-physical-plan",
"datafusion-sql",
"futures",
"goosefs-sdk",
"half",
"hf-hub",
"http 1.5.0",
+14 -14
View File
@@ -13,20 +13,20 @@ categories = ["database-implementations"]
rust-version = "1.91.0"
[workspace.dependencies]
lance = { "version" = "=11.0.0-beta.3", default-features = false, "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-core = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-datagen = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-file = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-io = { "version" = "=11.0.0-beta.3", default-features = false, "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-index = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-linalg = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-namespace = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-namespace-impls = { "version" = "=11.0.0-beta.3", default-features = false, "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-table = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-testing = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-datafusion = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-encoding = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
lance-arrow = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "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 }
+115
View File
@@ -6,7 +6,9 @@ import * as arrow17 from "apache-arrow-17";
import * as arrow18 from "apache-arrow-18";
import {
Vector as CurrentVector,
convertToTable,
tableFromIPC as currentTableFromIPC,
fromBufferToRecordBatch,
fromDataToBuffer,
fromRecordBatchToBuffer,
@@ -19,6 +21,7 @@ import {
FunctionOptions,
} from "../lancedb/embedding/embedding_function";
import { EmbeddingFunctionConfig } from "../lancedb/embedding/registry";
import { sanitizeTable } from "../lancedb/sanitize";
// biome-ignore lint/suspicious/noExplicitAny: skip
function sampleRecords(): Array<Record<string, any>> {
@@ -64,7 +67,11 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
tableFromIPC,
DataType,
Dictionary,
RecordBatch: ArrowRecordBatch,
Table: ArrowTable,
Uint8: ArrowUint8,
makeData: arrowMakeData,
vectorFromArray,
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
} = <any>arrow;
type Schema = ApacheArrow["Schema"];
@@ -1054,6 +1061,114 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
});
describe("when using two versions of arrow", function () {
it("preserves a dictionary shared by multiple fields", async function () {
const values = ["alpha", "beta", "alpha"];
const dictionaryVector = vectorFromArray(values);
const batch = new ArrowRecordBatch({
first: dictionaryVector.data[0],
second: dictionaryVector.data[0],
});
const table = new ArrowTable([batch]);
const sanitized = sanitizeTable(table);
expect([...sanitized.getChild("first")!]).toEqual(values);
expect([...sanitized.getChild("second")!]).toEqual(values);
const firstType = sanitized.schema.fields[0].type as {
dictionary: unknown;
};
const secondType = sanitized.schema.fields[1].type as {
dictionary: unknown;
};
expect(secondType.dictionary).toBe(firstType.dictionary);
expect(sanitized.batches[0].data.children[1].dictionary).toBe(
sanitized.batches[0].data.children[0].dictionary,
);
const buf = await fromDataToBuffer(table);
const actual = currentTableFromIPC(buf);
expect([...actual.getChild("first")!]).toEqual(values);
expect([...actual.getChild("second")!]).toEqual(values);
});
it("preserves shared dictionary data from another Arrow version", async function () {
const values = ["alpha", "beta", "alpha"];
const dictionaryVector = vectorFromArray(values);
const firstBatch = new ArrowRecordBatch({
label: dictionaryVector.slice(0, 2).data[0],
});
const secondBatch = new ArrowRecordBatch({
label: dictionaryVector.slice(2).data[0],
});
const table = new ArrowTable([firstBatch, secondBatch]);
const sanitized = sanitizeTable(table);
expect([...sanitized.getChild("label")!]).toEqual(values);
const dictionaries = sanitized.batches.map(
(batch) => batch.data.children[0].dictionary,
);
expect(dictionaries[0]).toBeInstanceOf(CurrentVector);
expect(dictionaries[1]).toBe(dictionaries[0]);
const buf = await fromDataToBuffer(table);
const actual = currentTableFromIPC(buf);
expect([...actual.getChild("label")!]).toEqual(values);
});
it("preserves shared chunks in growing dictionaries", async function () {
const type = new Dictionary(new Utf8(), new Int32(), 42, false);
const firstDictionary = vectorFromArray(["alpha", "beta"], new Utf8());
const secondDictionary = firstDictionary.concat(
vectorFromArray(["gamma"], new Utf8()),
);
const firstData = arrowMakeData({
type,
data: Int32Array.from([0, 1]),
dictionary: firstDictionary,
});
const secondData = arrowMakeData({
type,
data: Int32Array.from([2]),
dictionary: secondDictionary,
});
const table = new ArrowTable([
new ArrowRecordBatch({ label: firstData }),
new ArrowRecordBatch({ label: secondData }),
]);
const sanitized = sanitizeTable(table);
const expected = ["alpha", "beta", "gamma"];
expect([...sanitized.getChild("label")!]).toEqual(expected);
const firstLocalDictionary =
sanitized.batches[0].data.children[0].dictionary!;
const secondLocalDictionary =
sanitized.batches[1].data.children[0].dictionary!;
expect(secondLocalDictionary.data[0]).toBe(
firstLocalDictionary.data[0],
);
const buf = await fromTableToBuffer(sanitized);
const actual = currentTableFromIPC(buf);
expect([...actual.getChild("label")!]).toEqual(expected);
});
it("can serialize list data from another Arrow version", async function () {
const values = [["anime", "action"], [], null];
const vector = vectorFromArray(
values,
new List(new Field("item", new Utf8(), true)),
);
const table = new ArrowTable({ tags: vector });
const buf = await fromDataToBuffer(table);
const actual = currentTableFromIPC(buf);
const actualTags = actual.getChild("tags");
expect(actualTags?.get(0)?.toJSON()).toEqual(values[0]);
expect(actualTags?.get(1)?.toJSON()).toEqual(values[1]);
expect(actualTags?.get(2)).toBeNull();
});
it("can still import data", async function () {
const schema = new arrow15.Schema([
new arrow15.Field("id", new arrow15.Int32()),
+174 -29
View File
@@ -9,7 +9,7 @@
// comes from the exact same library instance. This is not always the case
// and so we must sanitize the input to ensure that it is compatible.
import { BufferType, Data } from "apache-arrow";
import { BufferType, Data, Vector } from "apache-arrow";
import type { IntBitWidth, TKeys, TimeBitWidth } from "apache-arrow/type";
import {
Binary,
@@ -74,6 +74,20 @@ import {
Utf8,
} from "./arrow";
type SanitizationContext = {
types: WeakMap<object, DataType>;
vectors: WeakMap<object, Vector>;
data: WeakMap<object, Data<DataType>>;
};
function createSanitizationContext(): SanitizationContext {
return {
types: new WeakMap(),
vectors: new WeakMap(),
data: new WeakMap(),
};
}
export function sanitizeMetadata(
metadataLike?: unknown,
): Map<string, string> | undefined {
@@ -186,6 +200,13 @@ export function sanitizeInterval(typeLike: object) {
}
export function sanitizeList(typeLike: object) {
return sanitizeListWithContext(typeLike, createSanitizationContext());
}
function sanitizeListWithContext(
typeLike: object,
context: SanitizationContext,
) {
if (!("children" in typeLike) || !Array.isArray(typeLike.children)) {
throw Error(
"Expected a List type to have an array-like `children` property",
@@ -194,19 +215,35 @@ export function sanitizeList(typeLike: object) {
if (typeLike.children.length !== 1) {
throw Error("Expected a List type to have exactly one child");
}
return new List(sanitizeField(typeLike.children[0]));
return new List(sanitizeFieldWithContext(typeLike.children[0], context));
}
export function sanitizeStruct(typeLike: object) {
return sanitizeStructWithContext(typeLike, createSanitizationContext());
}
function sanitizeStructWithContext(
typeLike: object,
context: SanitizationContext,
) {
if (!("children" in typeLike) || !Array.isArray(typeLike.children)) {
throw Error(
"Expected a Struct type to have an array-like `children` property",
);
}
return new Struct(typeLike.children.map((child) => sanitizeField(child)));
return new Struct(
typeLike.children.map((child) => sanitizeFieldWithContext(child, context)),
);
}
export function sanitizeUnion(typeLike: object) {
return sanitizeUnionWithContext(typeLike, createSanitizationContext());
}
function sanitizeUnionWithContext(
typeLike: object,
context: SanitizationContext,
) {
if (
!("typeIds" in typeLike) ||
!("mode" in typeLike) ||
@@ -226,7 +263,7 @@ export function sanitizeUnion(typeLike: object) {
typeLike.mode,
// biome-ignore lint/suspicious/noExplicitAny: skip
typeLike.typeIds as any,
typeLike.children.map((child) => sanitizeField(child)),
typeLike.children.map((child) => sanitizeFieldWithContext(child, context)),
);
}
@@ -234,6 +271,19 @@ export function sanitizeTypedUnion(
typeLike: object,
// eslint-disable-next-line @typescript-eslint/naming-convention
UnionType: typeof DenseUnion | typeof SparseUnion,
) {
return sanitizeTypedUnionWithContext(
typeLike,
UnionType,
createSanitizationContext(),
);
}
function sanitizeTypedUnionWithContext(
typeLike: object,
// eslint-disable-next-line @typescript-eslint/naming-convention
UnionType: typeof DenseUnion | typeof SparseUnion,
context: SanitizationContext,
) {
if (!("typeIds" in typeLike)) {
throw Error(
@@ -248,7 +298,7 @@ export function sanitizeTypedUnion(
return new UnionType(
typeLike.typeIds as Int32Array | number[],
typeLike.children.map((child) => sanitizeField(child)),
typeLike.children.map((child) => sanitizeFieldWithContext(child, context)),
);
}
@@ -262,6 +312,16 @@ export function sanitizeFixedSizeBinary(typeLike: object) {
}
export function sanitizeFixedSizeList(typeLike: object) {
return sanitizeFixedSizeListWithContext(
typeLike,
createSanitizationContext(),
);
}
function sanitizeFixedSizeListWithContext(
typeLike: object,
context: SanitizationContext,
) {
if (!("listSize" in typeLike) || typeof typeLike.listSize !== "number") {
throw Error("Expected a FixedSizeList type to have a `listSize` property");
}
@@ -275,11 +335,18 @@ export function sanitizeFixedSizeList(typeLike: object) {
}
return new FixedSizeList(
typeLike.listSize,
sanitizeField(typeLike.children[0]),
sanitizeFieldWithContext(typeLike.children[0], context),
);
}
export function sanitizeMap(typeLike: object) {
return sanitizeMapWithContext(typeLike, createSanitizationContext());
}
function sanitizeMapWithContext(
typeLike: object,
context: SanitizationContext,
) {
if (!("children" in typeLike) || !Array.isArray(typeLike.children)) {
throw Error(
"Expected a Map type to have an array-like `children` property",
@@ -292,7 +359,10 @@ export function sanitizeMap(typeLike: object) {
throw Error("Expected a Map type to have exactly one child");
}
return new Map_(sanitizeField(typeLike.children[0]), typeLike.keysSorted);
return new Map_(
sanitizeFieldWithContext(typeLike.children[0], context),
typeLike.keysSorted,
);
}
export function sanitizeDuration(typeLike: object) {
@@ -303,6 +373,13 @@ export function sanitizeDuration(typeLike: object) {
}
export function sanitizeDictionary(typeLike: object) {
return sanitizeDictionaryWithContext(typeLike, createSanitizationContext());
}
function sanitizeDictionaryWithContext(
typeLike: object,
context: SanitizationContext,
) {
if (!("id" in typeLike) || typeof typeLike.id !== "number") {
throw Error("Expected a Dictionary type to have an `id` property");
}
@@ -316,8 +393,8 @@ export function sanitizeDictionary(typeLike: object) {
throw Error("Expected a Dictionary type to have an `isOrdered` property");
}
return new Dictionary(
sanitizeType(typeLike.dictionary),
sanitizeType(typeLike.indices) as TKeys,
sanitizeTypeWithContext(typeLike.dictionary, context),
sanitizeTypeWithContext(typeLike.indices, context) as TKeys,
typeLike.id,
typeLike.isOrdered,
);
@@ -325,12 +402,23 @@ export function sanitizeDictionary(typeLike: object) {
// biome-ignore lint/suspicious/noExplicitAny: skip
export function sanitizeType(typeLike: unknown): DataType<any> {
return sanitizeTypeWithContext(typeLike, createSanitizationContext());
}
function sanitizeTypeWithContext(
typeLike: unknown,
context: SanitizationContext,
): DataType {
if (typeof typeLike === "string") {
return dataTypeFromName(typeLike);
}
if (typeof typeLike !== "object" || typeLike === null) {
throw Error("Expected a Type but object was null/undefined");
}
const cached = context.types.get(typeLike);
if (cached !== undefined) {
return cached;
}
if (
!("typeId" in typeLike) ||
!(
@@ -349,6 +437,16 @@ export function sanitizeType(typeLike: unknown): DataType<any> {
throw Error("Type's typeId property was not a function or number");
}
const type = sanitizeTypeById(typeLike, typeId, context);
context.types.set(typeLike, type);
return type;
}
function sanitizeTypeById(
typeLike: object,
typeId: Type,
context: SanitizationContext,
): DataType {
switch (typeId) {
case Type.NONE:
throw Error("Received a Type with a typeId of NONE");
@@ -375,21 +473,21 @@ export function sanitizeType(typeLike: unknown): DataType<any> {
case Type.Interval:
return sanitizeInterval(typeLike);
case Type.List:
return sanitizeList(typeLike);
return sanitizeListWithContext(typeLike, context);
case Type.Struct:
return sanitizeStruct(typeLike);
return sanitizeStructWithContext(typeLike, context);
case Type.Union:
return sanitizeUnion(typeLike);
return sanitizeUnionWithContext(typeLike, context);
case Type.FixedSizeBinary:
return sanitizeFixedSizeBinary(typeLike);
case Type.FixedSizeList:
return sanitizeFixedSizeList(typeLike);
return sanitizeFixedSizeListWithContext(typeLike, context);
case Type.Map:
return sanitizeMap(typeLike);
return sanitizeMapWithContext(typeLike, context);
case Type.Duration:
return sanitizeDuration(typeLike);
case Type.Dictionary:
return sanitizeDictionary(typeLike);
return sanitizeDictionaryWithContext(typeLike, context);
case Type.Int8:
return new Int8();
case Type.Int16:
@@ -433,9 +531,9 @@ export function sanitizeType(typeLike: unknown): DataType<any> {
case Type.TimestampSecond:
return sanitizeTypedTimestamp(typeLike, TimestampSecond);
case Type.DenseUnion:
return sanitizeTypedUnion(typeLike, DenseUnion);
return sanitizeTypedUnionWithContext(typeLike, DenseUnion, context);
case Type.SparseUnion:
return sanitizeTypedUnion(typeLike, SparseUnion);
return sanitizeTypedUnionWithContext(typeLike, SparseUnion, context);
case Type.IntervalDayTime:
return new IntervalDayTime();
case Type.IntervalYearMonth:
@@ -454,6 +552,13 @@ export function sanitizeType(typeLike: unknown): DataType<any> {
}
export function sanitizeField(fieldLike: unknown): Field {
return sanitizeFieldWithContext(fieldLike, createSanitizationContext());
}
function sanitizeFieldWithContext(
fieldLike: unknown,
context: SanitizationContext,
): Field {
if (fieldLike instanceof Field) {
return fieldLike;
}
@@ -471,7 +576,7 @@ export function sanitizeField(fieldLike: unknown): Field {
}
let type: DataType;
try {
type = sanitizeType(fieldLike.type);
type = sanitizeTypeWithContext(fieldLike.type, context);
} catch (error: unknown) {
throw Error(
`Unable to sanitize type for field: ${fieldLike.name} due to error: ${error}`,
@@ -501,6 +606,13 @@ export function sanitizeField(fieldLike: unknown): Field {
* than lancedb is using.
*/
export function sanitizeSchema(schemaLike: SchemaLike): Schema {
return sanitizeSchemaWithContext(schemaLike, createSanitizationContext());
}
function sanitizeSchemaWithContext(
schemaLike: SchemaLike,
context: SanitizationContext,
): Schema {
if (schemaLike instanceof Schema) {
return schemaLike;
}
@@ -522,7 +634,7 @@ export function sanitizeSchema(schemaLike: SchemaLike): Schema {
);
}
const sanitizedFields = schemaLike.fields.map((field) =>
sanitizeField(field),
sanitizeFieldWithContext(field, context),
);
return new Schema(sanitizedFields, metadata);
}
@@ -544,13 +656,18 @@ export function sanitizeTable(tableLike: TableLike): Table {
"The table passed in does not appear to be a table (no 'columns' property)",
);
}
const schema = sanitizeSchema(tableLike.schema);
const batches = tableLike.batches.map(sanitizeRecordBatch);
const context = createSanitizationContext();
const schema = sanitizeSchemaWithContext(tableLike.schema, context);
const batches = tableLike.batches.map((batch) =>
sanitizeRecordBatch(batch, context),
);
return new Table(schema, batches);
}
function sanitizeRecordBatch(batchLike: RecordBatchLike): RecordBatch {
function sanitizeRecordBatch(
batchLike: RecordBatchLike,
context: SanitizationContext,
): RecordBatch {
if (batchLike instanceof RecordBatch) {
return batchLike;
}
@@ -567,19 +684,43 @@ function sanitizeRecordBatch(batchLike: RecordBatchLike): RecordBatch {
"The record batch passed in does not appear to be a record batch (no 'data' property)",
);
}
const schema = sanitizeSchema(batchLike.schema);
const data = sanitizeData(batchLike.data);
const schema = sanitizeSchemaWithContext(batchLike.schema, context);
const data = sanitizeData(batchLike.data, context) as Data<Struct>;
return new RecordBatch(schema, data);
}
type DictionaryVectorLike = {
data: readonly DataLike[];
};
type DictionaryDataLike = DataLike & {
dictionary?: DictionaryVectorLike;
};
function sanitizeData(
dataLike: DataLike,
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
): import("apache-arrow").Data<Struct<any>> {
context: SanitizationContext,
): Data<DataType> {
if (dataLike instanceof Data) {
return dataLike;
}
return new Data(
dataLike.type,
const cachedData = context.data.get(dataLike);
if (cachedData !== undefined) {
return cachedData;
}
const dictionaryLike = (dataLike as DictionaryDataLike).dictionary;
let dictionary: Vector | undefined;
if (dictionaryLike !== undefined) {
dictionary = context.vectors.get(dictionaryLike);
if (dictionary === undefined) {
dictionary = new Vector(
dictionaryLike.data.map((data) => sanitizeData(data, context)),
);
context.vectors.set(dictionaryLike, dictionary);
}
}
const data = new Data(
sanitizeTypeWithContext(dataLike.type, context),
dataLike.offset,
dataLike.length,
dataLike.nullCount,
@@ -589,7 +730,11 @@ function sanitizeData(
[BufferType.VALIDITY]: dataLike.nullBitmap,
[BufferType.TYPE]: dataLike.typeIds,
},
dataLike.children.map((child) => sanitizeData(child, context)),
dictionary,
);
context.data.set(dataLike, data);
return data;
}
const constructorsByTypeName = {
+11 -1
View File
@@ -42,9 +42,19 @@ impl Job {
}
/// Wait until the operation reaches a terminal state.
///
/// Jobs that complete without a resource result resolve successfully.
/// Resource results are not exposed on this binding yet; unsupported
/// success results reject with a generic error.
#[napi(catch_unwind)]
pub async fn wait(&self) -> napi::Result<()> {
self.inner.wait().await.default_error()
match self.inner.wait().await.default_error()? {
lancedb::JobResult::None => Ok(()),
// JobResult is non_exhaustive; Function and future variants fail closed.
_ => Err(napi::Error::from_reason(
"unsupported job result".to_string(),
)),
}
}
/// Request cancellation. Cancelling a finished operation is a no-op.
+5
View File
@@ -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
@@ -23,6 +24,7 @@ from .schema import blob, vector, BlobType
from .job import AsyncJob, Job
from .table import AsyncTable, Table
from .types import BaseTokenizerType
from ._udf import FunctionCapability, udf
from ._lancedb import Session
from .namespace import (
connect_namespace,
@@ -507,6 +509,8 @@ __all__ = [
"FtsToken",
"col",
"Expr",
"Function",
"FunctionCapability",
"func",
"lit",
"URI",
@@ -521,5 +525,6 @@ __all__ = [
"RemoteDBConnection",
"Session",
"Table",
"udf",
"__version__",
]
+108
View File
@@ -0,0 +1,108 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Private first-class Function namespace facades for database connections.
These helpers are internal submission and lookup surfaces. They are not durable
resources and are not part of the public top-level export surface.
"""
from __future__ import annotations
from collections.abc import Callable
from typing import TYPE_CHECKING
from . import _udf
from ._lancedb import Function
from .job import AsyncJob, Job
if TYPE_CHECKING:
from .db import AsyncConnection, DBConnection
class _SyncFunctions:
"""Synchronous `db.functions` facade."""
__slots__ = ("_connection",)
def __init__(self, connection: DBConnection) -> None:
self._connection = connection
def __repr__(self) -> str:
return "_SyncFunctions()"
def register(self, name: str, decorated_udf: Callable[..., object]) -> Job:
"""Register a decorated UDF and return a synchronous [Job][lancedb.job.Job]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = self._connection._submit_register_function(name, definition)
return Job(AsyncJob(native_job))
def replace(
self, name: str, current: Function, decorated_udf: Callable[..., object]
) -> Job:
"""Conditionally replace a Function; return sync [Job][lancedb.job.Job]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = self._connection._submit_replace_function(
name, current, definition
)
return Job(AsyncJob(native_job))
def get(self, name: str) -> Function:
"""Return the Function currently bound to a database-scoped name."""
return self._connection._lookup_function_by_name(name)
def get_by_id(self, function_id: str) -> Function:
"""Return the immutable Function for an exact Function ID."""
return self._connection._lookup_function_by_id(function_id)
def remove(self, name: str, current: Function) -> None:
"""Conditionally remove a Function catalog name binding."""
return self._connection._remove_function_name(name, current)
def revoke(self, function: Function) -> None:
"""Revoke an exact immutable Function by administrator set-bit."""
return self._connection._revoke_function(function)
class _AsyncFunctions:
"""Asynchronous `async_db.functions` facade."""
__slots__ = ("_connection",)
def __init__(self, connection: AsyncConnection) -> None:
self._connection = connection
def __repr__(self) -> str:
return "_AsyncFunctions()"
async def register(
self, name: str, decorated_udf: Callable[..., object]
) -> AsyncJob:
"""Register a decorated UDF and return an [AsyncJob][lancedb.job.AsyncJob]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = await self._connection._register_function(name, definition)
return AsyncJob(native_job)
async def replace(
self, name: str, current: Function, decorated_udf: Callable[..., object]
) -> AsyncJob:
"""Conditionally replace a Function; return [AsyncJob][lancedb.job.AsyncJob]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = await self._connection._replace_function(name, current, definition)
return AsyncJob(native_job)
async def get(self, name: str) -> Function:
"""Return the Function currently bound to a database-scoped name."""
return await self._connection._lookup_function_by_name(name)
async def get_by_id(self, function_id: str) -> Function:
"""Return the immutable Function for an exact Function ID."""
return await self._connection._lookup_function_by_id(function_id)
async def remove(self, name: str, current: Function) -> None:
"""Conditionally remove a Function catalog name binding."""
return await self._connection._remove_function_name(name, current)
async def revoke(self, function: Function) -> None:
"""Revoke an exact immutable Function by administrator set-bit."""
return await self._connection._revoke_function(function)
+59 -1
View File
@@ -153,6 +153,16 @@ class Connection(object):
async def job_history(
self, job_id: Optional[str] = None
) -> List[pa.RecordBatch]: ...
async def _register_function(
self, name: str, definition: "_FunctionDefinition"
) -> Job: ...
async def _replace_function(
self, name: str, current: Function, definition: "_FunctionDefinition"
) -> Job: ...
async def _lookup_function_by_name(self, name: str) -> Function: ...
async def _lookup_function_by_id(self, function_id: str) -> Function: ...
async def _remove_function_name(self, name: str, current: Function) -> None: ...
async def _revoke_function(self, function: Function) -> None: ...
async def create_table(
self,
name: str,
@@ -216,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:
@@ -242,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
@@ -256,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: ...
@@ -318,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]): ...
+538
View File
@@ -0,0 +1,538 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Local authoring declaration surface for first-class UDFs.
This module snapshots declaration metadata onto a Python function and privately
validates packagable callables into a source snapshot. It does not mint durable
identity or register anything with a database.
"""
from __future__ import annotations
import ast
import inspect
import stat
import symtable
import sys
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from types import CodeType, FunctionType
from typing import NoReturn, ParamSpec, TypeVar
import pyarrow as pa
from . import _lancedb
__all__ = ["FunctionCapability", "udf"]
_P = ParamSpec("_P")
_R = TypeVar("_R")
_CONFIG_ATTR = "__lancedb_udf_config__"
_SYNTHETIC_SOURCE_FILENAME = "<lancedb-udf>"
_PACKAGING_ERROR = "udf is not packagable"
_ALLOWED_PARAM_KINDS = (
inspect.Parameter.POSITIONAL_OR_KEYWORD,
inspect.Parameter.KEYWORD_ONLY,
)
class FunctionCapability:
"""Local capability declaration for a first-class UDF.
Construct via :meth:`network` or :meth:`secret`. Direct construction is
rejected so callers cannot create an uninitialized capability.
"""
__slots__ = ("_kind", "_origin", "_reference", "_environment_variable")
def __new__(cls, *args: object, **kwargs: object) -> FunctionCapability:
raise TypeError(
"FunctionCapability cannot be constructed directly; "
"use FunctionCapability.network() or FunctionCapability.secret()"
)
@classmethod
def _create(
cls,
kind: str,
origin: str | None,
reference: str | None,
environment_variable: str | None,
) -> FunctionCapability:
obj = object.__new__(cls)
object.__setattr__(obj, "_kind", kind)
object.__setattr__(obj, "_origin", origin)
object.__setattr__(obj, "_reference", reference)
object.__setattr__(obj, "_environment_variable", environment_variable)
return obj
@classmethod
def network(cls, origin: str) -> FunctionCapability:
if not isinstance(origin, str):
raise TypeError("origin must be a string")
if origin == "":
raise ValueError("origin must be non-empty")
return cls._create("network", origin, None, None)
@classmethod
def secret(cls, reference: str, *, environment_variable: str) -> FunctionCapability:
if not isinstance(reference, str):
raise TypeError("reference must be a string")
if not isinstance(environment_variable, str):
raise TypeError("environment_variable must be a string")
if reference == "":
raise ValueError("reference must be non-empty")
if environment_variable == "":
raise ValueError("environment_variable must be non-empty")
return cls._create("secret", None, reference, environment_variable)
@property
def kind(self) -> str:
return self._kind
@property
def origin(self) -> str | None:
return self._origin
@property
def reference(self) -> str | None:
return self._reference
@property
def environment_variable(self) -> str | None:
return self._environment_variable
def __setattr__(self, name: str, value: object) -> None:
raise AttributeError(
f"{type(self).__name__!r} object attribute {name!r} is read-only"
)
def __delattr__(self, name: str) -> None:
raise AttributeError(
f"{type(self).__name__!r} object attribute {name!r} is read-only"
)
def __eq__(self, other: object) -> bool:
if not isinstance(other, FunctionCapability):
return NotImplemented
return (
self._kind == other._kind
and self._origin == other._origin
and self._reference == other._reference
and self._environment_variable == other._environment_variable
)
def __hash__(self) -> int:
return hash(
(
self._kind,
self._origin,
self._reference,
self._environment_variable,
)
)
def __repr__(self) -> str:
if self._kind == "network":
return f"FunctionCapability.network({self._origin!r})"
return (
"FunctionCapability.secret("
f"environment_variable={self._environment_variable!r})"
)
@dataclass(frozen=True, slots=True)
class _UdfConfig:
"""Private frozen snapshot of a ``@udf`` declaration."""
inputs: tuple[tuple[str, pa.DataType], ...]
output: pa.DataType
output_nullable: bool
python: str
packages: tuple[str, ...]
capabilities: tuple[FunctionCapability, ...]
@dataclass(frozen=True, slots=True)
class _PackagedUdf:
"""Private frozen snapshot of a validated packagable UDF."""
source: str
module: str
callable_name: str
config: _UdfConfig
def __repr__(self) -> str:
return (
f"_PackagedUdf(source=<redacted>, module={self.module!r}, "
f"callable_name={self.callable_name!r}, config={self.config!r})"
)
def _validate_inputs(
inputs: object,
) -> tuple[tuple[str, pa.DataType], ...]:
if not isinstance(inputs, Mapping):
raise TypeError("udf inputs must be a Mapping of name to pyarrow DataType")
snapshot: list[tuple[str, pa.DataType]] = []
for key, value in inputs.items():
if not isinstance(key, str):
raise TypeError("udf input names must be strings")
if key == "":
raise ValueError("udf input names must be non-empty")
if not isinstance(value, pa.DataType):
raise TypeError("udf input types must be pyarrow DataType values")
snapshot.append((key, value))
return tuple(snapshot)
def _validate_packages(packages: object) -> tuple[str, ...]:
if isinstance(packages, (str, bytes, bytearray)):
raise TypeError("udf packages must be a sequence of strings, not a string")
if not isinstance(packages, Sequence):
raise TypeError("udf packages must be a sequence of strings")
snapshot: list[str] = []
seen: set[str] = set()
for package in packages:
if not isinstance(package, str):
raise TypeError("udf packages must contain only strings")
if package == "":
raise ValueError("udf packages must be non-empty strings")
if package in seen:
raise ValueError(f"duplicate udf package: {package}")
seen.add(package)
snapshot.append(package)
return tuple(snapshot)
def _reject_non_exact_capability() -> NoReturn:
# Exact-type only: subclasses are authoring inputs we never accept. Keep the
# message fixed so hostile markers never enter exception text.
raise TypeError(
"udf capabilities must contain only FunctionCapability values"
) from None
def _require_exact_capability(capability: object) -> FunctionCapability:
if type(capability) is not FunctionCapability:
_reject_non_exact_capability()
return capability
def _validate_capabilities(capabilities: object) -> tuple[FunctionCapability, ...]:
if isinstance(capabilities, (str, bytes, bytearray)):
raise TypeError(
"udf capabilities must be a sequence of FunctionCapability, not a string"
)
if not isinstance(capabilities, Sequence):
raise TypeError("udf capabilities must be a sequence of FunctionCapability")
return tuple(_require_exact_capability(capability) for capability in capabilities)
def udf(
*,
inputs: Mapping[str, pa.DataType],
output: pa.DataType,
python: str,
packages: Sequence[str] = (),
output_nullable: bool = True,
capabilities: Sequence[FunctionCapability] = (),
) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]:
"""Declare a local UDF without packaging or registration.
Applying the returned decorator attaches a private frozen config snapshot
and returns the exact same function object.
"""
input_snapshot = _validate_inputs(inputs)
if not isinstance(output, pa.DataType):
raise TypeError("udf output must be a pyarrow DataType")
if not isinstance(python, str):
raise TypeError("udf python must be a string")
if python == "":
raise ValueError("udf python must be a non-empty string")
package_snapshot = _validate_packages(packages)
if not isinstance(output_nullable, bool):
raise TypeError("udf output_nullable must be a bool")
capability_snapshot = _validate_capabilities(capabilities)
config = _UdfConfig(
inputs=input_snapshot,
output=output,
output_nullable=output_nullable,
python=python,
packages=package_snapshot,
capabilities=capability_snapshot,
)
def decorator(fn: Callable[_P, _R]) -> Callable[_P, _R]:
if not inspect.isfunction(fn):
raise TypeError("udf can only decorate a Python function")
if hasattr(fn, _CONFIG_ATTR):
raise ValueError("function is already decorated with @udf")
setattr(fn, _CONFIG_ATTR, config)
return fn
return decorator
def _get_udf_config(fn: object) -> _UdfConfig:
"""Return the private declaration snapshot for a ``@udf``-decorated function."""
config = getattr(fn, _CONFIG_ATTR, None)
if config is None:
raise TypeError("function is not decorated with @udf")
if not isinstance(config, _UdfConfig):
raise TypeError("function is not decorated with @udf")
return config
def _packaging_reject() -> NoReturn:
raise ValueError(_PACKAGING_ERROR) from None
def _is_ordinary_function(fn: FunctionType) -> bool:
if fn.__name__ == "<lambda>":
return False
if fn.__qualname__ != fn.__name__:
return False
if inspect.iscoroutinefunction(fn) or inspect.isasyncgenfunction(fn):
return False
if inspect.isgeneratorfunction(fn):
return False
return True
def _resolve_source_path(fn: FunctionType, module: object) -> Path:
try:
fn_source: str | None = inspect.getsourcefile(fn)
except TypeError:
fn_source = None
source_lookup_failed = True
else:
source_lookup_failed = False
if source_lookup_failed:
_packaging_reject()
module_file = vars(module).get("__file__")
if not fn_source or not isinstance(module_file, str) or module_file == "":
_packaging_reject()
try:
resolved_paths: tuple[Path, Path] | None = (
Path(fn_source).resolve(),
Path(module_file).resolve(),
)
except (OSError, RuntimeError):
resolved_paths = None
if resolved_paths is None:
_packaging_reject()
fn_path, module_path = resolved_paths
if fn_path != module_path:
_packaging_reject()
if fn_path.suffix != ".py":
_packaging_reject()
try:
mode: int | None = fn_path.stat().st_mode
except OSError:
mode = None
if mode is None:
_packaging_reject()
if not stat.S_ISREG(mode):
_packaging_reject()
return fn_path
def _validate_source(
source: str, callable_name: str
) -> tuple[CodeType, symtable.SymbolTable]:
try:
module_code = compile(
source,
_SYNTHETIC_SOURCE_FILENAME,
"exec",
optimize=sys.flags.optimize,
)
ast.parse(source, filename=_SYNTHETIC_SOURCE_FILENAME, mode="exec")
table = symtable.symtable(source, _SYNTHETIC_SOURCE_FILENAME, "exec")
parsed: tuple[CodeType, symtable.SymbolTable] | None = (module_code, table)
except Exception:
parsed = None
if parsed is None:
_packaging_reject()
module_code, table = parsed
for child in table.get_children():
if child.get_name() == callable_name and child.get_type() == "function":
return module_code, table
_packaging_reject()
def _source_bound_names(table: symtable.SymbolTable) -> set[str]:
names: set[str] = set()
for symbol in table.get_symbols():
if symbol.is_imported() or symbol.is_assigned() or symbol.is_namespace():
names.add(symbol.get_name())
return names
def _code_fingerprint(code: CodeType) -> tuple[object, ...]:
"""Structural fingerprint ignoring only location/debug fields."""
constants = tuple(
_code_fingerprint(constant) if isinstance(constant, CodeType) else constant
for constant in code.co_consts
)
return (
code.co_name,
getattr(code, "co_qualname", code.co_name),
code.co_argcount,
code.co_posonlyargcount,
code.co_kwonlyargcount,
code.co_flags,
code.co_code,
code.co_names,
code.co_varnames,
code.co_freevars,
code.co_cellvars,
getattr(code, "co_exceptiontable", b""),
constants,
)
def _toplevel_code_candidates(
module_code: CodeType, callable_name: str
) -> list[CodeType]:
candidates: list[CodeType] = []
for constant in module_code.co_consts:
if not isinstance(constant, CodeType):
continue
if constant.co_name != callable_name:
continue
if getattr(constant, "co_qualname", callable_name) != callable_name:
continue
candidates.append(constant)
return candidates
def _validate_loaded_code_matches_source(
fn: FunctionType, module_code: CodeType
) -> None:
candidates = _toplevel_code_candidates(module_code, fn.__name__)
if not candidates:
_packaging_reject()
target = _code_fingerprint(fn.__code__)
if not any(_code_fingerprint(candidate) == target for candidate in candidates):
_packaging_reject()
def _validate_signature(fn: FunctionType, config: _UdfConfig) -> None:
try:
signature: inspect.Signature | None = inspect.signature(fn)
except (TypeError, ValueError):
signature = None
if signature is None:
_packaging_reject()
parameters = list(signature.parameters.values())
expected = [name for name, _ in config.inputs]
actual = [parameter.name for parameter in parameters]
if actual != expected:
_packaging_reject()
for parameter in parameters:
if parameter.kind not in _ALLOWED_PARAM_KINDS:
_packaging_reject()
def _validate_ambient_globals(fn: FunctionType, table: symtable.SymbolTable) -> None:
try:
closure_vars: inspect.ClosureVars | None = inspect.getclosurevars(fn)
except (TypeError, ValueError):
closure_vars = None
if closure_vars is None:
_packaging_reject()
if closure_vars.nonlocals:
_packaging_reject()
bound_names = _source_bound_names(table)
for name in closure_vars.globals:
if name not in bound_names:
_packaging_reject()
def _package_udf(fn: object) -> _PackagedUdf:
"""Validate and snapshot a packagable ``@udf``-decorated function."""
config = _get_udf_config(fn)
if not isinstance(fn, FunctionType) or not _is_ordinary_function(fn):
_packaging_reject()
module_name = fn.__module__
if (
not isinstance(module_name, str)
or module_name == ""
or module_name == "__main__"
):
_packaging_reject()
module = sys.modules.get(module_name)
if module is None:
_packaging_reject()
callable_name = fn.__name__
if vars(module).get(callable_name) is not fn:
_packaging_reject()
source_path = _resolve_source_path(fn, module)
try:
source: str | None = source_path.read_text(encoding="utf-8")
except (OSError, UnicodeError):
source = None
if source is None:
_packaging_reject()
module_code, table = _validate_source(source, callable_name)
_validate_signature(fn, config)
_validate_ambient_globals(fn, table)
_validate_loaded_code_matches_source(fn, module_code)
return _PackagedUdf(
source=source,
module=module_name,
callable_name=callable_name,
config=config,
)
def _normalize_capability_triple(
capability: FunctionCapability,
) -> tuple[str, str, str | None]:
"""Normalize a local capability declaration to the native triple shape."""
# Private config is untrusted; re-check exact type before any property access.
capability = _require_exact_capability(capability)
if capability.kind == "network":
origin = capability.origin
if origin is None:
raise ValueError("invalid network capability") from None
return ("network", origin, None)
if capability.kind == "secret":
reference = capability.reference
environment_variable = capability.environment_variable
if reference is None or environment_variable is None:
raise ValueError("invalid secret capability") from None
return ("secret", reference, environment_variable)
# Fail closed without echoing the unknown kind.
raise ValueError("unsupported capability kind") from None
def _build_function_definition(fn: object) -> _lancedb._FunctionDefinition:
"""Package a ``@udf`` and bridge it to the private native definition."""
packaged = _package_udf(fn)
config = packaged.config
capabilities = [
_normalize_capability_triple(capability) for capability in config.capabilities
]
return _lancedb._new_function_definition(
parameters=list(config.inputs),
output_type=config.output,
output_nullable=config.output_nullable,
module=packaged.module,
callable_name=packaged.callable_name,
source=packaged.source,
python=config.python,
packages=list(config.packages),
capabilities=capabilities,
)
+126
View File
@@ -63,8 +63,12 @@ if TYPE_CHECKING:
import pyarrow as pa
from .pydantic import LanceModel
from ._functions import _AsyncFunctions, _SyncFunctions
from ._lancedb import Connection as LanceDbConnection
from ._lancedb import Function
from ._lancedb import Job as NativeJob
from ._lancedb import JobDescription, JobInfo
from ._lancedb import _FunctionDefinition
from .common import DATA, URI
from .embeddings import EmbeddingFunctionConfig
from ._lancedb import Session
@@ -650,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):
"""
@@ -1267,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.
@@ -2013,6 +2110,35 @@ class AsyncConnection(object):
"""
return await self._inner.job_history(job_id)
@property
def functions(self) -> "_AsyncFunctions":
"""First-class Function operations for this connection."""
from ._functions import _AsyncFunctions
return _AsyncFunctions(self)
async def _register_function(
self, name: str, definition: "_FunctionDefinition"
) -> NativeJob:
return await self._inner._register_function(name, definition)
async def _replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> NativeJob:
return await self._inner._replace_function(name, current, definition)
async def _lookup_function_by_name(self, name: str) -> "Function":
return await self._inner._lookup_function_by_name(name)
async def _lookup_function_by_id(self, function_id: str) -> "Function":
return await self._inner._lookup_function_by_id(function_id)
async def _remove_function_name(self, name: str, current: "Function") -> None:
return await self._inner._remove_function_name(name, current)
async def _revoke_function(self, function: "Function") -> None:
return await self._inner._revoke_function(function)
async def namespace_client(self) -> LanceNamespace:
"""Get the equivalent namespace client for this connection.
+39 -2
View File
@@ -3,6 +3,8 @@
"""Custom exception handling"""
from typing import Optional
class MissingValueError(ValueError):
"""Exception raised when a required value is missing."""
@@ -26,12 +28,47 @@ class MissingColumnError(KeyError):
class JobFailedError(RuntimeError):
"""Exception raised when an asynchronous job reaches the failed state."""
"""Exception raised when an asynchronous job reaches the failed state.
pass
``error_code`` is the optional exact category string projected from the
native job failure when the backend supplied one. The RuntimeError
message remains the existing diagnostic text and must not be used to
recover or override the code.
"""
__slots__ = ("_error_code",)
def __init__(self, message: str, error_code: Optional[str] = None) -> None:
super().__init__(message)
self._error_code = error_code
@property
def error_code(self) -> Optional[str]:
"""Exact job failure error category string, when supplied."""
return self._error_code
class JobCancelledError(RuntimeError):
"""Exception raised when an asynchronous job was cancelled."""
pass
class FunctionError(RuntimeError):
"""Exception raised when a first-class Function operation fails.
``code`` is the stable semantic category from the native error. The
message is a sanitized client diagnostic and must not be used to recover
or override the code.
"""
__slots__ = ("_code",)
def __init__(self, message: str, code: str) -> None:
super().__init__(message)
self._code = code
@property
def code(self) -> str:
"""Stable Function error category string."""
return self._code
+16 -7
View File
@@ -10,6 +10,7 @@ from typing import Optional
from lancedb.background_loop import LOOP
from . import _lancedb
from ._lancedb import Function
class AsyncJob:
@@ -44,18 +45,22 @@ class AsyncJob:
return "finished"
return await self._inner.status()
async def wait(self, timeout: Optional[timedelta] = None):
async def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
"""Wait until the operation reaches a terminal state.
Returns the success result when present (currently a
:class:`~lancedb.Function`), or `None` when the job finished without
one.
Raises `JobFailedError` if the operation failed, `JobCancelledError`
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
"""
if self._inner is None:
return
return None
if timeout is None:
await self._inner.wait()
return await self._inner.wait()
else:
await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
return await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
async def cancel(self):
"""Request cancellation. Cancelling a finished operation is a no-op."""
@@ -88,15 +93,19 @@ class Job:
return "finished"
return LOOP.run(self._inner.status())
def wait(self, timeout: Optional[timedelta] = None):
def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
"""Block until the operation reaches a terminal state.
Returns the success result when present (currently a
:class:`~lancedb.Function`), or `None` when the job finished without
one.
Raises `JobFailedError` if the operation failed, `JobCancelledError`
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
"""
if self._inner is None:
return
LOOP.run(self._inner.wait(timeout))
return None
return LOOP.run(self._inner.wait(timeout))
def cancel(self):
"""Request cancellation. Cancelling a finished operation is a no-op."""
+31 -1
View File
@@ -26,7 +26,9 @@ from ..db import DBConnection, LOOP
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,
@@ -734,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.
+43
View File
@@ -7,6 +7,7 @@ import logging
from functools import cached_property
import os
from typing import (
TYPE_CHECKING,
Any,
Callable,
Dict,
@@ -67,6 +68,9 @@ from ..query import (
from ..table import AsyncTable, BlobMode, Branches, IndexStatistics, Query, Table, Tags
from ..types import BaseTokenizerType
if TYPE_CHECKING:
from lancedb._lancedb import _FunctionCall
class RemoteTable(Table):
def __init__(
@@ -570,6 +574,45 @@ class RemoteTable(Table):
)
)
def add_generated_column(self, column_name: str, call: "_FunctionCall") -> Job:
"""Add a generated column from an authored Function call.
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
of the Job does not publish the column; callers must wait and re-read
the table to observe the new definition and values.
"""
return Job(LOOP.run(self._table.add_generated_column(column_name, call)))
def generated_column_status(
self, column_name: str
) -> Literal["complete", "incomplete"]:
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
Projection-only: reads the named column's stored definition status.
Does not refresh values, submit a Job, or mutate table state.
"""
return LOOP.run(self._table.generated_column_status(column_name))
def refresh_generated_column(self, column_name: str) -> Job:
"""Refresh values for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
of the Job does not publish new values; callers must wait and re-read
the table to observe refreshed results.
"""
return Job(LOOP.run(self._table.refresh_generated_column(column_name)))
def alter_generated_column(
self, column_name: str, new_call: "_FunctionCall"
) -> Job:
"""Alter the Function call for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
of the Job does not publish the new definition; callers must wait and
re-read the table to observe the updated column.
"""
return Job(LOOP.run(self._table.alter_generated_column(column_name, new_call)))
def _is_legacy_create_index_call(
self,
first_arg: str,
+315 -27
View File
@@ -11,6 +11,11 @@ Provides StreamingDataset, a PyTorch IterableDataset that guarantees:
- **Resumability**: state_dict / load_state_dict capture per-split consumption
counts so training can resume from an exact mid-epoch position even when the
distributed topology changes between runs.
Transform failures on bad rows (e.g. nulls or NaNs from incomplete data) can
be tolerated with ``on_transform_error="skip"``; see the parameter
documentation on StreamingDataset for how this interacts with the guarantees
above.
"""
import ctypes
@@ -22,7 +27,7 @@ import time
from collections import deque
from concurrent.futures import ThreadPoolExecutor
from multiprocessing import RawArray
from typing import Any, Callable, Iterator, Optional
from typing import Any, Callable, Iterator, Optional, Union
from torch.utils.data import IterableDataset, get_worker_info
@@ -127,6 +132,49 @@ class StreamingDataset(IterableDataset):
Maximum number of transforms to run concurrently. Must be greater
than zero. When ``None`` (the default), uses ``os.cpu_count()`` or 1
when the CPU count is unavailable.
on_transform_error:
What to do when the transform raises an exception:
- ``"raise"`` (the default): the exception propagates and iteration
aborts.
- ``"skip"``: the failing rows are dropped and iteration continues.
- ``"warn"``: like ``"skip"``, but a warning is logged for each
failing batch.
- a callable ``handler(exc) -> bool``: called with the exception;
return ``True`` to skip the failing rows or ``False`` to re-raise.
Useful to skip only expected error types (compatible with
``webdataset.handlers`` style handlers).
When a batch fails, the transform is re-invoked on each single-row
slice of the batch so that only the rows that actually fail are
dropped. Transforms should therefore be deterministic and accept
batches of any size (including one row). Skipped rows are counted in
``rows_skipped``.
Skipping weakens the elastic-determinism guarantee at the end of the
epoch: splits that lose more rows than others run dry earlier, and
each rank's iterator ends at the last cycle where every split *it
owns* still has a row. Because bad rows are not distributed evenly
across splits, this means one rank's iterator can yield noticeably
fewer or more steps than another rank's *in the same run* — there is
no cross-rank coordination that stops every rank at the same global
step. This is generally safe for asynchronous or single-rank use,
but synchronous distributed training (e.g. ranks that call
``all_reduce`` every step) can hang or deadlock if one rank's
iterator is exhausted while others are still stepping; callers doing
synchronous multi-rank training with ``on_transform_error != "raise"``
are responsible for their own cross-rank stopping mechanism (e.g.
broadcasting a stop signal on ``StopIteration``). The final few
global steps can also differ across topologies (bounded by the skew
in bad-row counts across splits). The sequence of samples yielded
from each split remains deterministic. Mid-epoch
checkpoints remain exact provided the transform fails
deterministically; in multi-rank training each rank must save its
own ``state_dict`` and the states must be combined with
``merge_state_dicts`` before resuming on a different topology.
Prefer the ``filter`` parameter when bad rows can be expressed as a
SQL predicate (e.g. ``"col IS NOT NULL"``) — filtering happens before
splits are built, so every guarantee is fully preserved.
worker_info_override:
If set, used in place of ``torch.utils.data.get_worker_info()`` to
determine the DataLoader worker assignment. Intended for unit tests
@@ -152,6 +200,7 @@ class StreamingDataset(IterableDataset):
filter: Optional[str] = None,
transform: Optional[Callable] = None,
transform_parallelism: Optional[int] = None,
on_transform_error: Union[str, Callable[[Exception], bool]] = "raise",
connection_factory: Optional[Callable[[str], Any]] = None,
worker_info_override=None,
):
@@ -167,6 +216,13 @@ class StreamingDataset(IterableDataset):
)
if transform_parallelism is not None and transform_parallelism <= 0:
raise ValueError("transform_parallelism must be greater than 0")
if on_transform_error not in ("raise", "skip", "warn") and not callable(
on_transform_error
):
raise ValueError(
"on_transform_error must be 'raise', 'skip', 'warn', or a "
f"callable, got {on_transform_error!r}"
)
self._table = table
self._num_splits = num_splits
@@ -182,6 +238,7 @@ class StreamingDataset(IterableDataset):
self._filter = filter
self._transform = transform
self._transform_parallelism = transform_parallelism
self._on_transform_error = on_transform_error
self._connection_factory = connection_factory
self._worker_info_override = worker_info_override
@@ -199,19 +256,28 @@ class StreamingDataset(IterableDataset):
# in the main process. RawArray is picklable via the forkserver
# reduction protocol so it survives the dataset pickle round-trip.
# Layout: [unscanned_rows, raw_rows, cooked_rows, consumed_rows,
# bytes_loaded, fetch_time_us, transform_time_us]
self._worker_stats: RawArray = RawArray(ctypes.c_int64, 7)
# bytes_loaded, fetch_time_us, transform_time_us,
# rows_skipped]
self._worker_stats: RawArray = RawArray(ctypes.c_int64, 8)
# Cumulative bytes of Arrow buffer data fetched across all iterations.
self._bytes_loaded: int = 0
# Cumulative seconds spent in LanceDB I/O and in transform functions.
self._fetch_time: float = 0.0
self._transform_time: float = 0.0
# Cumulative rows dropped by on_transform_error across all iterations.
self._rows_skipped: int = 0
# Number of samples each split has already been consumed. At global
# step boundaries all splits have consumed this many samples, so a
# single scalar captures the topology-independent checkpoint state.
self._resume_offset: int = 0
# Permutation position each split has consumed through, keyed by
# global split index. Equal to _resume_offset for every split unless
# on_transform_error skipped rows, in which case skipped positions
# push the watermark of the affected splits further ahead. Splits
# this instance has never iterated have no entry.
self._resume_positions: dict[int, int] = {}
# Build the permutation table once, deterministically.
builder = permutation_builder(table)
@@ -275,6 +341,7 @@ class StreamingDataset(IterableDataset):
# Set identity transform on each Permutation so __getitems__ returns
# the raw RecordBatch. Stage 2 applies the real transform.
permutations: list[Permutation] = []
initial_positions: list[int] = []
for split_idx in my_splits:
perm = Permutation.from_tables(
self._table, self._perm_table, split=split_idx
@@ -282,14 +349,20 @@ class StreamingDataset(IterableDataset):
if self._columns is not None:
perm = perm.select_columns(self._columns)
perm = perm.with_transform(lambda batch: batch)
if self._resume_offset > 0:
perm = perm.with_skip(self._resume_offset)
start_pos = self._resume_positions.get(split_idx, self._resume_offset)
if start_pos > 0:
perm = perm.with_skip(start_pos)
initial_positions.append(start_pos)
permutations.append(perm)
n = len(permutations)
split_sizes = [perm.num_rows for perm in permutations]
initial_offset = self._resume_offset
local_consumed = [0] * n
# Permutation position each split has consumed through (absolute,
# i.e. counted from the start of the unskipped split). Runs ahead of
# initial + local_consumed when rows are skipped.
pos_consumed = list(initial_positions)
batch_size = self._read_batch_size
max_prefetch = self._prefetch_batches
@@ -302,12 +375,14 @@ class StreamingDataset(IterableDataset):
self._transform if self._transform is not None else Transforms.arrow2python
)
# Per-split pipeline state.
# Per-split pipeline state. Batches are paired with the absolute
# permutation position of their first row so that skipped rows can be
# accounted for in pos_consumed.
fetch_head = [0] * n
io_pending = [deque() for _ in range(n)] # Future[RecordBatch]
raw_batches = [deque() for _ in range(n)] # RecordBatch — fetched, awaiting tx
tx_pending = [deque() for _ in range(n)] # Future[list[Any]]
cooked = [deque() for _ in range(n)] # rows ready to yield
io_pending = [deque() for _ in range(n)] # (abs_start, Future[RecordBatch])
raw_batches = [deque() for _ in range(n)] # (abs_start, RecordBatch)
tx_pending = [deque() for _ in range(n)] # Future[list[(abs_pos, row)]]
cooked = [deque() for _ in range(n)] # (abs_pos, row) ready to yield
# Limit simultaneous transforms to transform_workers across all splits.
tx_semaphore = threading.Semaphore(transform_workers)
@@ -330,7 +405,8 @@ class StreamingDataset(IterableDataset):
fetch_head[i] += fetch
perm_i = permutations[i]
indices = list(range(start, start + fetch))
io_pending[i].append(io_pool.submit(_io_call, perm_i, indices))
abs_start = initial_positions[i] + start
io_pending[i].append((abs_start, io_pool.submit(_io_call, perm_i, indices)))
def _fill_io(i: int) -> None:
while len(io_pending[i]) < max_prefetch and fetch_head[i] < split_sizes[i]:
@@ -338,15 +414,72 @@ class StreamingDataset(IterableDataset):
def _drain_io(i: int) -> None:
"""Move completed I/O futures into raw_batches non-blockingly."""
while io_pending[i] and io_pending[i][0].done():
raw_batches[i].append(io_pending[i].popleft().result())
while io_pending[i] and io_pending[i][0][1].done():
abs_start, fut = io_pending[i].popleft()
raw_batches[i].append((abs_start, fut.result()))
# ── Stage 2 helpers ───────────────────────────────────────────────────
def _tx_call_guarded(batch):
on_error = self._on_transform_error
def _should_skip(exc: Exception) -> bool:
if on_error == "raise":
return False
if callable(on_error):
return bool(on_error(exc))
return True # "skip" or "warn"
def _check_row_count(rows: list, num_rows: int) -> None:
if len(rows) != num_rows:
raise ValueError(
f"transform returned {len(rows)} rows for a batch of "
f"{num_rows}; transforms must return exactly one output "
"row per input row. To drop bad rows, raise inside the "
"transform and pass on_transform_error='skip'."
)
def _transform_isolated(abs_start, batch, batch_exc):
"""Re-run the transform on single-row slices, dropping failures."""
out = []
skipped = 0
first_exc = None
for j in range(batch.num_rows):
try:
rows = list(final_transform(batch.slice(j, 1)))
except Exception as exc:
if not _should_skip(exc):
raise
skipped += 1
if first_exc is None:
first_exc = exc
continue
_check_row_count(rows, 1)
out.append((abs_start + j, rows[0]))
self._rows_skipped += skipped
if skipped and on_error == "warn":
logger.warning(
"Skipped %d of %d rows whose transform failed (first error: %r)",
skipped,
batch.num_rows,
first_exc if first_exc is not None else batch_exc,
)
return out
def _transform_batch(abs_start, batch):
"""Apply the transform, returning [(abs_pos, row), ...]."""
try:
rows = list(final_transform(batch))
except Exception as exc:
if not _should_skip(exc):
raise
return _transform_isolated(abs_start, batch, exc)
_check_row_count(rows, batch.num_rows)
return [(abs_start + j, row) for j, row in enumerate(rows)]
def _tx_call_guarded(abs_start, batch):
try:
t0 = time.perf_counter()
result = final_transform(batch)
result = _transform_batch(abs_start, batch)
self._transform_time += time.perf_counter() - t0
return result
finally:
@@ -355,8 +488,8 @@ class StreamingDataset(IterableDataset):
def _try_submit_tx(i: int) -> None:
"""Submit transforms for raw_batches[i] up to available capacity."""
while raw_batches[i] and tx_semaphore.acquire(blocking=False):
batch = raw_batches[i].popleft()
tx_pending[i].append(tx_pool.submit(_tx_call_guarded, batch))
abs_start, batch = raw_batches[i].popleft()
tx_pending[i].append(tx_pool.submit(_tx_call_guarded, abs_start, batch))
def _drain_tx(i: int) -> None:
"""Move completed transform futures into cooked non-blockingly."""
@@ -384,11 +517,14 @@ class StreamingDataset(IterableDataset):
# Acquire a transform slot (may block briefly if all
# transform_workers are busy with other splits).
tx_semaphore.acquire()
batch = raw_batches[i].popleft()
tx_pending[i].append(tx_pool.submit(_tx_call_guarded, batch))
abs_start, batch = raw_batches[i].popleft()
tx_pending[i].append(
tx_pool.submit(_tx_call_guarded, abs_start, batch)
)
elif io_pending[i]:
# Block on the oldest in-flight I/O fetch.
raw_batches[i].append(io_pending[i].popleft().result())
abs_start, fut = io_pending[i].popleft()
raw_batches[i].append((abs_start, fut.result()))
_advance(i)
else:
break # split exhausted
@@ -407,15 +543,28 @@ class StreamingDataset(IterableDataset):
_fill_io(i)
while True:
# Stop when any split is exhausted (all exhaust
# simultaneously: equal split sizes + round-robin).
if any(local_consumed[i] >= split_sizes[i] for i in range(n)):
# A cycle only runs if every split can still produce a
# row. Without skips all splits exhaust simultaneously
# (equal split sizes + round-robin); when
# on_transform_error drops rows a split can run dry
# early, ending the epoch at the last complete cycle.
# This check only sees splits owned by this rank/worker
# (my_splits) — there is no cross-rank coordination, so
# a different rank with fewer skipped rows keeps going;
# see the on_transform_error docstring.
exhausted = False
for i in range(n):
_ensure_cooked(i)
if not cooked[i]:
exhausted = True
break
if exhausted:
break
for i in range(n):
_ensure_cooked(i)
row = cooked[i].popleft()
pos, row = cooked[i].popleft()
local_consumed[i] += 1
pos_consumed[i] = pos + 1
_advance(i)
# After the last split in each cycle: update the
@@ -424,21 +573,39 @@ class StreamingDataset(IterableDataset):
# even when __iter__ runs in a worker process.
if i == n - 1:
self._resume_offset = initial_offset + local_consumed[i]
for j, split_idx in enumerate(my_splits):
self._resume_positions[split_idx] = pos_consumed[j]
ws = self._worker_stats
ws[0] = sum(
split_sizes[j] - fetch_head[j] for j in range(n)
)
ws[1] = sum(
batch.num_rows for q in raw_batches for batch in q
batch.num_rows
for q in raw_batches
for _, batch in q
)
ws[2] = sum(len(q) for q in cooked)
ws[3] = sum(local_consumed)
ws[4] = self._bytes_loaded
ws[5] = int(self._fetch_time * 1_000_000)
ws[6] = int(self._transform_time * 1_000_000)
ws[7] = self._rows_skipped
yield row
finally:
# Final stats flush: the per-cycle write above never runs
# when iteration ends mid-cycle (e.g. a split whose rows
# were all skipped before completing a single cycle), so
# counters like rows_skipped would otherwise be stale.
ws = self._worker_stats
ws[0] = sum(split_sizes[j] - fetch_head[j] for j in range(n))
ws[1] = 0 # queue-depth properties document 0 when idle
ws[2] = 0
ws[3] = sum(local_consumed)
ws[4] = self._bytes_loaded
ws[5] = int(self._fetch_time * 1_000_000)
ws[6] = int(self._transform_time * 1_000_000)
ws[7] = self._rows_skipped
self._raw_batches_ref = None
self._cooked_ref = None
self._fetch_head_ref = None
@@ -492,7 +659,7 @@ class StreamingDataset(IterableDataset):
batches. Returns 0 when not iterating.
"""
if self._raw_batches_ref is not None:
return sum(batch.num_rows for q in self._raw_batches_ref for batch in q)
return sum(batch.num_rows for q in self._raw_batches_ref for _, batch in q)
return int(self._worker_stats[1])
@property
@@ -522,6 +689,19 @@ class StreamingDataset(IterableDataset):
)
return int(self._worker_stats[0])
@property
def rows_skipped(self) -> int:
"""Number of rows dropped because their transform raised an exception.
Only ever non-zero when ``on_transform_error`` is set to ``"skip"``,
``"warn"``, or a callable that returned ``True``. Accumulates across
multiple iterations of the same dataset instance and is never reset
automatically.
"""
if self._raw_batches_ref is not None:
return self._rows_skipped
return int(self._worker_stats[7])
@property
def consumed_rows(self) -> int:
"""Number of rows already yielded to the caller across all splits.
@@ -587,12 +767,27 @@ class StreamingDataset(IterableDataset):
every split has been consumed the same number of times (by the
round-robin design), so the per-split count is a single uniform value
that is identical across all ranks and DataLoader workers.
``positions_consumed_per_split`` records how far into each split's
permutation iteration has advanced. It only differs from
``samples_consumed_per_split`` when ``on_transform_error`` skipped
rows, in which case entries are exact for the splits this instance
iterated and a lower bound (the sample count) for splits owned by
other ranks or workers. Combine the state dicts from all ranks with
[merge_state_dicts][lancedb.streaming.StreamingDataset.merge_state_dicts]
to recover the exact value for every split before resuming on a
different topology.
"""
positions = [
self._resume_positions.get(split, self._resume_offset)
for split in range(self._num_splits)
]
return {
"shuffle_seed": self._shuffle_seed,
"num_splits": self._num_splits,
"epoch": self._epoch,
"samples_consumed_per_split": [self._resume_offset] * self._num_splits,
"positions_consumed_per_split": positions,
}
def load_state_dict(self, state: dict) -> None:
@@ -618,3 +813,96 @@ class StreamingDataset(IterableDataset):
self._resume_offset = consumed[0] if consumed else 0
else:
self._resume_offset = int(consumed)
# Older checkpoints predate positions_consumed_per_split; without
# skipped rows positions equal sample counts, so falling back to
# _resume_offset (the .get default in __iter__) is exact.
positions = state.get("positions_consumed_per_split")
if positions is None:
self._resume_positions = {}
else:
self._resume_positions = {
split: int(pos) for split, pos in enumerate(positions)
}
@staticmethod
def merge_state_dicts(states: list[dict]) -> dict:
"""Merge state dicts saved by different ranks into one exact state.
Only needed when ``on_transform_error`` skips rows in multi-rank
training: each rank then knows the exact permutation position only for
its own splits, and records a lower bound for the rest. Because
exactly one rank owns each split, the elementwise maximum across all
ranks' ``positions_consumed_per_split`` recovers the exact position of
every split. Without skipped rows every rank's state is already
identical and merging is a no-op.
Raises ``ValueError`` if the states are empty or were not produced by
the same run (mismatched seed, split count, epoch, or sample counts).
The merge is always all-to-all and topology-agnostic: collect the
``state_dict()`` from every rank of the *previous* run into one list,
merge that whole list, and hand the identical merged result to every
rank of the *next* run — regardless of whether the rank count grew,
shrank, or stayed the same. There is no pairwise or subset merging
step, because each split's exact position is only known to whichever
rank owned that split, and the elementwise maximum needs every rank's
contribution to be correct.
For example, checkpointing 8 ranks and resuming on 4 (the same
pattern applies when growing, e.g. 4 ranks resuming on 8)::
states = [ds.state_dict() for ds in previous_run_datasets] # 8
merged = StreamingDataset.merge_state_dicts(states)
for ds in resumed_datasets: # now only 4 ranks
ds.load_state_dict(merged) # same dict on every rank
The rank count on either side never affects the merge itself, since
``merge_state_dicts`` only cares about the list of states it is
given. Each split's position is recovered by elementwise maximum;
here rank 0 owned split 0 (and skipped two rows there) while rank 1
owned split 1 (and skipped one row):
>>> rank0 = {
... "shuffle_seed": 0, "num_splits": 2, "epoch": 0,
... "samples_consumed_per_split": [3, 3],
... "positions_consumed_per_split": [5, 3],
... }
>>> rank1 = {
... "shuffle_seed": 0, "num_splits": 2, "epoch": 0,
... "samples_consumed_per_split": [3, 3],
... "positions_consumed_per_split": [3, 4],
... }
>>> merged = StreamingDataset.merge_state_dicts([rank0, rank1])
>>> merged["positions_consumed_per_split"]
[5, 4]
"""
if not states:
raise ValueError("merge_state_dicts requires at least one state dict")
first = states[0]
for state in states[1:]:
for key in ("shuffle_seed", "num_splits", "epoch"):
if state[key] != first[key]:
raise ValueError(
f"{key} mismatch across state dicts: "
f"{state[key]} != {first[key]}"
)
if (
state["samples_consumed_per_split"]
!= first["samples_consumed_per_split"]
):
raise ValueError(
"samples_consumed_per_split mismatch across state dicts; "
"state_dict() must be called at the same global step "
"boundary on every rank"
)
merged = dict(first)
all_positions = [
state.get(
"positions_consumed_per_split", state["samples_consumed_per_split"]
)
for state in states
]
merged["positions_consumed_per_split"] = [
max(per_split) for per_split in zip(*all_positions)
]
return merged
+119
View File
@@ -185,6 +185,7 @@ if TYPE_CHECKING:
LsmWriteSpec,
MergeResult,
UpdateResult,
_FunctionCall,
)
from .index import IndexConfig
import pandas
@@ -1007,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.
@@ -2849,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,
@@ -5066,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.
@@ -1456,6 +1456,408 @@ def test_shuffle_clump_size_yields_all_rows(lance_table):
)
# ---------------------------------------------------------------------------
# on_transform_error tests
# ---------------------------------------------------------------------------
class BadRowError(ValueError):
"""Raised by the failing transforms below when a batch contains a bad id."""
def _failing_transform(bad_ids: set):
"""A transform that raises BadRowError whenever the batch has a bad id.
Raises on the full batch and on any single-row slice containing a bad id,
so per-row isolation drops exactly the bad rows.
"""
def transform(batch: pa.RecordBatch) -> list:
ids = batch.column("id").to_pylist()
bad = sorted(set(ids) & bad_ids)
if bad:
raise BadRowError(f"bad ids in batch: {bad}")
return [{"id": i} for i in ids]
return transform
def _sequential_split_members(table) -> list[list[int]]:
"""Return each split's ids in yield order for shuffle=False.
With a single rank and no workers the round-robin yields one row per split
per cycle, so item k of a clean run belongs to split k % NUM_SPLITS.
"""
ds = StreamingDataset(table, num_splits=NUM_SPLITS, shuffle=False)
members: list[list[int]] = [[] for _ in range(NUM_SPLITS)]
for k, row in enumerate(ds):
members[k % NUM_SPLITS].append(row["id"])
return members
def test_on_transform_error_default_raises(lance_table):
"""By default a transform exception propagates and aborts iteration."""
ds = StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle_seed=SHUFFLE_SEED,
transform=_failing_transform({7}),
)
with pytest.raises(BadRowError):
list(ds)
def test_on_transform_error_invalid_value(lance_table):
with pytest.raises(ValueError, match="on_transform_error"):
StreamingDataset(lance_table, num_splits=NUM_SPLITS, on_transform_error="bogus")
def test_on_transform_error_skip_drops_bad_rows(lance_table):
"""With one bad row per split, 'skip' yields every good row exactly once
and counts the dropped rows in rows_skipped."""
members = _sequential_split_members(lance_table)
bad_ids = {members[i][4] for i in range(NUM_SPLITS)}
ds = StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle=False,
transform=_failing_transform(bad_ids),
on_transform_error="skip",
)
assert ds.rows_skipped == 0
ids = [row["id"] for row in ds]
assert sorted(ids) == sorted(set(range(NUM_ROWS)) - bad_ids)
assert ds.rows_skipped == NUM_SPLITS
def test_on_transform_error_skip_uneven_ends_at_last_complete_cycle(lance_table):
"""When one split loses more rows than the others, the epoch ends at the
last cycle where every split still has a row no crash, no bad rows, and
every step remains one sample per split."""
members = _sequential_split_members(lance_table)
bad_ids = set(members[0][:3]) # all 3 bad rows in split 0
ds = StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle=False,
transform=_failing_transform(bad_ids),
on_transform_error="skip",
)
items = [row["id"] for row in ds]
rows_per_split = NUM_ROWS // NUM_SPLITS
expected_cycles = rows_per_split - len(bad_ids)
assert len(items) == expected_cycles * NUM_SPLITS
assert len(set(items)) == len(items), "duplicate samples yielded"
assert not set(items) & bad_ids, "a bad row was yielded"
# Split 0 contributed exactly its surviving rows, in order, one per cycle.
survivors = [i for i in members[0] if i not in bad_ids]
assert items[0::NUM_SPLITS] == survivors[:expected_cycles]
def test_on_transform_error_warn_logs(lance_table, caplog):
"""'warn' skips like 'skip' but logs a warning for the failing batch."""
members = _sequential_split_members(lance_table)
bad_ids = {members[i][3] for i in range(NUM_SPLITS)}
ds = StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle=False,
transform=_failing_transform(bad_ids),
on_transform_error="warn",
)
with caplog.at_level(logging.WARNING, logger="lancedb.streaming"):
items = list(ds)
assert len(items) == NUM_ROWS - NUM_SPLITS
assert ds.rows_skipped == NUM_SPLITS
assert "Skipped" in caplog.text
assert "BadRowError" in caplog.text
def test_on_transform_error_callable_selective(lance_table):
"""A callable handler can skip expected errors and re-raise the rest."""
members = _sequential_split_members(lance_table)
bad_ids = {members[i][0] for i in range(NUM_SPLITS)}
handled: list[Exception] = []
def handler(exc: Exception) -> bool:
handled.append(exc)
return isinstance(exc, BadRowError)
ds = StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle=False,
transform=_failing_transform(bad_ids),
on_transform_error=handler,
)
items = list(ds)
assert len(items) == NUM_ROWS - NUM_SPLITS
assert handled and all(isinstance(exc, BadRowError) for exc in handled)
def broken_transform(batch: pa.RecordBatch) -> list:
raise TypeError("boom")
ds2 = StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle=False,
transform=broken_transform,
on_transform_error=handler,
)
with pytest.raises(TypeError, match="boom"):
list(ds2)
def test_transform_wrong_row_count_raises(lance_table):
"""A transform that returns the wrong number of rows is an error even with
on_transform_error='skip' silent shrinkage would corrupt accounting."""
def drops_rows(batch: pa.RecordBatch) -> list:
return batch.column("id").to_pylist()[:-1]
ds = StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle_seed=SHUFFLE_SEED,
transform=drops_rows,
on_transform_error="skip",
)
with pytest.raises(ValueError, match="one output row per input row"):
list(ds)
def test_skip_deterministic_across_runs(lance_table):
"""With a fixed seed, skipping produces the identical sample sequence on
every run skips are data-dependent, not run-dependent."""
bad_ids = {5, 17, 46}
def run() -> tuple[list[int], int]:
ds = StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle_seed=SHUFFLE_SEED,
transform=_failing_transform(bad_ids),
on_transform_error="skip",
)
return [row["id"] for row in ds], ds.rows_skipped
ids_a, skipped_a = run()
ids_b, skipped_b = run()
assert ids_a == ids_b
assert skipped_a == skipped_b
assert not set(ids_a) & bad_ids
def test_skip_elastic_det_across_world_sizes(lance_table):
"""With equal bad-row counts per split, skipping preserves the full
elastic-determinism guarantee: identical global batches at every step for
every compatible world_size."""
members = _sequential_split_members(lance_table)
bad_ids = {members[i][6] for i in range(NUM_SPLITS)}
def collect(world_size: int) -> list[frozenset[int]]:
micro = GLOBAL_BATCH_SIZE // world_size
iters = [
iter(
StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle=False,
rank=rank,
world_size=world_size,
transform=_failing_transform(bad_ids),
on_transform_error="skip",
)
)
for rank in range(world_size)
]
_STOP = object()
batches: list[frozenset[int]] = []
while True:
step_samples: set[int] = set()
exhausted = 0
for it in iters:
for _ in range(micro):
val = next(it, _STOP)
if val is _STOP:
exhausted += 1
break
step_samples.add(val["id"])
if exhausted == len(iters):
break
assert exhausted == 0, (
"Rank iterators exhausted at different steps despite equal "
"bad-row counts per split"
)
batches.append(frozenset(step_samples))
return batches
reference = collect(1)
assert len(reference) == NUM_ROWS // NUM_SPLITS - 1
for ws in (2, 3, 4):
assert collect(ws) == reference, f"world_size={ws} diverged"
def test_resumability_with_skips_same_topology(lance_table):
"""Checkpointing mid-epoch with skipped rows resumes exactly: no sample
repeated, no sample lost, skipped rows stay skipped."""
members = _sequential_split_members(lance_table)
# Uneven skips: positions diverge across splits (2 bad in split 0, 1 in
# split 5), which only a position-based checkpoint can resume exactly.
bad_ids = {members[0][2], members[0][3], members[5][7]}
kwargs = dict(
num_splits=NUM_SPLITS,
shuffle=False,
transform=_failing_transform(bad_ids),
on_transform_error="skip",
)
reference = [row["id"] for row in StreamingDataset(lance_table, **kwargs)]
rows_per_split = NUM_ROWS // NUM_SPLITS
assert len(reference) == (rows_per_split - 2) * NUM_SPLITS
steps = 3
ds = StreamingDataset(lance_table, **kwargs)
it = iter(ds)
consumed = [next(it)["id"] for _ in range(steps * NUM_SPLITS)]
checkpoint = ds.state_dict()
it.close()
# Split 0 skipped positions 2 and 3 within its first 3 yields; split 5's
# bad row is beyond the checkpoint. Everything else is at 3 = the sample
# count.
positions = checkpoint["positions_consumed_per_split"]
assert positions[0] == 5
assert positions[1:] == [3] * (NUM_SPLITS - 1)
assert checkpoint["samples_consumed_per_split"] == [3] * NUM_SPLITS
ds2 = StreamingDataset(lance_table, **kwargs)
ds2.load_state_dict(checkpoint)
resumed = [row["id"] for row in ds2]
assert consumed == reference[: steps * NUM_SPLITS]
assert resumed == reference[steps * NUM_SPLITS :]
def test_resumability_with_skips_elastic_merge(lance_table):
"""Elastic resume with skips: each rank's checkpoint knows exact positions
only for its own splits; merge_state_dicts recovers the global state, and
a run on a different world_size continues exactly."""
members = _sequential_split_members(lance_table)
# Bad rows early in split 0 (rank 0) and split 6 (rank 1 of a ws=2 run) so
# both ranks' position vectors diverge before the checkpoint.
bad_ids = {members[0][0], members[0][2], members[6][1]}
kwargs = dict(
num_splits=NUM_SPLITS,
shuffle=False,
transform=_failing_transform(bad_ids),
on_transform_error="skip",
)
reference = [row["id"] for row in StreamingDataset(lance_table, **kwargs)]
steps = 3
world_size = 2
micro = GLOBAL_BATCH_SIZE // world_size
datasets = [
StreamingDataset(lance_table, rank=rank, world_size=world_size, **kwargs)
for rank in range(world_size)
]
iters = [iter(ds) for ds in datasets]
seen: list[frozenset[int]] = []
for _ in range(steps):
step_samples = set()
for it in iters:
for _ in range(micro):
step_samples.add(next(it)["id"])
seen.append(frozenset(step_samples))
states = [ds.state_dict() for ds in datasets]
for it in iters:
it.close()
merged = StreamingDataset.merge_state_dicts(states)
expected_positions = [3] * NUM_SPLITS
expected_positions[0] = 5 # skipped positions 0 and 2
expected_positions[6] = 4 # skipped position 1
assert merged["positions_consumed_per_split"] == expected_positions
# The first 3 global batches match the world_size=1 reference.
ref_batches = [
frozenset(reference[s * NUM_SPLITS : (s + 1) * NUM_SPLITS])
for s in range(len(reference) // NUM_SPLITS)
]
assert seen == ref_batches[:steps]
# Resume on world_size=1 from the merged state.
ds_resume = StreamingDataset(lance_table, **kwargs)
ds_resume.load_state_dict(merged)
resumed = [row["id"] for row in ds_resume]
assert resumed == reference[steps * NUM_SPLITS :]
def test_rows_skipped_flushed_when_split_entirely_bad(lance_table):
"""A split whose rows all fail never completes a cycle, so the epoch ends
immediately but rows_skipped must still report the drops after the
iterator exits (the shared-memory counter is flushed on exhaustion)."""
members = _sequential_split_members(lance_table)
bad_ids = set(members[0]) # every row of split 0 is bad
ds = StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle=False,
transform=_failing_transform(bad_ids),
on_transform_error="skip",
)
assert list(ds) == []
assert ds.rows_skipped == len(bad_ids)
def test_merge_state_dicts_validates_consistency(lance_table):
ds = StreamingDataset(lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED)
state = ds.state_dict()
other = dict(state, shuffle_seed=SHUFFLE_SEED + 1)
with pytest.raises(ValueError, match="shuffle_seed mismatch"):
StreamingDataset.merge_state_dicts([state, other])
with pytest.raises(ValueError, match="at least one"):
StreamingDataset.merge_state_dicts([])
def test_load_state_dict_without_positions_key(lance_table):
"""Checkpoints from before positions_consumed_per_split existed still
resume exactly (positions equal sample counts when nothing is skipped)."""
reference = [
row["id"]
for row in StreamingDataset(
lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED
)
]
steps = 4
ds = StreamingDataset(lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED)
it = iter(ds)
for _ in range(steps * NUM_SPLITS):
next(it)
checkpoint = ds.state_dict()
it.close()
del checkpoint["positions_consumed_per_split"]
ds2 = StreamingDataset(
lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED
)
ds2.load_state_dict(checkpoint)
resumed = [row["id"] for row in ds2]
assert resumed == reference[steps * NUM_SPLITS :]
def test_num_splits_defaults_to_world_size(lance_table):
"""Omitting num_splits gives world_size splits (one per rank)."""
ds = StreamingDataset(
@@ -0,0 +1,372 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python exact Function handle call authoring (FF-028)."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import Iterator
from typing import Any, Callable
import pyarrow as pa
import pytest
import lancedb
from lancedb import _lancedb as _native
from lancedb.expr import Expr, col, func, lit
_CALL_PATH = "/v1/functions/lookup"
_CALL_CATALOG_NAME = "text.normalize.call-name"
_CALL_FUNCTION_ID = "fn.exact.call-handle"
_LITERAL_PAYLOAD_SENTINEL = "LITERAL_PAYLOAD_SENTINEL_call_xyz_42"
_INT_PAYLOAD_SENTINEL = 2_147_000_123
# Pinned Rust-canonical schema-only type IPC (base64).
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_LIST_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////+4AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAABAAAANz///8c"
"AAAADAAAAAAAAQxcAAAAAQAAABwAAAAEAAQABAAAABAAFAAQAA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECH"
"AAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAQAAABpdGVtAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////8AAAAAFAAAAAAAAAAMABQAEgAMAAgABAAMAAAAnAAAAKAAAAAQAAAAAAAEAAgACAAAAAQACAAAAAQAAAA"
"BAAAABAAAANz///8cAAAADAAAAAAAAQxcAAAAAQAAABwAAAAEAAQABAAAABAAFAAQAA4ADwAEAAAACAAQAAAA"
"GAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAQAAABpdGVtAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAAAwAAAAEFSUk9XMQ=="
)
_OVERDESIGN_ATTRS = (
"id",
"function_id",
"name",
"connection",
"table",
"snapshot",
"field_id",
"field_ids",
"job",
"job_id",
"artifact",
"digest",
"retry_key",
"idempotency_key",
"user_version",
"execute",
"status",
"wait",
"cancel",
"to_json",
"_to_json",
"serialize",
"geneva",
)
def _sample_function_wire(
*,
function_id: str = _CALL_FUNCTION_ID,
parameters: list[dict[str, str]] | None = None,
output_type_ipc: str = _UTF8_TYPE_IPC_B64,
) -> dict[str, Any]:
return {
"format_version": 1,
"id": function_id,
"signature": {
"parameters": parameters
or [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": output_type_ipc,
"nullable": True,
},
},
}
def _lookup_success_body(function: dict[str, Any] | None = None) -> bytes:
return json.dumps({"function": function or _sample_function_wire()}).encode("utf-8")
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _lookup_function(function: dict[str, Any] | None = None):
body = _lookup_success_body(function)
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _CALL_PATH
_read_body(request)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(body)
with _mock_remote_db(handler) as db:
return db.functions.get(_CALL_CATALOG_NAME)
def _authored_call_type():
cls = getattr(_native, "_FunctionCall", None)
if cls is None:
pytest.fail("lancedb._lancedb._FunctionCall is missing")
return cls
def _exception_text(exc: BaseException) -> str:
parts = [str(exc), repr(exc)]
current: BaseException | None = exc
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(f"{type(current).__name__}: {current}")
current = current.__cause__ or current.__context__
return "\n".join(parts)
def test_function_keyword_call_returns_private_frozen_authored_value():
function = _lookup_function()
assert callable(function)
authored = function(text=col("text"), limit=8)
authored_type = _authored_call_type()
assert type(authored) is authored_type
assert authored_type.__module__ == "lancedb._lancedb"
assert authored_type.__name__ == "_FunctionCall"
# Keyword order must not matter; bindings store/render in signature order.
authored_reversed = function(limit=8, text=col("text"))
assert type(authored_reversed) is authored_type
rendered = repr(authored_reversed)
assert rendered.index("text=") < rendered.index("limit=")
assert 'text=field("text")' in rendered
assert "limit=literal(Int32, null=false)" in rendered
def test_function_call_rejects_positional_missing_and_unknown_args():
function = _lookup_function()
with pytest.raises(TypeError, match="keyword"):
function(col("text"), 8)
with pytest.raises((TypeError, ValueError), match="limit"):
function(text=col("text"))
with pytest.raises((TypeError, ValueError), match="text"):
function(limit=8)
with pytest.raises((TypeError, ValueError), match="unknown|extra"):
function(text=col("text"), limit=8, extra=1)
def test_function_call_accepts_direct_case_sensitive_column_and_rejects_complex_exprs():
function = _lookup_function()
authored = function(text=col("firstName"), limit=1)
assert type(authored) is _authored_call_type()
rendered = repr(authored)
assert 'text=field("firstName")' in rendered
assert "limit=literal(Int32, null=false)" in rendered
complex_exprs = (
col("text") + lit("x"),
col("text").cast(pa.string()),
func("lower", col("text")),
col("text") == lit("x"),
col("text").lower(),
)
for expr in complex_exprs:
with pytest.raises((TypeError, ValueError)):
function(text=expr, limit=1)
# Raw native PyExpr is not the public col() wrapper.
with pytest.raises((TypeError, ValueError)):
function(text=col("text")._inner, limit=1)
# Non-expression / non-literal objects are rejected for field-shaped misuse
# when a column binding is required; plain strings are literals for utf8.
with pytest.raises((TypeError, ValueError)):
function(text=object(), limit=1)
def test_function_call_plain_literal_declared_type_null_and_nested():
function = _lookup_function()
authored = function(text="hello", limit=7)
assert type(authored) is _authored_call_type()
rendered = repr(authored)
assert "text=literal(Utf8, null=false)" in rendered
assert "limit=literal(Int32, null=false)" in rendered
# Plain Python int normalizes to declared Int32 and non-null.
authored_int32 = function(text="hello", limit=2_147_483_647)
assert type(authored_int32) is _authored_call_type()
rendered_int32 = repr(authored_int32)
assert "limit=literal(Int32, null=false)" in rendered_int32
assert "Int64" not in rendered_int32
assert "2147483647" not in rendered_int32
# Plain None keeps each declared parameter type with null=true.
authored_null = function(text=None, limit=None)
assert type(authored_null) is _authored_call_type()
rendered_null = repr(authored_null)
assert "text=literal(Utf8, null=true)" in rendered_null
assert "limit=literal(Int32, null=true)" in rendered_null
list_function = _lookup_function(
_sample_function_wire(
parameters=[
{"name": "values", "data_type_ipc": _LIST_INT32_TYPE_IPC_B64},
]
)
)
authored_list = list_function(values=[1, 2, 3])
assert type(authored_list) is _authored_call_type()
rendered_list = repr(authored_list)
assert "values=literal(List(Int32), null=false)" in rendered_list
assert "[1, 2, 3]" not in rendered_list
authored_list_null = list_function(values=None)
assert type(authored_list_null) is _authored_call_type()
rendered_list_null = repr(authored_list_null)
assert "values=literal(List(Int32), null=true)" in rendered_list_null
def test_function_call_direct_literal_expr_exact_type_only():
function = _lookup_function()
# lit(int) is Int64 in the expression builder; int32 parameter must reject it.
with pytest.raises((TypeError, ValueError), match="limit|int32|type") as raised:
function(text="hello", limit=lit(8))
reject_text = _exception_text(raised.value)
assert "Int64" in reject_text or "int64" in reject_text.lower()
assert "Int32" in reject_text or "int32" in reject_text.lower()
# Exact utf8 literal expression is accepted and stored as Utf8/non-null.
authored = function(text=lit("hello"), limit=8)
assert type(authored) is _authored_call_type()
rendered = repr(authored)
assert "text=literal(Utf8, null=false)" in rendered
assert "limit=literal(Int32, null=false)" in rendered
assert "hello" not in rendered
# Cast / arithmetic around a literal is not a direct Literal node.
with pytest.raises((TypeError, ValueError)):
function(text=lit("hello").cast(pa.string()), limit=8)
def test_function_call_conversion_error_and_repr_are_payload_free():
function = _lookup_function()
with pytest.raises((TypeError, ValueError)) as raised:
function(text="ok", limit=_LITERAL_PAYLOAD_SENTINEL)
text = _exception_text(raised.value)
assert _LITERAL_PAYLOAD_SENTINEL not in text
assert "limit" in text
assert "int32" in text.lower() or "Int32" in text
authored = function(text=_LITERAL_PAYLOAD_SENTINEL, limit=_INT_PAYLOAD_SENTINEL)
rendered = f"{authored!r}\n{authored!s}"
assert _LITERAL_PAYLOAD_SENTINEL not in rendered
assert str(_INT_PAYLOAD_SENTINEL) not in rendered
assert "text=literal(Utf8, null=false)" in rendered
assert "limit=literal(Int32, null=false)" in rendered
assert type(authored).__name__ == "_FunctionCall"
assert "_FunctionCall" in rendered
def test_function_call_private_type_nonconstructible_immutable_and_not_exported():
function = _lookup_function()
authored = function(text=col("text"), limit=1)
authored_type = _authored_call_type()
assert "_FunctionCall" not in getattr(lancedb, "__all__", [])
assert not hasattr(lancedb, "_FunctionCall")
assert getattr(_native, "_FunctionCall", None) is authored_type
with pytest.raises(TypeError):
authored_type()
for attr in _OVERDESIGN_ATTRS:
assert not hasattr(authored, attr)
for attr in ("function", "bindings", "arguments", "parameters", "text", "limit"):
with pytest.raises(AttributeError):
setattr(authored, attr, None)
# Existing Function handle stays frozen / connection-free / name-free.
assert not hasattr(function, "name")
assert not hasattr(function, "connection")
with pytest.raises(AttributeError):
function.id = "mutated"
def test_function_call_does_not_change_col_query_expression_behavior():
# Regression guard: authoring must not alter public col()/Expr query behavior.
expr = col("firstName") > lit(1)
assert isinstance(expr, Expr)
assert expr.to_sql() == "(`firstName` > 1)"
@@ -0,0 +1,181 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
from __future__ import annotations
import pyarrow as pa
from lancedb import udf
@udf(
inputs={"value": pa.int64()},
output=pa.int64(),
python="3.12",
packages=["pyarrow==24.0.0"],
output_nullable=True,
)
def double_nullable(value):
if value is None:
return None
return value * 2
def test_first_class_function_enterprise_lifecycle():
import json
import os
import uuid
from datetime import timedelta
import pytest
import lancedb
from lancedb.exceptions import FunctionError
from lancedb.expr import col
host = os.environ.get("LANCEDB_FCF_E2E_HOST")
if not host:
pytest.skip("LANCEDB_FCF_E2E_HOST is required for the live enterprise test")
database_uri = os.environ.get("LANCEDB_FCF_E2E_DB_URI", "db://fcf-e2e-local")
api_key = os.environ.get("LANCEDB_FCF_E2E_API_KEY", "fake")
run_suffix = uuid.uuid4().hex[:12]
table_name = f"fcf_e2e_{run_suffix}"
function_name = f"fcf_e2e.double_{run_suffix}"
job_timeout = timedelta(minutes=5)
query_timeout = timedelta(seconds=30)
def connect():
return lancedb.connect(
database_uri,
api_key=api_key,
host_override=host,
)
setup_db = connect()
setup_db.create_table(
table_name,
data=pa.Table.from_pylist(
[
{"row_id": 1, "value": 2},
{"row_id": 2, "value": 5},
{"row_id": 3, "value": None},
],
schema=pa.schema(
[
pa.field("row_id", pa.int64(), nullable=False),
pa.field("value", pa.int64(), nullable=True),
]
),
),
)
registration_job = setup_db.functions.register(function_name, double_nullable)
registration_job_id = registration_job.id
assert isinstance(registration_job_id, str) and registration_job_id
registered_function = registration_job.wait(timeout=job_timeout)
assert type(registered_function) is lancedb.Function
assert isinstance(registered_function.id, str) and registered_function.id
with pytest.raises(AttributeError):
registered_function.id = "mutated"
catalog_reader = connect()
function_by_name = catalog_reader.functions.get(function_name)
function_by_id = catalog_reader.functions.get_by_id(registered_function.id)
expected_signature = ((("value", pa.int64()),), pa.int64(), True)
expected_identity = (
registered_function.id,
*expected_signature,
)
for function in (registered_function, function_by_name, function_by_id):
assert type(function) is lancedb.Function
assert (
function.id,
function.parameters,
function.output_type,
function.output_nullable,
) == expected_identity
generated_column_table = catalog_reader.open_table(table_name)
generated_column_job = generated_column_table.add_generated_column(
"derived",
registered_function(value=col("value")),
)
generated_column_job_id = generated_column_job.id
assert isinstance(generated_column_job_id, str) and generated_column_job_id
assert generated_column_job.wait(timeout=job_timeout) is None
complete_reader = connect().open_table(table_name)
complete_status = complete_reader.generated_column_status("derived")
assert complete_status == "complete"
initial_rows = sorted(
complete_reader.search()
.select(["row_id", "value", "derived"])
.limit(3)
.to_list(timeout=query_timeout),
key=lambda row: row["row_id"],
)
assert initial_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 5, "derived": 10},
{"row_id": 3, "value": None, "derived": None},
]
update_result = complete_reader.update(
where="row_id = 2",
values={"value": 7},
)
assert update_result.rows_updated == 1
incomplete_reader = connect().open_table(table_name)
incomplete_status = incomplete_reader.generated_column_status("derived")
assert incomplete_status == "incomplete"
with pytest.raises(FunctionError) as raised:
(
incomplete_reader.search()
.select(["row_id", "derived"])
.limit(3)
.to_list(timeout=query_timeout)
)
assert raised.value.code == "generated_column_incomplete"
refresh_job = incomplete_reader.refresh_generated_column("derived")
refresh_job_id = refresh_job.id
assert isinstance(refresh_job_id, str) and refresh_job_id
assert refresh_job.wait(timeout=job_timeout) is None
refreshed_reader = connect().open_table(table_name)
refreshed_status = refreshed_reader.generated_column_status("derived")
assert refreshed_status == "complete"
final_rows = sorted(
refreshed_reader.search()
.select(["row_id", "value", "derived"])
.limit(3)
.to_list(timeout=query_timeout),
key=lambda row: row["row_id"],
)
assert final_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 7, "derived": 14},
{"row_id": 3, "value": None, "derived": None},
]
evidence = {
"run_suffix": run_suffix,
"database": database_uri.removeprefix("db://"),
"table": table_name,
"function": function_name,
"function_id": registered_function.id,
"job_ids": {
"register": registration_job_id,
"add_generated_column": generated_column_job_id,
"refresh_generated_column": refresh_job_id,
},
"status_transitions": [
complete_status,
incomplete_status,
refreshed_status,
],
"final_rows": final_rows,
}
print(json.dumps(evidence, sort_keys=True, separators=(",", ":")))
@@ -0,0 +1,595 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
from __future__ import annotations
import pyarrow as pa
from lancedb import udf
_RUNNING_DEADLINE_SECONDS = 30
@udf(
inputs={"value": pa.int64()},
output=pa.int64(),
python="3.12",
packages=["pyarrow==24.0.0"],
output_nullable=True,
)
def reliable_double(value):
if value is None:
return None
return value * 2
@udf(
inputs={"value": pa.int64()},
output=pa.int64(),
python="3.12",
packages=["pyarrow==24.0.0"],
output_nullable=True,
)
def terminate_worker_on_input(value):
if value is None:
return None
try:
if len(value) == 0:
return value
except TypeError:
pass
import os
os._exit(73)
@udf(
inputs={"value": pa.int64()},
output=pa.int64(),
python="3.12",
packages=["pyarrow==24.0.0"],
output_nullable=False,
)
def slow_triple(value):
import time
time.sleep(0.02)
return value * 3
def _require_live() -> str:
import os
import pytest
host = os.environ.get("LANCEDB_FCF_E2E_HOST")
if not host:
pytest.skip(
"LANCEDB_FCF_E2E_HOST is required for live enterprise reliability tests"
)
return host
def _job_timeout():
from datetime import timedelta
return timedelta(minutes=5)
def _query_timeout():
from datetime import timedelta
return timedelta(seconds=30)
def _connect():
import os
import lancedb
return lancedb.connect(
os.environ.get("LANCEDB_FCF_E2E_DB_URI", "db://fcf-e2e-local"),
api_key=os.environ.get("LANCEDB_FCF_E2E_API_KEY", "fake"),
host_override=_require_live(),
)
def _run_names(case: str) -> tuple[str, str]:
import uuid
suffix = uuid.uuid4().hex[:12]
return f"fcf_rel_{case}_{suffix}", f"fcf_rel.{case}_{suffix}"
def _read_rows(table, columns: list[str], row_count: int) -> list[dict]:
return sorted(
table.search()
.select(columns)
.limit(row_count)
.to_list(timeout=_query_timeout()),
key=lambda row: row["row_id"],
)
def _emit_evidence(case: str, evidence: dict) -> None:
import json
print(
json.dumps(
{"case": case, **evidence},
sort_keys=True,
separators=(",", ":"),
)
)
def test_enterprise_reliability_core_lifecycle():
import pytest
import lancedb
from lancedb.exceptions import FunctionError
from lancedb.expr import col
_require_live()
table_name, function_name = _run_names("lifecycle")
setup_db = _connect()
setup_db.create_table(
table_name,
data=pa.Table.from_pylist(
[
{"row_id": 1, "value": 2},
{"row_id": 2, "value": 5},
{"row_id": 3, "value": None},
],
schema=pa.schema(
[
pa.field("row_id", pa.int64(), nullable=False),
pa.field("value", pa.int64(), nullable=True),
]
),
),
)
registration_job = setup_db.functions.register(function_name, reliable_double)
registration_job_id = registration_job.id
assert isinstance(registration_job_id, str) and registration_job_id
registered = registration_job.wait(timeout=_job_timeout())
assert type(registered) is lancedb.Function
assert isinstance(registered.id, str) and registered.id
with pytest.raises(AttributeError):
registered.id = "mutated"
catalog_reader = _connect()
by_name = catalog_reader.functions.get(function_name)
by_id = catalog_reader.functions.get_by_id(registered.id)
expected_identity = (
registered.id,
(("value", pa.int64()),),
pa.int64(),
True,
)
for function in (registered, by_name, by_id):
assert type(function) is lancedb.Function
assert (
function.id,
function.parameters,
function.output_type,
function.output_nullable,
) == expected_identity
table = catalog_reader.open_table(table_name)
create_job = table.add_generated_column(
"derived",
registered(value=col("value")),
)
create_job_id = create_job.id
assert isinstance(create_job_id, str) and create_job_id
assert create_job.wait(timeout=_job_timeout()) is None
complete_reader = _connect().open_table(table_name)
complete_status = complete_reader.generated_column_status("derived")
assert complete_status == "complete"
initial_rows = _read_rows(
complete_reader,
["row_id", "value", "derived"],
3,
)
assert initial_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 5, "derived": 10},
{"row_id": 3, "value": None, "derived": None},
]
complete_reader.update(where="row_id = 2", values={"value": 7})
incomplete_reader = _connect().open_table(table_name)
changed_rows = _read_rows(incomplete_reader, ["row_id", "value"], 3)
assert changed_rows == [
{"row_id": 1, "value": 2},
{"row_id": 2, "value": 7},
{"row_id": 3, "value": None},
]
incomplete_status = incomplete_reader.generated_column_status("derived")
assert incomplete_status == "incomplete"
with pytest.raises(FunctionError) as raised:
(
incomplete_reader.search()
.select(["row_id", "derived"])
.limit(3)
.to_list(timeout=_query_timeout())
)
assert raised.value.code == "generated_column_incomplete"
refresh_job = incomplete_reader.refresh_generated_column("derived")
refresh_job_id = refresh_job.id
assert isinstance(refresh_job_id, str) and refresh_job_id
assert refresh_job.wait(timeout=_job_timeout()) is None
refreshed_reader = _connect().open_table(table_name)
refreshed_status = refreshed_reader.generated_column_status("derived")
assert refreshed_status == "complete"
final_rows = _read_rows(
refreshed_reader,
["row_id", "value", "derived"],
3,
)
assert final_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 7, "derived": 14},
{"row_id": 3, "value": None, "derived": None},
]
_emit_evidence(
"core_lifecycle",
{
"final_rows": final_rows,
"function_id": registered.id,
"job_ids": {
"create": create_job_id,
"refresh": refresh_job_id,
"register": registration_job_id,
},
"status": [
complete_status,
incomplete_status,
refreshed_status,
],
"table": table_name,
},
)
def test_enterprise_reliability_restart_retention():
import json
import os
import pytest
import lancedb
_require_live()
raw_evidence = os.environ.get("LANCEDB_FCF_E2E_RESTART_EVIDENCE")
if not raw_evidence:
pytest.skip(
"LANCEDB_FCF_E2E_RESTART_EVIDENCE is required for restart retention"
)
try:
evidence = json.loads(raw_evidence)
except json.JSONDecodeError as error:
pytest.fail(f"LANCEDB_FCF_E2E_RESTART_EVIDENCE must be valid JSON: {error.msg}")
assert isinstance(evidence, dict), (
"LANCEDB_FCF_E2E_RESTART_EVIDENCE must be a JSON object"
)
table_name = evidence.get("table")
function_id = evidence.get("function_id")
raw_job_ids = evidence.get("job_ids")
assert isinstance(table_name, str) and table_name, (
"restart evidence must contain a non-empty table"
)
assert isinstance(function_id, str) and function_id, (
"restart evidence must contain a non-empty function_id"
)
assert isinstance(raw_job_ids, dict), (
"restart evidence must contain a job_ids object"
)
job_ids = {}
for job_kind in ("register", "create", "refresh"):
job_id = raw_job_ids.get(job_kind)
assert isinstance(job_id, str) and job_id, (
f"restart evidence must contain a non-empty job_ids.{job_kind}"
)
job_ids[job_kind] = job_id
db = _connect()
function = db.functions.get_by_id(function_id)
expected_identity = (
function_id,
(("value", pa.int64()),),
pa.int64(),
True,
)
assert type(function) is lancedb.Function
assert (
function.id,
function.parameters,
function.output_type,
function.output_nullable,
) == expected_identity
jobs = {}
for job_kind in ("register", "create", "refresh"):
job = db.get_job(job_ids[job_kind])
assert job is not None
assert job.job_id == job_ids[job_kind]
assert job.state == "finished"
assert job.failure is None
jobs[job_kind] = job
registered_result = jobs["register"].result
assert type(registered_result) is lancedb.Function
assert (
registered_result.id,
registered_result.parameters,
registered_result.output_type,
registered_result.output_nullable,
) == expected_identity
assert jobs["create"].result is None
assert jobs["refresh"].result is None
table = db.open_table(table_name)
status = table.generated_column_status("derived")
assert status == "complete"
assert table.count_rows() == 3
final_rows = _read_rows(table, ["row_id", "value", "derived"], 3)
assert final_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 7, "derived": 14},
{"row_id": 3, "value": None, "derived": None},
]
_emit_evidence(
"restart_retention",
{
"final_rows": final_rows,
"function_id": function_id,
"generated_column_status": status,
"job_ids": job_ids,
"job_states": {
job_kind: jobs[job_kind].state
for job_kind in ("register", "create", "refresh")
},
"table": table_name,
},
)
def test_enterprise_reliability_failure_atomicity_and_worker_recovery():
import pytest
import lancedb
from lancedb.exceptions import JobFailedError
from lancedb.expr import col
_require_live()
table_name, failing_function_name = _run_names("worker_failure")
_, healthy_function_name = _run_names("worker_recovery")
row_count = 4
setup_db = _connect()
setup_db.create_table(
table_name,
data=pa.table(
{
"row_id": list(range(row_count)),
"value": [1, 2, 3, 4],
}
),
)
registration_job = setup_db.functions.register(
failing_function_name,
terminate_worker_on_input,
)
failing_function = registration_job.wait(timeout=_job_timeout())
assert type(failing_function) is lancedb.Function
table = setup_db.open_table(table_name)
failed_create_job = table.add_generated_column(
"must_not_publish",
failing_function(value=col("value")),
)
failed_job_id = failed_create_job.id
assert isinstance(failed_job_id, str) and failed_job_id
with pytest.raises(JobFailedError) as raised:
failed_create_job.wait(timeout=_job_timeout())
assert raised.value.error_code == "udf_execution_failure"
first_description = _connect().get_job(failed_job_id)
second_description = _connect().get_job(failed_job_id)
for description in (first_description, second_description):
assert description is not None
assert description.job_id == failed_job_id
assert description.state == "failed"
assert description.failure is not None
assert description.failure.error_code == "udf_execution_failure"
atomic_reader = _connect().open_table(table_name)
assert "must_not_publish" not in atomic_reader.schema.names
assert _read_rows(atomic_reader, ["row_id", "value"], row_count) == [
{"row_id": 0, "value": 1},
{"row_id": 1, "value": 2},
{"row_id": 2, "value": 3},
{"row_id": 3, "value": 4},
]
healthy_registration_job = setup_db.functions.register(
healthy_function_name,
reliable_double,
)
healthy_function = healthy_registration_job.wait(timeout=_job_timeout())
assert type(healthy_function) is lancedb.Function
recovery_job = atomic_reader.add_generated_column(
"recovered",
healthy_function(value=col("value")),
)
recovery_job_id = recovery_job.id
assert isinstance(recovery_job_id, str) and recovery_job_id
assert recovery_job.wait(timeout=_job_timeout()) is None
recovered_reader = _connect().open_table(table_name)
assert "must_not_publish" not in recovered_reader.schema.names
assert recovered_reader.generated_column_status("recovered") == "complete"
recovered_rows = _read_rows(
recovered_reader,
["row_id", "value", "recovered"],
row_count,
)
assert recovered_rows == [
{"row_id": 0, "value": 1, "recovered": 2},
{"row_id": 1, "value": 2, "recovered": 4},
{"row_id": 2, "value": 3, "recovered": 6},
{"row_id": 3, "value": 4, "recovered": 8},
]
_emit_evidence(
"failure_atomicity_and_worker_recovery",
{
"failure_code": first_description.failure.error_code,
"failed_job_id": failed_job_id,
"recovered_rows": recovered_rows,
"recovery_job_id": recovery_job_id,
"table": table_name,
},
)
def test_enterprise_reliability_concurrent_refresh_fencing():
import time
import pytest
import lancedb
from lancedb.exceptions import FunctionError, JobFailedError
from lancedb.expr import col
_require_live()
table_name, function_name = _run_names("refresh_fencing")
row_count = 1024
setup_db = _connect()
setup_db.create_table(
table_name,
data=pa.table(
{
"row_id": list(range(row_count)),
"value": list(range(row_count)),
}
),
)
registration_job = setup_db.functions.register(function_name, slow_triple)
function = registration_job.wait(timeout=_job_timeout())
assert type(function) is lancedb.Function
table = setup_db.open_table(table_name)
create_job = table.add_generated_column(
"derived",
function(value=col("value")),
)
assert create_job.wait(timeout=_job_timeout()) is None
initial_reader = _connect().open_table(table_name)
assert initial_reader.generated_column_status("derived") == "complete"
initial_reader.update(where="row_id = 0", values={"value": 10_000})
incomplete_reader = _connect().open_table(table_name)
assert incomplete_reader.generated_column_status("derived") == "incomplete"
refresh_job = incomplete_reader.refresh_generated_column("derived")
refresh_job_id = refresh_job.id
assert isinstance(refresh_job_id, str) and refresh_job_id
deadline = time.monotonic() + _RUNNING_DEADLINE_SECONDS
observed_states = []
running_observations = 0
while running_observations < 2:
state = refresh_job.status()
if not observed_states or observed_states[-1] != state:
observed_states.append(state)
if state == "running":
running_observations += 1
else:
running_observations = 0
assert state not in {"finished", "failed", "cancelled"}
assert time.monotonic() < deadline
if running_observations < 2:
time.sleep(0.05)
concurrent_writer = _connect().open_table(table_name)
concurrent_writer.update(where="row_id = 1", values={"value": 20_000})
with pytest.raises(JobFailedError) as raised:
refresh_job.wait(timeout=_job_timeout())
assert raised.value.error_code == "stale_or_conflicting_input"
stale_job = _connect().get_job(refresh_job_id)
assert stale_job is not None
assert stale_job.job_id == refresh_job_id
assert stale_job.state == "failed"
assert stale_job.failure is not None
assert stale_job.failure.error_code == raised.value.error_code
if observed_states[-1] != stale_job.state:
observed_states.append(stale_job.state)
stale_reader = _connect().open_table(table_name)
stale_rows = _read_rows(stale_reader, ["row_id", "value"], row_count)
assert len(stale_rows) == row_count
for row_id, row in enumerate(stale_rows):
expected_value = 10_000 if row_id == 0 else 20_000 if row_id == 1 else row_id
assert (row["row_id"], row["value"]) == (row_id, expected_value)
assert stale_reader.generated_column_status("derived") == "incomplete"
with pytest.raises(FunctionError) as incomplete:
(
stale_reader.search()
.select(["row_id", "derived"])
.limit(row_count)
.to_list(timeout=_query_timeout())
)
assert incomplete.value.code == "generated_column_incomplete"
resubmitted_job = stale_reader.refresh_generated_column("derived")
resubmitted_job_id = resubmitted_job.id
assert isinstance(resubmitted_job_id, str) and resubmitted_job_id
assert resubmitted_job.wait(timeout=_job_timeout()) is None
final_reader = _connect().open_table(table_name)
final_status = final_reader.generated_column_status("derived")
assert final_status == "complete"
final_rows = _read_rows(
final_reader,
["row_id", "value", "derived"],
row_count,
)
assert len(final_rows) == row_count
final_checksum = 0
for row_id, row in enumerate(final_rows):
expected_value = 10_000 if row_id == 0 else 20_000 if row_id == 1 else row_id
assert (row["row_id"], row["value"], row["derived"]) == (
row_id,
expected_value,
expected_value * 3,
)
final_checksum += row["derived"]
_emit_evidence(
"concurrent_refresh_fencing",
{
"failure_code": stale_job.failure.error_code,
"final_checksum": final_checksum,
"final_status": final_status,
"observed_states": observed_states,
"resubmitted_job_id": resubmitted_job_id,
"row_count": row_count,
"stale_job_id": refresh_job_id,
"table": table_name,
},
)
@@ -0,0 +1,268 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract: Python projection of JobFailure.error_code / JobFailedError.error_code.
Public Function failures expose eight stable string categories. Asynchronous
errors remain the unified JobFailedError and JobFailureInfo. Python must
project the optional exact error_code string already supplied structurally by
Rust: preserve a known code, preserve an unknown nonempty future code
byte-for-byte, and return None for legacy failure payloads without error_code.
Never infer or override a code from message, phase, retryable, HTTP status,
job type, or state.
"""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from datetime import timedelta
from typing import Any, Callable, Optional
import pytest
import lancedb
from lancedb.exceptions import JobFailedError
_DESCRIBE_PATH = "/v1/jobs/describe"
_KNOWN_CODE = "name_or_function_not_found"
_CONFLICTING_STABLE_IN_MESSAGE = "definition_validation_failure"
_UNKNOWN_CODE = "enterprise_future_category_xyz"
_WAIT_KNOWN_CODE = "unsupported_runtime_or_capability"
_WAIT_CONFLICTING_IN_MESSAGE = "revoked_function"
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _failed_describe_body(
*,
job_id: str,
error_code: Optional[str] = None,
include_error_code: bool = True,
phase: str = "execute",
message: str = "worker died",
retryable: bool = False,
job_type: str = "create_index",
) -> dict[str, Any]:
failure: dict[str, Any] = {
"phase": phase,
"message": message,
"retryable": retryable,
}
if include_error_code:
failure["error_code"] = error_code
return {
"job_id": job_id,
"job_type": job_type,
"job_state": "FAILED",
"creation_ms": 1000,
"spec": {},
"failure": failure,
}
def _describe_handler(bodies_by_job_id: dict[str, dict[str, Any]]):
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _DESCRIBE_PATH
payload = json.loads(_read_body(request).decode("utf-8") or "{}")
job_id = payload["job_id"]
body = bodies_by_job_id.get(job_id)
if body is None:
request.send_response(404)
request.end_headers()
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
return handler
def test_get_job_failure_error_code_known_not_inferred_from_message():
"""Structural error_code wins; conflicting message text must not override."""
body = _failed_describe_body(
job_id="job-known",
error_code=_KNOWN_CODE,
phase="validate",
message=f"looks like {_CONFLICTING_STABLE_IN_MESSAGE}",
retryable=False,
)
with _mock_remote_db(_describe_handler({"job-known": body})) as db:
description = db.get_job("job-known")
assert description is not None
failure = description.failure
assert failure is not None
assert failure.error_code == _KNOWN_CODE
assert failure.error_code != _CONFLICTING_STABLE_IN_MESSAGE
assert failure.phase == "validate"
assert failure.message == f"looks like {_CONFLICTING_STABLE_IN_MESSAGE}"
assert failure.retryable is False
def test_get_job_failure_error_code_unknown_preserved_byte_for_byte():
body = _failed_describe_body(
job_id="job-unknown",
error_code=_UNKNOWN_CODE,
phase="execute",
message=f"new category mentioning {_KNOWN_CODE}",
retryable=True,
)
with _mock_remote_db(_describe_handler({"job-unknown": body})) as db:
failure = db.get_job("job-unknown").failure
assert failure.error_code == _UNKNOWN_CODE
assert failure.error_code != _KNOWN_CODE
def test_get_job_failure_error_code_absent_is_none():
"""Legacy describe payloads without error_code must not invent a category."""
body = _failed_describe_body(
job_id="job-legacy",
include_error_code=False,
phase="execute",
message=f"{_KNOWN_CODE} in logs",
retryable=True,
)
with _mock_remote_db(_describe_handler({"job-legacy": body})) as db:
failure = db.get_job("job-legacy").failure
assert failure.error_code is None
assert failure.phase == "execute"
assert failure.retryable is True
def test_sync_job_wait_job_failed_error_code_known_not_inferred():
body = _failed_describe_body(
job_id="job-wait-known",
error_code=_WAIT_KNOWN_CODE,
phase="dispatch",
message=f"{_WAIT_CONFLICTING_IN_MESSAGE} in transport logs",
retryable=False,
)
with _mock_remote_db(_describe_handler({"job-wait-known": body})) as db:
with pytest.raises(JobFailedError) as exc_info:
db.job("job-wait-known").wait(timeout=timedelta(seconds=5))
err = exc_info.value
assert isinstance(err, JobFailedError)
assert err.error_code == _WAIT_KNOWN_CODE
assert err.error_code != _WAIT_CONFLICTING_IN_MESSAGE
def test_sync_job_wait_job_failed_error_code_absent_is_none():
body = _failed_describe_body(
job_id="job-wait-legacy",
include_error_code=False,
phase="execute",
message=f"{_WAIT_KNOWN_CODE} mentioned only in message",
retryable=True,
)
with _mock_remote_db(_describe_handler({"job-wait-legacy": body})) as db:
with pytest.raises(JobFailedError) as exc_info:
db.job("job-wait-legacy").wait(timeout=timedelta(seconds=5))
assert exc_info.value.error_code is None
@pytest.mark.asyncio
async def test_async_job_wait_job_failed_error_code_unknown_preserved():
body = _failed_describe_body(
job_id="job-wait-unknown",
error_code=_UNKNOWN_CODE,
phase="execute",
message=f"future code with {_WAIT_KNOWN_CODE} in text",
retryable=False,
)
async with _mock_remote_db_async(
_describe_handler({"job-wait-unknown": body})
) as db:
with pytest.raises(JobFailedError) as exc_info:
await db.job("job-wait-unknown").wait(timeout=timedelta(seconds=5))
err = exc_info.value
assert err.error_code == _UNKNOWN_CODE
assert err.error_code != _WAIT_KNOWN_CODE
def test_job_failed_error_legacy_message_construction_error_code_is_none():
err = JobFailedError("legacy construction with only a message")
assert err.error_code is None
def test_job_failed_error_error_code_is_read_only():
err = JobFailedError("message")
with pytest.raises(AttributeError):
err.error_code = _KNOWN_CODE
@@ -0,0 +1,634 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python first-class Function catalog lookup."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable
import pyarrow as pa
import pytest
import lancedb
from lancedb import _lancedb as _native
from lancedb.remote.errors import HttpError
def _function_error_cls(*, required: bool = True):
"""Resolve FunctionError from the live module (records RED when absent)."""
from lancedb import exceptions as exc_mod
cls = getattr(exc_mod, "FunctionError", None)
if cls is None:
if required:
pytest.fail("lancedb.exceptions.FunctionError is missing")
return type("MissingFunctionError", (), {})
return cls
_LOOKUP_PATH = "/v1/functions/lookup"
_LOOKUP_CATALOG_NAME = "text.normalize.lookup-name"
_LOOKUP_FUNCTION_ID = "fn.exact.lookup-handle"
_LOOKUP_SERVER_MESSAGE_MARKER = (
"SERVER_LOOKUP_DIAGNOSTIC_MARKER name=text.normalize.lookup-name "
"id=fn.exact.lookup-handle"
)
_SENSITIVE_BODY_MARKER = "SENSITIVE_LOOKUP_BODY_MARKER"
_UNKNOWN_CODE = "enterprise_future_lookup_category_xyz"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as job-result
# tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust FileWriter.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_LOOKUP_KEYWORDS = (
"idempotency_key",
"retry_key",
"user_version",
"replace",
"expected_current_function_id",
"list",
"alias",
"lineage",
"FunctionVersion",
)
def _sample_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _LOOKUP_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _lookup_success_body(
*,
function: dict[str, Any] | None = None,
extra_outer: dict[str, Any] | None = None,
) -> bytes:
body: dict[str, Any] = {"function": function or _sample_function_wire()}
if extra_outer:
body.update(extra_outer)
return json.dumps(body).encode("utf-8")
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _exception_chain_text(exc: BaseException) -> str:
parts: list[str] = []
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(str(current))
parts.append(repr(current))
current = current.__cause__
return "\n".join(parts)
def _assert_payload_free(exc: BaseException) -> None:
text = _exception_chain_text(exc)
assert _LOOKUP_SERVER_MESSAGE_MARKER not in text
assert _LOOKUP_CATALOG_NAME not in text
assert _LOOKUP_FUNCTION_ID not in text
assert _SENSITIVE_BODY_MARKER not in text
def _assert_exact_lookup_function(function: object) -> None:
assert type(function) is lancedb.Function
assert function.id == _LOOKUP_FUNCTION_ID
assert not hasattr(function, "name")
assert function.parameters == (
("text", pa.string()),
("limit", pa.int32()),
)
assert function.output_type == pa.string()
assert function.output_nullable is True
assert _LOOKUP_CATALOG_NAME not in repr(function)
assert _LOOKUP_CATALOG_NAME not in str(function)
def _assert_name_request(raw: bytes, body: dict[str, Any]) -> None:
assert raw
assert body == {"name": _LOOKUP_CATALOG_NAME}
assert "function_id" not in body
def _assert_id_request(raw: bytes, body: dict[str, Any]) -> None:
assert raw
assert body == {"function_id": _LOOKUP_FUNCTION_ID}
assert "name" not in body
def _assert_native_lookup_methods_present() -> None:
assert hasattr(_native.Connection, "_lookup_function_by_name")
assert hasattr(_native.Connection, "_lookup_function_by_id")
assert callable(getattr(_native.Connection, "_lookup_function_by_name"))
assert callable(getattr(_native.Connection, "_lookup_function_by_id"))
def test_native_connection_exposes_private_lookup_methods():
_assert_native_lookup_methods_present()
def test_sync_remote_get_by_name_exact_request_and_function_shape():
_assert_native_lookup_methods_present()
seen: dict[str, Any] = {}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _LOOKUP_PATH
raw = _read_body(request)
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
with _mock_remote_db(handler) as db:
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
assert not hasattr(db, "get_function")
function = db.functions.get(_LOOKUP_CATALOG_NAME)
_assert_name_request(seen["raw"], seen["body"])
_assert_exact_lookup_function(function)
def test_sync_remote_get_by_id_exact_request_and_function_shape():
_assert_native_lookup_methods_present()
seen: dict[str, Any] = {}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _LOOKUP_PATH
raw = _read_body(request)
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
with _mock_remote_db(handler) as db:
function = db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
_assert_id_request(seen["raw"], seen["body"])
_assert_exact_lookup_function(function)
@pytest.mark.asyncio
async def test_async_remote_get_by_name_and_id():
_assert_native_lookup_methods_present()
name_seen: dict[str, Any] = {}
id_seen: dict[str, Any] = {}
stage = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _LOOKUP_PATH
raw = _read_body(request)
body = json.loads(raw.decode("utf-8"))
stage["n"] += 1
if stage["n"] == 1:
name_seen["raw"] = raw
name_seen["body"] = body
else:
id_seen["raw"] = raw
id_seen["body"] = body
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
async with _mock_remote_db_async(handler) as db:
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
by_name = await db.functions.get(_LOOKUP_CATALOG_NAME)
by_id = await db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
_assert_name_request(name_seen["raw"], name_seen["body"])
_assert_id_request(id_seen["raw"], id_seen["body"])
_assert_exact_lookup_function(by_name)
_assert_exact_lookup_function(by_id)
def test_sync_remote_get_accepts_additive_outer_success_fields():
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
_lookup_success_body(
extra_outer={
"server_extra": {"ok": True},
"request_echo_name": _LOOKUP_CATALOG_NAME,
}
)
)
with _mock_remote_db(handler) as db:
function = db.functions.get(_LOOKUP_CATALOG_NAME)
_assert_exact_lookup_function(function)
def test_empty_name_and_id_reject_before_transport():
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
_read_body(request)
request.send_response(500)
request.end_headers()
request.wfile.write(b"should not be reached")
with _mock_remote_db(handler) as db:
with pytest.raises(ValueError):
db.functions.get("")
with pytest.raises(ValueError):
db.functions.get_by_id("")
assert received["n"] == 0
def test_local_sync_lookup_not_implemented_without_table_mutation(tmp_path):
_assert_native_lookup_methods_present()
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
with pytest.raises(NotImplementedError):
db.functions.get(_LOOKUP_CATALOG_NAME)
with pytest.raises(NotImplementedError):
db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
assert db.list_tables().tables == before
@pytest.mark.asyncio
async def test_local_async_lookup_not_implemented_without_table_mutation(tmp_path):
_assert_native_lookup_methods_present()
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
with pytest.raises(NotImplementedError):
await db.functions.get(_LOOKUP_CATALOG_NAME)
with pytest.raises(NotImplementedError):
await db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
assert (await db.list_tables()).tables == before
def test_explicit_known_code_is_function_error_with_exact_code():
body = {
"error_code": "name_or_function_not_found",
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
"looks_like": "definition_validation_failure",
}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
request.send_response(404)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(handler) as db:
with pytest.raises(function_error) as exc_info:
db.functions.get(_LOOKUP_CATALOG_NAME)
err = exc_info.value
assert isinstance(err, function_error)
assert err.code == "name_or_function_not_found"
assert err.code != "definition_validation_failure"
_assert_payload_free(err)
def test_explicit_unknown_code_preserved_despite_status_and_message():
body = {
"error_code": _UNKNOWN_CODE,
"message": (
f"{_LOOKUP_SERVER_MESSAGE_MARKER} revoked_function "
"name_or_function_not_found"
),
}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
raw = _read_body(request)
assert json.loads(raw.decode("utf-8")) == {"function_id": _LOOKUP_FUNCTION_ID}
request.send_response(409)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(handler) as db:
with pytest.raises(function_error) as exc_info:
db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
err = exc_info.value
assert err.code == _UNKNOWN_CODE
assert err.code != "revoked_function"
assert err.code != "name_or_function_not_found"
_assert_payload_free(err)
@pytest.mark.parametrize(
"label,status,response_body",
[
(
"missing_code_404",
404,
{
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
},
),
(
"empty_code",
400,
{
"error_code": "",
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
},
),
(
"wrong_type_code",
400,
{
"error_code": 123,
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
},
),
(
"null_code",
404,
{
"error_code": None,
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
},
),
(
"non_json",
404,
f"not-json {_LOOKUP_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
),
],
)
def test_invalid_or_missing_error_code_is_payload_free_http(
label: str, status: int, response_body: object
):
del label # parametrize label for failure diagnosis only
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
request.send_response(status)
request.send_header("Content-Type", "application/json")
request.end_headers()
if isinstance(response_body, str):
request.wfile.write(response_body.encode("utf-8"))
else:
request.wfile.write(json.dumps(response_body).encode("utf-8"))
with _mock_remote_db(handler) as db:
with pytest.raises(HttpError) as exc_info:
db.functions.get(_LOOKUP_CATALOG_NAME)
err = exc_info.value
assert isinstance(err, HttpError)
assert not isinstance(err, _function_error_cls(required=False))
_assert_payload_free(err)
@pytest.mark.parametrize(
"label,response_body",
[
(
"missing_function",
{
"server_extra": True,
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
},
),
(
"null_function",
{
"function": None,
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
},
),
(
"wrong_type_function",
{
"function": "not-an-object",
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
},
),
(
"invalid_function_shape",
{
"function": {
"format_version": 1,
"id": _LOOKUP_FUNCTION_ID,
# missing signature
_SENSITIVE_BODY_MARKER: True,
}
},
),
],
)
def test_malformed_success_is_payload_free_http(label: str, response_body: dict):
del label
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(response_body).encode("utf-8"))
with _mock_remote_db(handler) as db:
with pytest.raises(HttpError) as exc_info:
db.functions.get(_LOOKUP_CATALOG_NAME)
err = exc_info.value
assert isinstance(err, HttpError)
assert not isinstance(err, _function_error_cls(required=False))
_assert_payload_free(err)
def test_function_error_surface_omits_server_marker_name_and_id():
body = {
"error_code": "name_or_function_not_found",
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
"function_id": _LOOKUP_FUNCTION_ID,
"name": _LOOKUP_CATALOG_NAME,
_SENSITIVE_BODY_MARKER: True,
}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
_read_body(request)
request.send_response(404)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(handler) as db:
with pytest.raises(function_error) as exc_info:
db.functions.get(_LOOKUP_CATALOG_NAME)
err = exc_info.value
_assert_payload_free(err)
assert getattr(err, "code", None) == "name_or_function_not_found"
def test_no_direct_db_lookup_methods_and_no_deleted_keywords():
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
_read_body(request)
request.send_response(500)
request.end_headers()
request.wfile.write(b"should not be reached")
with _mock_remote_db(handler) as db:
assert not hasattr(db, "lookup_function")
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
assert not hasattr(db, "get_function")
assert not hasattr(db.functions, "get_by_name")
assert not hasattr(db.functions, "list")
for keyword in _DELETED_LOOKUP_KEYWORDS:
with pytest.raises(TypeError):
db.functions.get(_LOOKUP_CATALOG_NAME, **{keyword: True})
with pytest.raises(TypeError):
db.functions.get_by_id(_LOOKUP_FUNCTION_ID, **{keyword: True})
assert received["n"] == 0
def test_function_error_is_not_top_level_export():
assert not hasattr(lancedb, "FunctionError")
function_error = _function_error_cls()
assert issubclass(function_error, RuntimeError)
@@ -0,0 +1,398 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for Python first-class Function registration."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from datetime import timedelta
from typing import Any, Callable
from unittest import mock
import pyarrow as pa
import pytest
import lancedb
import lancedb._udf as _udf_mod
import lancedb.job
from lancedb import FunctionCapability, udf
from lancedb.remote.errors import HttpError
_SOURCE_MARKER = "registration-source-marker-unique-xyz"
_SECRET_REFERENCE = "secret://team/registration-redact-token-xyz"
_SECRET_ENV = "REGISTER_API_TOKEN"
_NETWORK_ORIGIN = "https://api.registration-example.com"
_FUNCTION_NAME = "text.normalize"
_FUNCTION_ID_RETRY = "fn.register-retry-1"
_JOB_ID_RETRY = "job-register-retry-1"
_JOB_ID_ASYNC = "job-register-async-1"
_REGISTER_PATH = "/v1/functions/register"
_DESCRIBE_PATH = "/v1/jobs/describe"
_DELETED_REGISTER_KEYWORDS = (
"idempotency_key",
"retry_key",
"user_version",
"deterministic",
"null_policy",
"replace",
"expected_current_function_id",
)
_SPEC_KEYS = {
"format_version",
"name",
"definition",
"expected_current_function_id",
}
@udf(
inputs={"text": pa.string(), "limit": pa.int32()},
output=pa.string(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=True,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
],
)
def packable_register_normalize(text, limit):
"""registration-source-marker-unique-xyz."""
return text[:limit]
def _definition_json(fn: object) -> dict[str, Any]:
payload = _udf_mod._build_function_definition(fn)._to_json()
if isinstance(payload, bytes):
return json.loads(payload.decode("utf-8"))
assert isinstance(payload, str)
return json.loads(payload)
def _expected_register_spec(name: str, fn: object) -> dict[str, Any]:
return {
"format_version": 1,
"name": name,
"definition": _definition_json(fn),
"expected_current_function_id": None,
}
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _exception_chain_text(exc: BaseException) -> str:
parts: list[str] = []
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(str(current))
parts.append(repr(current))
current = current.__cause__
return "\n".join(parts)
def _assert_markers_absent_from_exception(exc: BaseException) -> None:
text = _exception_chain_text(exc)
assert _SOURCE_MARKER not in text
assert _SECRET_REFERENCE not in text
def _assert_exact_register_spec(body: dict[str, Any], expected: dict[str, Any]) -> None:
assert set(body) == _SPEC_KEYS
assert body == expected
assert body["format_version"] == 1
assert body["expected_current_function_id"] is None
assert _SOURCE_MARKER in json.dumps(body["definition"])
assert any(
capability.get("reference") == _SECRET_REFERENCE
for capability in body["definition"]["capabilities"]
)
def test_sync_remote_register_retries_exact_wire_and_returns_job():
expected_spec = _expected_register_spec(_FUNCTION_NAME, packable_register_normalize)
attempts: list[dict[str, Any]] = []
describe_calls: list[dict[str, Any]] = []
function_result_wire = {
"kind": "function",
"format_version": 1,
"function": {
"format_version": 1,
"id": _FUNCTION_ID_RETRY,
"signature": expected_spec["definition"]["signature"],
},
}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _REGISTER_PATH:
request_id = request.headers.get("x-request-id")
attempts.append(
{
"request_id": request_id,
"raw": raw,
"body": json.loads(raw.decode("utf-8")),
}
)
if len(attempts) == 1:
request.send_response(500)
request.end_headers()
request.wfile.write(b"transient register failure")
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_RETRY}).encode("utf-8"))
return
assert request.path == _DESCRIBE_PATH
body = json.loads(raw.decode("utf-8"))
assert body["job_id"] == _JOB_ID_RETRY
describe_calls.append(body)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"job_id": _JOB_ID_RETRY,
"job_state": "DONE",
"job_type": "register_function",
"creation_ms": 1,
"spec": {},
"result": function_result_wire,
}
).encode("utf-8")
)
package_calls = {"n": 0}
original_package = _udf_mod._package_udf
def counting_package(fn: object):
package_calls["n"] += 1
return original_package(fn)
with _mock_remote_db(handler) as db:
assert not hasattr(db, "register_function")
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
job = db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_RETRY
waited = job.wait(timeout=timedelta(seconds=5))
assert package_calls["n"] == 1
assert len(attempts) == 2
first, second = attempts
assert isinstance(first["request_id"], str) and first["request_id"]
assert first["request_id"] == second["request_id"]
assert first["raw"] == second["raw"]
assert first["raw"]
_assert_exact_register_spec(first["body"], expected_spec)
_assert_exact_register_spec(second["body"], expected_spec)
assert len(describe_calls) == 1
assert describe_calls[0]["job_id"] == _JOB_ID_RETRY
assert type(waited) is lancedb.Function
assert waited.id == _FUNCTION_ID_RETRY
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
assert waited.output_type == pa.string()
assert waited.output_nullable is True
@pytest.mark.asyncio
async def test_async_remote_register_returns_async_job_with_exact_spec():
expected_spec = _expected_register_spec(_FUNCTION_NAME, packable_register_normalize)
seen: dict[str, Any] = {}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _REGISTER_PATH
raw = _read_body(request)
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
async with _mock_remote_db_async(handler) as db:
assert not hasattr(db, "register_function")
job = await db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert seen.get("raw")
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
_assert_exact_register_spec(seen["body"], expected_spec)
assert type(job) is lancedb.job.AsyncJob
assert job.id == _JOB_ID_ASYNC
def test_sync_remote_register_http_error_omits_source_and_secret_markers():
echoed = f"register failed with {_SOURCE_MARKER} and {_SECRET_REFERENCE}"
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
assert request.path == _REGISTER_PATH
_read_body(request)
request.send_response(400)
request.end_headers()
request.wfile.write(echoed.encode("utf-8"))
with _mock_remote_db(handler) as db:
with pytest.raises(HttpError) as exc_info:
db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert received["n"] == 1
err = exc_info.value
assert isinstance(err, HttpError)
assert err.status_code == 400
_assert_markers_absent_from_exception(err)
def test_empty_name_rejects_before_http():
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
_read_body(request)
request.send_response(500)
request.end_headers()
request.wfile.write(b"should not be reached")
with _mock_remote_db(handler) as db:
with pytest.raises(ValueError):
db.functions.register("", packable_register_normalize)
assert received["n"] == 0
def test_local_sync_register_not_implemented_without_table_mutation(tmp_path):
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "register_function")
with pytest.raises(NotImplementedError):
db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert db.list_tables().tables == before
@pytest.mark.asyncio
async def test_local_async_register_not_implemented_without_table_mutation(tmp_path):
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "register_function")
with pytest.raises(NotImplementedError):
await db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert (await db.list_tables()).tables == before
@pytest.mark.parametrize("keyword", _DELETED_REGISTER_KEYWORDS)
def test_register_rejects_deleted_overdesign_keywords_before_submission(keyword):
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
_read_body(request)
request.send_response(500)
request.end_headers()
request.wfile.write(b"should not be reached")
with _mock_remote_db(handler) as db:
with pytest.raises(TypeError):
db.functions.register(
_FUNCTION_NAME,
packable_register_normalize,
**{keyword: True},
)
assert received["n"] == 0
@@ -0,0 +1,719 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python conditional first-class Function name removal."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable
import pyarrow as pa
import pytest
import lancedb
from lancedb import _lancedb as _native
from lancedb.remote.errors import HttpError
_REMOVE_PATH = "/v1/functions/remove"
_LOOKUP_PATH = "/v1/functions/lookup"
_REMOVE_CATALOG_NAME = "text.normalize.remove-name"
_REMOVE_FUNCTION_ID = "fn.exact.remove-handle"
_REMOVE_SERVER_MESSAGE_MARKER = (
"SERVER_REMOVE_DIAGNOSTIC_MARKER name=text.normalize.remove-name "
"id=fn.exact.remove-handle"
)
_SENSITIVE_BODY_MARKER = "SENSITIVE_REMOVE_BODY_MARKER"
_CONFLICTING_MESSAGE_CODE = "revoked_function"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
# replace tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_REMOVE_KEYWORDS = (
"expected_current_function_id",
"function_id",
"idempotency_key",
"retry_key",
"user_version",
"version",
"force",
"if_exists",
"revoke",
"delete",
)
def _function_error_cls(*, required: bool = True):
"""Resolve FunctionError from the live module (records RED when absent)."""
from lancedb import exceptions as exc_mod
cls = getattr(exc_mod, "FunctionError", None)
if cls is None:
if required:
pytest.fail("lancedb.exceptions.FunctionError is missing")
return type("MissingFunctionError", (), {})
return cls
def _sample_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _REMOVE_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _lookup_success_body() -> bytes:
return json.dumps({"function": _sample_function_wire()}).encode("utf-8")
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
def _close_db(db: Any) -> None:
with contextlib.suppress(Exception):
inner = getattr(db, "_conn", None)
if inner is not None:
inner.close()
return
close = getattr(db, "close", None)
if callable(close):
close()
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
db = None
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
if db is not None:
_close_db(db)
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
db = None
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
if db is not None:
_close_db(db)
server.shutdown()
thread.join()
def _exception_chain_text(exc: BaseException) -> str:
parts: list[str] = []
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(str(current))
parts.append(repr(current))
current = current.__cause__
return "\n".join(parts)
def _assert_payload_free(exc: BaseException) -> None:
text = _exception_chain_text(exc)
assert _REMOVE_SERVER_MESSAGE_MARKER not in text
assert _REMOVE_CATALOG_NAME not in text
assert _REMOVE_FUNCTION_ID not in text
assert _SENSITIVE_BODY_MARKER not in text
def _assert_exact_remove_function(function: object) -> None:
assert type(function) is lancedb.Function
assert function.id == _REMOVE_FUNCTION_ID
assert not hasattr(function, "name")
assert function.parameters == (
("text", pa.string()),
("limit", pa.int32()),
)
assert function.output_type == pa.string()
assert function.output_nullable is True
assert _REMOVE_CATALOG_NAME not in repr(function)
assert _REMOVE_CATALOG_NAME not in str(function)
def _assert_exact_remove_request(
request: http.server.BaseHTTPRequestHandler,
raw: bytes,
body: dict[str, Any],
*,
expected_id: str,
) -> None:
assert request.command == "POST"
assert request.path == _REMOVE_PATH
assert "?" not in request.path
assert raw
assert body == {
"name": _REMOVE_CATALOG_NAME,
"expected_current_function_id": expected_id,
}
assert set(body) == {"name", "expected_current_function_id"}
assert "format_version" not in body
assert "function_id" not in body
assert "function" not in body
assert "signature" not in body
assert "job_id" not in body
assert "idempotency_key" not in body
assert "user_version" not in body
assert "force" not in body
assert "if_exists" not in body
request_id = request.headers.get("x-request-id")
assert isinstance(request_id, str) and request_id
def _assert_native_remove_method_present() -> None:
assert hasattr(_native.Connection, "_remove_function_name")
assert callable(getattr(_native.Connection, "_remove_function_name"))
def _lookup_success_handler(
counters: dict[str, int],
*,
after_lookup: (
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
) = None,
):
"""Serve exact name lookup; optionally continue for remove."""
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
counters["lookup"] = counters.get("lookup", 0) + 1
body = json.loads(raw.decode("utf-8"))
assert body == {"name": _REMOVE_CATALOG_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
if after_lookup is not None:
after_lookup(request, raw)
return
counters["remove"] = counters.get("remove", 0) + 1
request.send_response(500)
request.end_headers()
request.wfile.write(b"unexpected remove")
return handler
def _observe_current(db) -> lancedb.Function:
current = db.functions.get(_REMOVE_CATALOG_NAME)
_assert_exact_remove_function(current)
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
return current
async def _observe_current_async(db) -> lancedb.Function:
current = await db.functions.get(_REMOVE_CATALOG_NAME)
_assert_exact_remove_function(current)
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
return current
def test_native_connection_exposes_private_remove_function_name():
_assert_native_remove_method_present()
def test_sync_remote_remove_exact_body_path_request_id_returns_none():
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
remove_attempts: list[dict[str, Any]] = []
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REMOVE_PATH
counters["remove"] += 1
body = json.loads(payload.decode("utf-8"))
remove_attempts.append(
{
"request": request,
"raw": payload,
"body": body,
"request_id": request.headers.get("x-request-id"),
}
)
# Illegal body on 204 must be ignored; success is status-driven only.
request.send_response(204)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
_SENSITIVE_BODY_MARKER: True,
"message": _REMOVE_SERVER_MESSAGE_MARKER,
}
).encode("utf-8")
)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
result = db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert result is None
assert counters["lookup"] == 1
assert counters["remove"] == 1
assert len(remove_attempts) == 1
attempt = remove_attempts[0]
_assert_exact_remove_request(
attempt["request"],
attempt["raw"],
attempt["body"],
expected_id=current.id,
)
assert attempt["body"]["expected_current_function_id"] == current.id
@pytest.mark.asyncio
async def test_async_remote_remove_exact_body_returns_none():
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
seen: dict[str, Any] = {}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
assert request.path == _REMOVE_PATH
counters["remove"] += 1
seen["request"] = request
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(204)
request.end_headers()
async with _mock_remote_db_async(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
current = await _observe_current_async(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
result = await db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert result is None
assert counters["lookup"] == 1
assert counters["remove"] == 1
assert seen.get("raw")
_assert_exact_remove_request(
seen["request"],
seen["raw"],
seen["body"],
expected_id=current.id,
)
def test_after_remove_name_lookup_not_found_id_lookup_same_function():
"""Catalog-pointer SDK sequence via a stateful fixture; not server atomicity."""
_assert_native_remove_method_present()
counters: dict[str, int] = {
"lookup_name": 0,
"lookup_id": 0,
"remove": 0,
}
removed = {"yes": False}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
body = json.loads(raw.decode("utf-8"))
if "name" in body:
counters["lookup_name"] += 1
assert body == {"name": _REMOVE_CATALOG_NAME}
if removed["yes"]:
request.send_response(404)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"error_code": "name_or_function_not_found",
"message": _REMOVE_SERVER_MESSAGE_MARKER,
}
).encode("utf-8")
)
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
counters["lookup_id"] += 1
assert body == {"function_id": _REMOVE_FUNCTION_ID}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
assert request.path == _REMOVE_PATH
counters["remove"] += 1
body = json.loads(raw.decode("utf-8"))
_assert_exact_remove_request(
request, raw, body, expected_id=_REMOVE_FUNCTION_ID
)
removed["yes"] = True
request.send_response(204)
request.end_headers()
function_error = _function_error_cls()
with _mock_remote_db(handler) as db:
current = _observe_current(db)
assert counters["lookup_name"] == 1
assert counters["lookup_id"] == 0
assert counters["remove"] == 0
result = db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert result is None
assert counters["lookup_name"] == 1
assert counters["remove"] == 1
with pytest.raises(function_error) as exc_info:
db.functions.get(_REMOVE_CATALOG_NAME)
err = exc_info.value
assert err.code == "name_or_function_not_found"
_assert_payload_free(err)
by_id = db.functions.get_by_id(_REMOVE_FUNCTION_ID)
assert counters["lookup_name"] == 2
assert counters["lookup_id"] == 1
assert counters["remove"] == 1
_assert_exact_remove_function(by_id)
assert by_id.id == current.id
assert by_id.parameters == current.parameters
assert by_id.output_type == current.output_type
assert by_id.output_nullable is current.output_nullable
def test_explicit_name_conflict_is_function_error_payload_free():
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
body = {
"error_code": "name_conflict",
"message": (
f"{_REMOVE_SERVER_MESSAGE_MARKER} looks_like {_CONFLICTING_MESSAGE_CODE}"
),
"name": _REMOVE_CATALOG_NAME,
"function_id": _REMOVE_FUNCTION_ID,
_SENSITIVE_BODY_MARKER: True,
}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REMOVE_PATH
counters["remove"] += 1
parsed = json.loads(payload.decode("utf-8"))
_assert_exact_remove_request(
request, payload, parsed, expected_id=_REMOVE_FUNCTION_ID
)
request.send_response(409)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(function_error) as exc_info:
db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert counters["lookup"] == 1
assert counters["remove"] == 1
err = exc_info.value
assert isinstance(err, function_error)
assert err.code == "name_conflict"
assert err.code != _CONFLICTING_MESSAGE_CODE
_assert_payload_free(err)
@pytest.mark.parametrize(
"label,status,response_body",
[
(
"200_with_body",
200,
{
"ok": True,
"message": _REMOVE_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
"job_id": "must-not-infer-job",
},
),
(
"202_empty",
202,
f"{_REMOVE_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
),
("200_empty", 200, ""),
],
)
def test_http_200_202_cannot_return_success(
label: str, status: int, response_body: object
):
del label
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REMOVE_PATH
counters["remove"] += 1
parsed = json.loads(payload.decode("utf-8"))
_assert_exact_remove_request(
request, payload, parsed, expected_id=_REMOVE_FUNCTION_ID
)
request.send_response(status)
request.send_header("Content-Type", "application/json")
request.end_headers()
if isinstance(response_body, str):
request.wfile.write(response_body.encode("utf-8"))
else:
request.wfile.write(json.dumps(response_body).encode("utf-8"))
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(HttpError) as exc_info:
db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert counters["lookup"] == 1
assert counters["remove"] == 1
err = exc_info.value
assert isinstance(err, HttpError)
assert not isinstance(err, _function_error_cls(required=False))
_assert_payload_free(err)
def test_empty_name_rejects_before_remove_transport():
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(ValueError):
db.functions.remove("", current)
assert counters["lookup"] == 1
assert counters["remove"] == 0
@pytest.mark.parametrize(
"bad_current",
[
_REMOVE_FUNCTION_ID,
{"id": _REMOVE_FUNCTION_ID},
object(),
123,
],
)
def test_raw_id_or_arbitrary_current_rejected_without_remove(bad_current):
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
# Observe a real handle separately so the bad-current path is isolated.
_ = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(TypeError):
db.functions.remove(_REMOVE_CATALOG_NAME, bad_current)
assert counters["lookup"] == 1
assert counters["remove"] == 0
def test_local_sync_remove_not_implemented_without_table_mutation(tmp_path):
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
assert type(current) is lancedb.Function
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
with pytest.raises(NotImplementedError):
db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert db.list_tables().tables == before
_close_db(db)
@pytest.mark.asyncio
async def test_local_async_remove_not_implemented_without_table_mutation(tmp_path):
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
assert type(current) is lancedb.Function
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
with pytest.raises(NotImplementedError):
await db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert (await db.list_tables()).tables == before
db.close()
@pytest.mark.parametrize("keyword", _DELETED_REMOVE_KEYWORDS)
def test_remove_rejects_deleted_cas_retry_version_kwargs(keyword):
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(TypeError):
db.functions.remove(
_REMOVE_CATALOG_NAME,
current,
**{keyword: True},
)
assert counters["lookup"] == 1
assert counters["remove"] == 0
def test_no_direct_remove_methods_and_function_has_no_remove_facade_private():
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
assert callable(getattr(db.functions, "remove", None))
assert not hasattr(lancedb, "_SyncFunctions")
assert not hasattr(lancedb, "_AsyncFunctions")
assert type(db.functions).__name__.startswith("_")
assert "_SyncFunctions" not in getattr(lancedb, "__all__", [])
assert "_AsyncFunctions" not in getattr(lancedb, "__all__", [])
assert counters["lookup"] == 1
assert counters["remove"] == 0
@@ -0,0 +1,579 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python conditional first-class Function replacement."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from datetime import timedelta
from typing import Any, Callable
from unittest import mock
import pyarrow as pa
import pytest
import lancedb
import lancedb._udf as _udf_mod
import lancedb.job
from lancedb import FunctionCapability, udf
from lancedb.exceptions import JobFailedError
_SOURCE_MARKER = "replace-source-marker-unique-xyz"
_SECRET_REFERENCE = "secret://team/replace-redact-token-xyz"
_SECRET_ENV = "REPLACE_API_TOKEN"
_NETWORK_ORIGIN = "https://api.replace-example.com"
_FUNCTION_NAME = "text.normalize"
_CURRENT_FUNCTION_ID = "fn.replace-current-1"
_REPLACED_FUNCTION_ID = "fn.replace-result-1"
_JOB_ID_SYNC = "job-replace-sync-1"
_JOB_ID_ASYNC = "job-replace-async-1"
_JOB_ID_CONFLICT = "job-replace-conflict-1"
_REGISTER_PATH = "/v1/functions/register"
_LOOKUP_PATH = "/v1/functions/lookup"
_DESCRIBE_PATH = "/v1/jobs/describe"
_CONFLICTING_MESSAGE_CODE = "definition_validation_failure"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
# job-result tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_REPLACE_KEYWORDS = (
"idempotency_key",
"retry_key",
"user_version",
"version",
"deterministic",
"null_policy",
"replace",
"expected_current_function_id",
"alias",
"lineage",
)
_SPEC_KEYS = {
"format_version",
"name",
"definition",
"expected_current_function_id",
}
@udf(
inputs={"text": pa.string(), "limit": pa.int32()},
output=pa.string(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=True,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
],
)
def packable_replace_normalize(text, limit):
"""replace-source-marker-unique-xyz."""
return text[:limit]
def _current_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _CURRENT_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _definition_json(fn: object) -> dict[str, Any]:
payload = _udf_mod._build_function_definition(fn)._to_json()
if isinstance(payload, bytes):
return json.loads(payload.decode("utf-8"))
assert isinstance(payload, str)
return json.loads(payload)
def _expected_replace_spec(name: str, current_id: str, fn: object) -> dict[str, Any]:
return {
"format_version": 1,
"name": name,
"definition": _definition_json(fn),
"expected_current_function_id": current_id,
}
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _assert_exact_replace_spec(
body: dict[str, Any], expected: dict[str, Any], current_id: str
) -> None:
assert set(body) == _SPEC_KEYS
assert body == expected
assert body["format_version"] == 1
assert body["expected_current_function_id"] == current_id
assert body["expected_current_function_id"] is not None
assert _SOURCE_MARKER in json.dumps(body["definition"])
assert any(
capability.get("reference") == _SECRET_REFERENCE
for capability in body["definition"]["capabilities"]
)
def _lookup_success_handler(
counters: dict[str, int],
*,
after_lookup: (
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
) = None,
):
"""Serve exact lookup; optionally continue for register/describe."""
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
counters["lookup"] = counters.get("lookup", 0) + 1
body = json.loads(raw.decode("utf-8"))
assert body == {"name": _FUNCTION_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps({"function": _current_function_wire()}).encode("utf-8")
)
return
if after_lookup is not None:
after_lookup(request, raw)
return
counters["register"] = counters.get("register", 0) + 1
request.send_response(500)
request.end_headers()
request.wfile.write(b"unexpected register")
return handler
def _observe_current(db) -> lancedb.Function:
current = db.functions.get(_FUNCTION_NAME)
assert type(current) is lancedb.Function
assert current.id == _CURRENT_FUNCTION_ID
assert not hasattr(current, "name")
assert not hasattr(current, "replace")
return current
async def _observe_current_async(db) -> lancedb.Function:
current = await db.functions.get(_FUNCTION_NAME)
assert type(current) is lancedb.Function
assert current.id == _CURRENT_FUNCTION_ID
assert not hasattr(current, "name")
assert not hasattr(current, "replace")
return current
def test_sync_remote_replace_exact_body_one_package_job_and_function_result():
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
expected_spec = _expected_replace_spec(
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
)
register_attempts: list[dict[str, Any]] = []
function_result_wire = {
"kind": "function",
"format_version": 1,
"function": {
"format_version": 1,
"id": _REPLACED_FUNCTION_ID,
"signature": expected_spec["definition"]["signature"],
},
}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
if request.path == _REGISTER_PATH:
counters["register"] += 1
register_attempts.append(
{
"request_id": request.headers.get("x-request-id"),
"raw": payload,
"body": json.loads(payload.decode("utf-8")),
}
)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_SYNC}).encode("utf-8"))
return
assert request.path == _DESCRIBE_PATH
counters["describe"] += 1
body = json.loads(payload.decode("utf-8"))
assert body["job_id"] == _JOB_ID_SYNC
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"job_id": _JOB_ID_SYNC,
"job_state": "DONE",
"job_type": "register_function",
"creation_ms": 1,
"spec": {},
"result": function_result_wire,
}
).encode("utf-8")
)
package_calls = {"n": 0}
original_package = _udf_mod._package_udf
def counting_package(fn: object):
package_calls["n"] += 1
return original_package(fn)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "replace_function")
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
job = db.functions.replace(
_FUNCTION_NAME, current, packable_replace_normalize
)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_SYNC
waited = job.wait(timeout=timedelta(seconds=5))
assert package_calls["n"] == 1
assert counters["lookup"] == 1
assert counters["register"] == 1
assert counters["describe"] == 1
assert len(register_attempts) == 1
attempt = register_attempts[0]
assert isinstance(attempt["request_id"], str) and attempt["request_id"]
assert attempt["raw"]
_assert_exact_replace_spec(attempt["body"], expected_spec, current.id)
assert type(waited) is lancedb.Function
assert waited.id == _REPLACED_FUNCTION_ID
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
assert waited.output_type == pa.string()
assert waited.output_nullable is True
@pytest.mark.asyncio
async def test_async_remote_replace_exact_body_returns_async_job():
counters: dict[str, int] = {"lookup": 0, "register": 0}
expected_spec = _expected_replace_spec(
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
)
seen: dict[str, Any] = {}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
assert request.path == _REGISTER_PATH
counters["register"] += 1
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
async with _mock_remote_db_async(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "replace_function")
current = await _observe_current_async(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
job = await db.functions.replace(
_FUNCTION_NAME, current, packable_replace_normalize
)
assert counters["lookup"] == 1
assert counters["register"] == 1
assert seen.get("raw")
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
_assert_exact_replace_spec(seen["body"], expected_spec, current.id)
assert type(job) is lancedb.job.AsyncJob
assert job.id == _JOB_ID_ASYNC
def test_sync_remote_replace_failed_name_conflict_raises_job_failed_error_code():
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
if request.path == _REGISTER_PATH:
counters["register"] += 1
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps({"job_id": _JOB_ID_CONFLICT}).encode("utf-8")
)
return
assert request.path == _DESCRIBE_PATH
counters["describe"] += 1
body = json.loads(raw.decode("utf-8"))
assert body["job_id"] == _JOB_ID_CONFLICT
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"job_id": _JOB_ID_CONFLICT,
"job_type": "register_function",
"job_state": "FAILED",
"creation_ms": 1,
"spec": {},
"failure": {
"phase": "validate",
"message": (
f"looks like {_CONFLICTING_MESSAGE_CODE} during CAS"
),
"retryable": False,
"error_code": "name_conflict",
},
}
).encode("utf-8")
)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
job = db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
assert type(job) is lancedb.job.Job
with pytest.raises(JobFailedError) as exc_info:
job.wait(timeout=timedelta(seconds=5))
assert counters["lookup"] == 1
assert counters["register"] == 1
assert counters["describe"] == 1
err = exc_info.value
assert isinstance(err, JobFailedError)
assert err.error_code == "name_conflict"
assert err.error_code != _CONFLICTING_MESSAGE_CODE
def test_empty_name_rejects_before_register_transport():
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with pytest.raises(ValueError):
db.functions.replace("", current, packable_replace_normalize)
assert counters["lookup"] == 1
assert counters["register"] == 0
@pytest.mark.parametrize(
"bad_current",
[
_CURRENT_FUNCTION_ID,
{"id": _CURRENT_FUNCTION_ID},
object(),
123,
],
)
def test_raw_id_or_arbitrary_current_rejected_without_register(bad_current):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
# Observe a real handle separately so the bad-current path is isolated.
_ = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with pytest.raises(TypeError):
db.functions.replace(
_FUNCTION_NAME, bad_current, packable_replace_normalize
)
assert counters["lookup"] == 1
assert counters["register"] == 0
def test_local_sync_replace_not_implemented_without_table_mutation(tmp_path):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["register"] == 0
assert type(current) is lancedb.Function
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "replace_function")
with pytest.raises(NotImplementedError):
db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
assert db.list_tables().tables == before
@pytest.mark.asyncio
async def test_local_async_replace_not_implemented_without_table_mutation(tmp_path):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["register"] == 0
assert type(current) is lancedb.Function
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "replace_function")
with pytest.raises(NotImplementedError):
await db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
assert (await db.list_tables()).tables == before
@pytest.mark.parametrize("keyword", _DELETED_REPLACE_KEYWORDS)
def test_replace_rejects_deleted_cas_retry_version_kwargs(keyword):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with pytest.raises(TypeError):
db.functions.replace(
_FUNCTION_NAME,
current,
packable_replace_normalize,
**{keyword: True},
)
assert counters["lookup"] == 1
assert counters["register"] == 0
def test_no_direct_replace_function_methods_and_function_has_no_replace():
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert not hasattr(db, "replace_function")
assert not hasattr(db, "register_function")
assert not hasattr(current, "replace")
assert not hasattr(current, "replace_function")
assert callable(getattr(db.functions, "replace", None))
assert counters["lookup"] == 1
assert counters["register"] == 0
@@ -0,0 +1,728 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python exact first-class Function revocation."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable
import pyarrow as pa
import pytest
import lancedb
from lancedb import _lancedb as _native
from lancedb.remote.errors import HttpError
_REVOKE_PATH = "/v1/functions/revoke"
_LOOKUP_PATH = "/v1/functions/lookup"
_REVOKE_CATALOG_NAME = "text.normalize.revoke-name"
_REVOKE_FUNCTION_ID = "fn.exact.revoke-handle"
_REVOKE_SERVER_MESSAGE_MARKER = (
"SERVER_REVOKE_DIAGNOSTIC_MARKER id=fn.exact.revoke-handle "
"name=text.normalize.revoke-name"
)
_SENSITIVE_BODY_MARKER = "SENSITIVE_REVOKE_BODY_MARKER"
_CONFLICTING_MESSAGE_CODE = "revoked_function"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
# remove tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_REVOKE_KEYWORDS = (
"function_id",
"name",
"idempotency_key",
"retry_key",
"user_version",
"version",
"reason",
"expiry",
"force",
"if_exists",
"remove",
"delete",
)
def _function_error_cls(*, required: bool = True):
"""Resolve FunctionError from the live module (records RED when absent)."""
from lancedb import exceptions as exc_mod
cls = getattr(exc_mod, "FunctionError", None)
if cls is None:
if required:
pytest.fail("lancedb.exceptions.FunctionError is missing")
return type("MissingFunctionError", (), {})
return cls
def _sample_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _REVOKE_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _lookup_success_body() -> bytes:
return json.dumps({"function": _sample_function_wire()}).encode("utf-8")
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
def _close_db(db: Any) -> None:
with contextlib.suppress(Exception):
inner = getattr(db, "_conn", None)
if inner is not None:
inner.close()
return
close = getattr(db, "close", None)
if callable(close):
close()
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
db = None
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
if db is not None:
_close_db(db)
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
db = None
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
if db is not None:
_close_db(db)
server.shutdown()
thread.join()
def _exception_chain_text(exc: BaseException) -> str:
parts: list[str] = []
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(str(current))
parts.append(repr(current))
current = current.__cause__
return "\n".join(parts)
def _assert_payload_free(exc: BaseException) -> None:
text = _exception_chain_text(exc)
assert _REVOKE_SERVER_MESSAGE_MARKER not in text
assert _REVOKE_CATALOG_NAME not in text
assert _REVOKE_FUNCTION_ID not in text
assert _SENSITIVE_BODY_MARKER not in text
def _assert_exact_revoke_function(function: object) -> None:
assert type(function) is lancedb.Function
assert function.id == _REVOKE_FUNCTION_ID
assert not hasattr(function, "name")
assert function.parameters == (
("text", pa.string()),
("limit", pa.int32()),
)
assert function.output_type == pa.string()
assert function.output_nullable is True
assert _REVOKE_CATALOG_NAME not in repr(function)
assert _REVOKE_CATALOG_NAME not in str(function)
def _assert_exact_revoke_request(
request: http.server.BaseHTTPRequestHandler,
raw: bytes,
body: dict[str, Any],
*,
expected_id: str,
) -> None:
assert request.command == "POST"
assert request.path == _REVOKE_PATH
assert "?" not in request.path
assert "remove" not in request.path
assert raw
assert body == {"function_id": expected_id}
assert set(body) == {"function_id"}
assert "name" not in body
assert "expected_current_function_id" not in body
assert "format_version" not in body
assert "function" not in body
assert "signature" not in body
assert "job_id" not in body
assert "idempotency_key" not in body
assert "user_version" not in body
assert "reason" not in body
assert "expiry" not in body
assert "force" not in body
request_id = request.headers.get("x-request-id")
assert isinstance(request_id, str) and request_id
def _assert_native_revoke_method_present() -> None:
assert hasattr(_native.Connection, "_revoke_function")
assert callable(getattr(_native.Connection, "_revoke_function"))
def _lookup_success_handler(
counters: dict[str, int],
*,
after_lookup: (
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
) = None,
):
"""Serve exact name lookup; optionally continue for revoke."""
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
counters["lookup"] = counters.get("lookup", 0) + 1
body = json.loads(raw.decode("utf-8"))
assert body == {"name": _REVOKE_CATALOG_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
if after_lookup is not None:
after_lookup(request, raw)
return
counters["revoke"] = counters.get("revoke", 0) + 1
request.send_response(500)
request.end_headers()
request.wfile.write(b"unexpected revoke")
return handler
def _observe_current(db) -> lancedb.Function:
current = db.functions.get(_REVOKE_CATALOG_NAME)
_assert_exact_revoke_function(current)
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
return current
async def _observe_current_async(db) -> lancedb.Function:
current = await db.functions.get(_REVOKE_CATALOG_NAME)
_assert_exact_revoke_function(current)
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
return current
def test_native_connection_exposes_private_revoke_function():
_assert_native_revoke_method_present()
def test_sync_remote_revoke_exact_body_path_request_id_returns_none():
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
revoke_attempts: list[dict[str, Any]] = []
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
body = json.loads(payload.decode("utf-8"))
revoke_attempts.append(
{
"request": request,
"raw": payload,
"body": body,
"request_id": request.headers.get("x-request-id"),
}
)
# Illegal body on 204 must be ignored; success is status-driven only.
request.send_response(204)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
_SENSITIVE_BODY_MARKER: True,
"message": _REVOKE_SERVER_MESSAGE_MARKER,
}
).encode("utf-8")
)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "revoke_function")
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
result = db.functions.revoke(current)
assert result is None
assert counters["lookup"] == 1
assert counters["revoke"] == 1
assert len(revoke_attempts) == 1
attempt = revoke_attempts[0]
_assert_exact_revoke_request(
attempt["request"],
attempt["raw"],
attempt["body"],
expected_id=current.id,
)
assert attempt["body"]["function_id"] == current.id
_assert_exact_revoke_function(current)
@pytest.mark.asyncio
async def test_async_remote_revoke_exact_body_returns_none():
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
seen: dict[str, Any] = {}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
seen["request"] = request
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(204)
request.end_headers()
async with _mock_remote_db_async(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "revoke_function")
current = await _observe_current_async(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
result = await db.functions.revoke(current)
assert result is None
assert counters["lookup"] == 1
assert counters["revoke"] == 1
assert seen.get("raw")
_assert_exact_revoke_request(
seen["request"],
seen["raw"],
seen["body"],
expected_id=current.id,
)
def test_repeated_remote_revoke_204_both_return_none():
"""Two logical calls each receiving 204 both succeed (Python outcome only)."""
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
revoke_request_ids: list[str] = []
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
body = json.loads(payload.decode("utf-8"))
_assert_exact_revoke_request(
request, payload, body, expected_id=_REVOKE_FUNCTION_ID
)
request_id = request.headers.get("x-request-id")
assert isinstance(request_id, str) and request_id
revoke_request_ids.append(request_id)
request.send_response(204)
request.end_headers()
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
first = db.functions.revoke(current)
second = db.functions.revoke(current)
assert first is None
assert second is None
assert counters["lookup"] == 1
assert counters["revoke"] == 2
assert len(revoke_request_ids) == 2
_assert_exact_revoke_function(current)
def test_after_revoke_name_and_id_lookup_still_return_function():
"""Revoke does not unlink names; SDK-visible sequence only, not Sophon proof."""
_assert_native_revoke_method_present()
counters: dict[str, int] = {
"lookup_name": 0,
"lookup_id": 0,
"revoke": 0,
}
revoked = {"yes": False}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
body = json.loads(raw.decode("utf-8"))
if "name" in body:
counters["lookup_name"] += 1
assert body == {"name": _REVOKE_CATALOG_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
counters["lookup_id"] += 1
assert body == {"function_id": _REVOKE_FUNCTION_ID}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
body = json.loads(raw.decode("utf-8"))
_assert_exact_revoke_request(
request, raw, body, expected_id=_REVOKE_FUNCTION_ID
)
revoked["yes"] = True
request.send_response(204)
request.end_headers()
with _mock_remote_db(handler) as db:
current = _observe_current(db)
assert counters["lookup_name"] == 1
assert counters["lookup_id"] == 0
assert counters["revoke"] == 0
assert not revoked["yes"]
result = db.functions.revoke(current)
assert result is None
assert counters["lookup_name"] == 1
assert counters["revoke"] == 1
assert revoked["yes"]
by_name = db.functions.get(_REVOKE_CATALOG_NAME)
by_id = db.functions.get_by_id(_REVOKE_FUNCTION_ID)
assert counters["lookup_name"] == 2
assert counters["lookup_id"] == 1
assert counters["revoke"] == 1
_assert_exact_revoke_function(by_name)
_assert_exact_revoke_function(by_id)
assert by_name.id == current.id
assert by_id.id == current.id
assert by_name.parameters == current.parameters
assert by_id.parameters == current.parameters
assert by_name.output_type == current.output_type
assert by_id.output_type == current.output_type
assert by_name.output_nullable is current.output_nullable
assert by_id.output_nullable is current.output_nullable
def test_explicit_name_or_function_not_found_is_function_error_payload_free():
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
body = {
"error_code": "name_or_function_not_found",
"message": (
f"{_REVOKE_SERVER_MESSAGE_MARKER} looks_like {_CONFLICTING_MESSAGE_CODE} "
"name_conflict"
),
"name": _REVOKE_CATALOG_NAME,
"function_id": _REVOKE_FUNCTION_ID,
_SENSITIVE_BODY_MARKER: True,
}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
parsed = json.loads(payload.decode("utf-8"))
_assert_exact_revoke_request(
request, payload, parsed, expected_id=_REVOKE_FUNCTION_ID
)
request.send_response(404)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
with pytest.raises(function_error) as exc_info:
db.functions.revoke(current)
assert counters["lookup"] == 1
assert counters["revoke"] == 1
err = exc_info.value
assert isinstance(err, function_error)
assert err.code == "name_or_function_not_found"
assert err.code != _CONFLICTING_MESSAGE_CODE
assert err.code != "name_conflict"
_assert_payload_free(err)
@pytest.mark.parametrize(
"label,status,response_body",
[
(
"200_with_body",
200,
{
"ok": True,
"message": _REVOKE_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
"job_id": "must-not-infer-job",
},
),
(
"202_empty",
202,
f"{_REVOKE_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
),
("200_empty", 200, ""),
],
)
def test_http_200_202_cannot_return_success(
label: str, status: int, response_body: object
):
del label
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
parsed = json.loads(payload.decode("utf-8"))
_assert_exact_revoke_request(
request, payload, parsed, expected_id=_REVOKE_FUNCTION_ID
)
request.send_response(status)
request.send_header("Content-Type", "application/json")
request.end_headers()
if isinstance(response_body, str):
request.wfile.write(response_body.encode("utf-8"))
else:
request.wfile.write(json.dumps(response_body).encode("utf-8"))
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
with pytest.raises(HttpError) as exc_info:
db.functions.revoke(current)
assert counters["lookup"] == 1
assert counters["revoke"] == 1
err = exc_info.value
assert isinstance(err, HttpError)
assert not isinstance(err, _function_error_cls(required=False))
_assert_payload_free(err)
@pytest.mark.parametrize(
"bad_function",
[
_REVOKE_FUNCTION_ID,
{"id": _REVOKE_FUNCTION_ID},
object(),
123,
],
)
def test_raw_id_or_arbitrary_function_rejected_without_revoke(bad_function):
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
# Observe a real handle separately so the bad-function path is isolated.
_ = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
with pytest.raises(TypeError):
db.functions.revoke(bad_function)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
def test_local_sync_revoke_not_implemented_without_table_mutation(tmp_path):
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
assert type(current) is lancedb.Function
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "revoke_function")
with pytest.raises(NotImplementedError):
db.functions.revoke(current)
assert db.list_tables().tables == before
_close_db(db)
@pytest.mark.asyncio
async def test_local_async_revoke_not_implemented_without_table_mutation(tmp_path):
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
assert type(current) is lancedb.Function
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "revoke_function")
with pytest.raises(NotImplementedError):
await db.functions.revoke(current)
assert (await db.list_tables()).tables == before
db.close()
@pytest.mark.parametrize("keyword", _DELETED_REVOKE_KEYWORDS)
def test_revoke_rejects_overdesigned_kwargs(keyword):
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
with pytest.raises(TypeError):
db.functions.revoke(current, **{keyword: True})
assert counters["lookup"] == 1
assert counters["revoke"] == 0
def test_no_direct_revoke_methods_and_function_has_no_revoke_facade_private():
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert not hasattr(db, "revoke_function")
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
assert callable(getattr(db.functions, "revoke", None))
assert not hasattr(lancedb, "_SyncFunctions")
assert not hasattr(lancedb, "_AsyncFunctions")
assert type(db.functions).__name__.startswith("_")
assert "_SyncFunctions" not in getattr(lancedb, "__all__", [])
assert "_AsyncFunctions" not in getattr(lancedb, "__all__", [])
assert counters["lookup"] == 1
assert counters["revoke"] == 0
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,899 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python ``table.add_generated_column`` (FF-032).
Public user shape under test:
job = table.add_generated_column(
"normalized_text",
normalize(text=col("text")),
)
job.wait()
These tests exercise the live worktree PyO3 extension and public sync/async
wrappers. While the public methods and hidden native bridge are absent they
fail against that extension; once present they freeze the public contract
below. They must not fake success paths.
"""
from __future__ import annotations
import contextlib
import http.server
import inspect
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable
import pytest
import lancedb
import lancedb.job
from lancedb import _lancedb as _native
from lancedb.expr import col
from lancedb.remote.table import RemoteTable
from lancedb.table import AsyncTable, LanceTable, Table
_LOOKUP_PATH = "/v1/functions/lookup"
_JOB_DESCRIBE_PATH = "/v1/jobs/describe"
_TABLE_NAME = "articles"
_DESCRIBE_PATH = f"/v1/table/{_TABLE_NAME}/describe/"
_CREATE_PATH = f"/v1/table/{_TABLE_NAME}/generated_columns/create/"
_BRANCHES_CREATE_PATH = f"/v1/table/{_TABLE_NAME}/branches/create/"
_BRANCHES_LIST_PATH = f"/v1/table/{_TABLE_NAME}/branches/list/"
_CATALOG_NAME = "text.normalize"
_FUNCTION_ID = "fn.exact.normalize.gen-col"
_JOB_ID_SYNC = "job-create-gen-col-sync-1"
_JOB_ID_ASYNC = "job-create-gen-col-async-1"
_JOB_ID_BRANCH = "job-create-gen-col-branch-1"
_SOURCE_TABLE_VERSION = 42
_TEXT_FIELD_ID = 7
_BRANCH_NAME = "exp"
_BRANCH_SOURCE_VERSION = 9
_BRANCH_TEXT_FIELD_ID = 11
_DESCRIBE_BODY_MARKER = "SENSITIVE_DESCRIBE_BODY_MARKER_gen_col_xyz"
_CREATE_RESPONSE_MARKER = "SENSITIVE_CREATE_RESPONSE_MARKER_gen_col_xyz"
_LITERAL_PAYLOAD_SENTINEL = "LITERAL_PAYLOAD_SENTINEL_gen_col_xyz"
# Pinned Rust-canonical schema-only Utf8 type IPC (base64), shared with FF-028.
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_FORBIDDEN_PUBLIC_NAMES = (
"FunctionCall",
"BoundFunctionCall",
"AuthoredFunctionCall",
"CreateGeneratedColumnRequest",
"CreateGeneratedColumnJobSpec",
"GeneratedColumnBindingSnapshot",
"GeneratedColumnCreateRequest",
"geneva",
"GenevaFunction",
"VirtualColumnDefinition",
)
_FORBIDDEN_METHOD_KWARGS = (
"source_table_version",
"version",
"field_id",
"field_ids",
"output",
"output_type",
"output_nullable",
"nullable",
"spec",
"retry_key",
"idempotency_key",
"request",
"envelope",
"table_ref",
"branch",
)
def _sample_function_wire(
*,
function_id: str = _FUNCTION_ID,
parameters: list[dict[str, str]] | None = None,
) -> dict[str, Any]:
return {
"format_version": 1,
"id": function_id,
"signature": {
"parameters": parameters
or [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _text_schema_fields(
*, arrow_type: str = "string", nullable: bool = True
) -> dict[str, Any]:
return {
"fields": [
{
"name": "text",
"type": {"type": arrow_type},
"nullable": nullable,
}
]
}
def _describe_body(
*,
version: int = _SOURCE_TABLE_VERSION,
field_ids: list[int] | None = None,
arrow_type: str = "string",
include_marker: bool = True,
) -> dict[str, Any]:
body: dict[str, Any] = {
"version": version,
"schema": _text_schema_fields(arrow_type=arrow_type),
"field_ids": field_ids if field_ids is not None else [_TEXT_FIELD_ID],
}
if include_marker:
body["server_diagnostic"] = _DESCRIBE_BODY_MARKER
return body
def _create_gen_column_done_body(job_id: str) -> dict[str, Any]:
# DONE with omitted result: create_gen_column projects JobResult::None.
return {
"job_id": job_id,
"job_state": "DONE",
"job_type": "create_gen_column",
"creation_ms": 1,
"spec": {},
}
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _exception_text(exc: BaseException) -> str:
parts = [str(exc), repr(exc)]
current: BaseException | None = exc
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(f"{type(current).__name__}: {current}")
current = current.__cause__ or current.__context__
return "\n".join(parts)
def _json_response(
request: http.server.BaseHTTPRequestHandler, body: dict[str, Any]
) -> None:
payload = json.dumps(body).encode("utf-8")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(payload)
def _lookup_function(db: Any) -> lancedb.Function:
return db.functions.get(_CATALOG_NAME)
class _RequestLog:
"""Track lookup/describe/create after setup; setup traffic is excluded."""
def __init__(self) -> None:
self.lookup: list[dict[str, Any]] = []
self.describe: list[dict[str, Any]] = []
self.create: list[dict[str, Any]] = []
self.other_table: list[str] = []
self.recording = False
def start(self) -> None:
# Drop setup's explicit Function lookup and open_table describe so
# operation accounting cannot be polluted by fixture traffic.
self.lookup.clear()
self.describe.clear()
self.create.clear()
self.other_table.clear()
self.recording = True
def note(self, path: str, body: dict[str, Any] | None = None) -> None:
if not self.recording:
return
if path == _LOOKUP_PATH:
self.lookup.append(body or {})
elif path == _DESCRIBE_PATH:
self.describe.append(body or {})
elif path == _CREATE_PATH:
self.create.append(body or {})
elif path.startswith(f"/v1/table/{_TABLE_NAME}/"):
self.other_table.append(path)
def _assert_no_operation_traffic(log: _RequestLog) -> None:
assert log.lookup == []
assert log.describe == []
assert log.create == []
assert log.other_table == []
def _assert_exact_public_signature(method: Any) -> None:
"""Freeze ``(self, column_name, call)`` with no varargs/kwargs escape hatches."""
params = list(inspect.signature(method).parameters.values())
assert [p.name for p in params] == ["self", "column_name", "call"]
for param in params:
assert param.kind in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
)
assert param.default is inspect.Parameter.empty
assert param.kind is not inspect.Parameter.VAR_POSITIONAL
assert param.kind is not inspect.Parameter.VAR_KEYWORD
assert param.kind is not inspect.Parameter.KEYWORD_ONLY
def _open_table_and_function(
*,
describe_body: dict[str, Any] | None = None,
on_create: Callable[[dict[str, Any], http.server.BaseHTTPRequestHandler], None]
| None = None,
job_id: str = _JOB_ID_SYNC,
support_branch_create: bool = False,
function_wire: dict[str, Any] | None = None,
):
"""Open remote table + immutable Function; return (db, table, function, log, cm)."""
log = _RequestLog()
binding_describe = describe_body or _describe_body()
open_describe = {
"version": 1,
"schema": _text_schema_fields(),
}
state = {"opened": False}
wire = function_wire or _sample_function_wire()
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _LOOKUP_PATH:
log.note(request.path, body)
_json_response(request, {"function": wire})
return
if request.path == _JOB_DESCRIBE_PATH:
assert body["job_id"] == job_id
_json_response(request, _create_gen_column_done_body(job_id))
return
if support_branch_create and request.path == _BRANCHES_CREATE_PATH:
log.note(request.path, body)
_json_response(request, {})
return
if support_branch_create and request.path == _BRANCHES_LIST_PATH:
log.note(request.path, body)
_json_response(
request,
{
"branches": {
_BRANCH_NAME: {
"parentBranch": None,
"parentVersion": 1,
"createAt": 1,
"manifestSize": 1,
}
}
},
)
return
if request.path == _DESCRIBE_PATH:
# First describe seeds open_table; later ones are binding snapshots.
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding_describe)
return
if request.path == _CREATE_PATH:
log.note(request.path, body)
if on_create is not None:
on_create(body, request)
return
_json_response(
request,
{
"job_id": job_id,
"server_extra": {"marker": _CREATE_RESPONSE_MARKER},
},
)
return
if request.path.startswith(f"/v1/table/{_TABLE_NAME}/"):
log.note(request.path, body)
request.send_response(404)
request.end_headers()
request.wfile.write(b"unexpected path")
cm = _mock_remote_db(handler)
db = cm.__enter__()
function = _lookup_function(db)
table = db.open_table(_TABLE_NAME)
assert isinstance(table, RemoteTable)
# open_table consumed the seed describe; binding/create accounting starts now.
# Setup's one explicit lookup is cleared here and must not pollute counts.
log.start()
return db, table, function, log, cm
def _assert_exact_create_envelope(
body: dict[str, Any],
*,
source_table_version: int,
column_name: str,
field_id: int,
branch: str | None = None,
) -> None:
expected_keys = {"source_table_version", "spec"}
if branch is not None:
expected_keys.add("branch")
assert set(body) == expected_keys
assert body["source_table_version"] == source_table_version
assert "table_ref" not in body
if branch is None:
assert "branch" not in body
else:
assert body["branch"] == branch
spec = body["spec"]
assert set(spec) == {"format_version", "column_name", "function_call"}
assert spec["format_version"] == 1
assert spec["column_name"] == column_name
for forbidden in (
"table_ref",
"source_table_version",
"version",
"output",
"output_type",
"output_field_id",
"dependency_epoch",
"materialized_epoch",
"idempotency_key",
"retry_key",
"name",
"handle",
"artifact",
"geneva",
):
assert forbidden not in spec
call = spec["function_call"]
assert set(call) == {"function_id", "arguments"}
assert call["function_id"] == _FUNCTION_ID
assert len(call["arguments"]) == 1
binding = call["arguments"][0]
assert binding["parameter"] == "text"
value = binding["value"]
assert value["kind"] == "field"
assert value["field_id"] == field_id
assert value["data_type_ipc"] == _UTF8_TYPE_IPC_B64
assert "name" not in value
assert "column_name" not in value
assert "text" not in value
# Serialized call must not late-bind by column name anywhere relevant.
dumped = json.dumps(call)
assert '"column_name"' not in dumped
assert "normalized_text" not in dumped
def test_public_and_native_add_generated_column_seams_must_exist():
"""Public sync/async methods and the private native bridge must exist."""
assert hasattr(_native.Table, "_add_generated_column"), (
"native private bridge Table._add_generated_column is missing"
)
assert hasattr(AsyncTable, "add_generated_column"), (
"AsyncTable.add_generated_column is missing"
)
assert hasattr(Table, "add_generated_column"), (
"Table.add_generated_column is missing"
)
assert hasattr(LanceTable, "add_generated_column"), (
"LanceTable.add_generated_column is missing"
)
assert hasattr(RemoteTable, "add_generated_column"), (
"RemoteTable.add_generated_column is missing"
)
# Once present, freeze the exact public positional surface.
_assert_exact_public_signature(Table.add_generated_column)
_assert_exact_public_signature(LanceTable.add_generated_column)
_assert_exact_public_signature(RemoteTable.add_generated_column)
_assert_exact_public_signature(AsyncTable.add_generated_column)
def test_sync_remote_add_generated_column_returns_job_without_eager_wrapper_mutation():
db, table, normalize, log, cm = _open_table_and_function(job_id=_JOB_ID_SYNC)
try:
# Capture public wrapper state before the operation window.
schema_before = table.schema
version_before = table.version
log.start()
call = normalize(text=col("text"))
# Exact public argument order from the frozen user example.
job = table.add_generated_column(
"normalized_text",
call,
)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_SYNC
# Exact success path stops after submit: one binding describe, one create,
# and no catalog re-lookup. Do not wait yet.
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert len(log.create) == 1
assert log.other_table == []
# Public schema/version through the existing wrapper must still reflect
# the pre-submit table: generated column is not published by Job accept.
# Access both before wait so eager wrapper cache invalidation / refresh /
# version advancement is observable.
schema_after = table.schema
assert "normalized_text" not in schema_after.names
assert schema_after == schema_before
# Schema must be served from the existing wrapper cache — no extra
# describe beyond the one binding snapshot.
assert len(log.describe) == 1
assert len(log.lookup) == 0
assert len(log.create) == 1
version_after = table.version
assert version_after == version_before
# Public Remote ``version`` always describes once by design; that probe
# must not drag a schema-cache miss, create, or catalog lookup with it.
assert len(log.describe) == 2
assert len(log.lookup) == 0
assert len(log.create) == 1
assert log.other_table == []
waited = job.wait()
assert waited is None
assert len(log.lookup) == 0
assert len(log.create) == 1
finally:
cm.__exit__(None, None, None)
@pytest.mark.asyncio
async def test_async_remote_add_generated_column_returns_async_job_and_wait_none():
log = _RequestLog()
state = {"opened": False}
binding_describe = _describe_body()
open_describe = {"version": 1, "schema": _text_schema_fields()}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _LOOKUP_PATH:
log.note(request.path, body)
_json_response(request, {"function": _sample_function_wire()})
return
if request.path == _JOB_DESCRIBE_PATH:
assert body["job_id"] == _JOB_ID_ASYNC
_json_response(request, _create_gen_column_done_body(_JOB_ID_ASYNC))
return
if request.path == _DESCRIBE_PATH:
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding_describe)
return
if request.path == _CREATE_PATH:
log.note(request.path, body)
_json_response(request, {"job_id": _JOB_ID_ASYNC})
return
request.send_response(404)
request.end_headers()
async with _mock_remote_db_async(handler) as db:
normalize = await db.functions.get(_CATALOG_NAME)
table = await db.open_table(_TABLE_NAME)
log.start()
call = normalize(text=col("text"))
job = await table.add_generated_column("normalized_text", call)
assert type(job) is lancedb.job.AsyncJob
assert job.id == _JOB_ID_ASYNC
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert len(log.create) == 1
waited = await job.wait()
assert waited is None
assert len(log.lookup) == 0
assert len(log.create) == 1
def test_remote_add_generated_column_one_describe_one_create_exact_envelope():
db, table, normalize, log, cm = _open_table_and_function(job_id=_JOB_ID_SYNC)
try:
call = normalize(text=col("text"))
job = table.add_generated_column("normalized_text", call)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_SYNC
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert len(log.create) == 1
assert log.other_table == []
_assert_exact_create_envelope(
log.create[0],
source_table_version=_SOURCE_TABLE_VERSION,
column_name="normalized_text",
field_id=_TEXT_FIELD_ID,
)
finally:
cm.__exit__(None, None, None)
def test_remote_branch_add_generated_column_includes_exact_branch_identity():
branch_describe = _describe_body(
version=_BRANCH_SOURCE_VERSION,
field_ids=[_BRANCH_TEXT_FIELD_ID],
)
db, table, normalize, log, cm = _open_table_and_function(
describe_body=branch_describe,
job_id=_JOB_ID_BRANCH,
support_branch_create=True,
)
try:
branched = table.branches.create(_BRANCH_NAME)
assert isinstance(branched, RemoteTable)
assert branched.current_branch() == _BRANCH_NAME
log.start()
call = normalize(text=col("text"))
job = branched.add_generated_column("normalized_text", call)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_BRANCH
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert len(log.create) == 1
assert log.describe[0].get("branch") == _BRANCH_NAME
_assert_exact_create_envelope(
log.create[0],
source_table_version=_BRANCH_SOURCE_VERSION,
column_name="normalized_text",
field_id=_BRANCH_TEXT_FIELD_ID,
branch=_BRANCH_NAME,
)
finally:
cm.__exit__(None, None, None)
def test_empty_column_name_fails_locally_with_zero_table_requests():
db, table, normalize, log, cm = _open_table_and_function()
try:
# Authored call owns a real literal so payload-free failure is not vacuous.
call = normalize(text=_LITERAL_PAYLOAD_SENTINEL)
with pytest.raises((ValueError, TypeError)) as raised:
table.add_generated_column("", call)
text = _exception_text(raised.value)
lowered = text.lower()
assert "column" in lowered or "empty" in lowered or "non-empty" in lowered
assert _DESCRIBE_BODY_MARKER not in text
assert _CREATE_RESPONSE_MARKER not in text
assert _LITERAL_PAYLOAD_SENTINEL not in text
_assert_no_operation_traffic(log)
finally:
cm.__exit__(None, None, None)
@pytest.mark.parametrize(
("column_ref", "expected_token"),
[
("missing_text", "missing_text"),
("Text", "Text"), # exact-case mismatch against schema field "text"
],
)
def test_missing_or_case_mismatch_column_one_describe_zero_create(
column_ref: str, expected_token: str
):
db, table, normalize, log, cm = _open_table_and_function()
try:
call = normalize(text=col(column_ref))
with pytest.raises(ValueError) as raised:
table.add_generated_column("normalized_text", call)
text = _exception_text(raised.value)
assert expected_token in text
assert "text" in text # parameter name from the Function signature
assert "missing" in text.lower() or "field" in text.lower()
assert _DESCRIBE_BODY_MARKER not in text
assert _CREATE_RESPONSE_MARKER not in text
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert log.create == []
assert log.other_table == []
finally:
cm.__exit__(None, None, None)
def test_type_mismatch_one_describe_zero_create_identifies_parameter():
db, table, normalize, log, cm = _open_table_and_function(
describe_body=_describe_body(arrow_type="int32"),
)
try:
call = normalize(text=col("text"))
with pytest.raises(ValueError) as raised:
table.add_generated_column("normalized_text", call)
text = _exception_text(raised.value)
assert "text" in text
assert "type" in text.lower() or "mismatch" in text.lower()
assert _DESCRIBE_BODY_MARKER not in text
assert _CREATE_RESPONSE_MARKER not in text
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert log.create == []
assert log.other_table == []
finally:
cm.__exit__(None, None, None)
def test_literal_payload_stays_out_of_field_binding_failure():
"""Authored call owns a real literal; later field binding fails payload-free."""
wire = _sample_function_wire(
parameters=[
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "prefix", "data_type_ipc": _UTF8_TYPE_IPC_B64},
]
)
db, table, normalize, log, cm = _open_table_and_function(function_wire=wire)
try:
call = normalize(text=col("missing_text"), prefix=_LITERAL_PAYLOAD_SENTINEL)
with pytest.raises(ValueError) as raised:
table.add_generated_column("normalized_text", call)
text = _exception_text(raised.value)
assert "missing_text" in text
assert _LITERAL_PAYLOAD_SENTINEL not in text
assert _DESCRIBE_BODY_MARKER not in text
assert _CREATE_RESPONSE_MARKER not in text
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert log.create == []
assert log.other_table == []
finally:
cm.__exit__(None, None, None)
@pytest.mark.asyncio
async def test_closed_async_table_fails_with_zero_operation_requests():
log = _RequestLog()
state = {"opened": False}
binding_describe = _describe_body()
open_describe = {"version": 1, "schema": _text_schema_fields()}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _LOOKUP_PATH:
log.note(request.path, body)
_json_response(request, {"function": _sample_function_wire()})
return
if request.path == _DESCRIBE_PATH:
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding_describe)
return
if request.path == _CREATE_PATH:
log.note(request.path, body)
_json_response(request, {"job_id": _JOB_ID_ASYNC})
return
request.send_response(404)
request.end_headers()
async with _mock_remote_db_async(handler) as db:
normalize = await db.functions.get(_CATALOG_NAME)
table = await db.open_table(_TABLE_NAME)
call = normalize(text=col("text"))
# Public close only — do not mutate private implementation fields.
table.close()
log.start()
try:
await table.add_generated_column("normalized_text", call)
except AttributeError:
# Method missing: re-raise so the failure names the public seam.
raise
except Exception as exc:
text = _exception_text(exc)
assert "closed" in text.lower()
else:
pytest.fail("closed AsyncTable must fail before transport")
_assert_no_operation_traffic(log)
def test_rejects_non_authored_call_before_any_operation_request():
db, table, normalize, log, cm = _open_table_and_function()
try:
bad_values = (
normalize, # exact Function handle itself
{"text": "x"},
col("text"), # direct query Expr
object(),
)
for bad in bad_values:
with pytest.raises(TypeError):
table.add_generated_column("normalized_text", bad)
_assert_no_operation_traffic(log)
finally:
cm.__exit__(None, None, None)
def test_native_valid_call_returns_not_supported_without_mutation(tmp_path):
# Immutable Function handle is connection-free; obtain it via remote lookup.
def lookup_only(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
_json_response(request, {"function": _sample_function_wire()})
with _mock_remote_db(lookup_only) as remote_db:
normalize = _lookup_function(remote_db)
db = lancedb.connect(tmp_path)
table = db.create_table(_TABLE_NAME, [{"text": "Hello"}, {"text": "World"}])
assert isinstance(table, LanceTable)
version_before = table.version
schema_before = table.schema
rows_before = table.to_arrow().to_pylist()
call = normalize(text=col("text"))
with pytest.raises(NotImplementedError) as raised:
table.add_generated_column("normalized_text", call)
text = _exception_text(raised.value)
assert "not supported" in text.lower() or "submit_create_generated_column" in text
assert "add_columns" not in text.lower()
assert table.version == version_before
assert table.schema == schema_before
assert "normalized_text" not in table.schema.names
assert table.to_arrow().to_pylist() == rows_before
def test_public_surface_is_minimal_and_private_call_stays_opaque():
for name in _FORBIDDEN_PUBLIC_NAMES:
assert name not in getattr(lancedb, "__all__", [])
assert not hasattr(lancedb, name)
assert not hasattr(lancedb, "_FunctionCall")
authored_type = getattr(_native, "_FunctionCall", None)
assert authored_type is not None
with pytest.raises(TypeError):
authored_type()
# When the public method exists, reject overdesign kwargs and keep the frozen
# positional surface: (self, column_name, call).
if hasattr(Table, "add_generated_column"):
_assert_exact_public_signature(Table.add_generated_column)
for keyword in _FORBIDDEN_METHOD_KWARGS:
assert (
keyword not in inspect.signature(Table.add_generated_column).parameters
)
db, table, normalize, log, cm = _open_table_and_function()
try:
call = normalize(text=col("text"))
for keyword in _FORBIDDEN_METHOD_KWARGS:
with pytest.raises(TypeError):
table.add_generated_column(
"normalized_text",
call,
**{keyword: object()},
)
_assert_no_operation_traffic(log)
finally:
cm.__exit__(None, None, None)
if hasattr(LanceTable, "add_generated_column"):
_assert_exact_public_signature(LanceTable.add_generated_column)
if hasattr(RemoteTable, "add_generated_column"):
_assert_exact_public_signature(RemoteTable.add_generated_column)
if hasattr(AsyncTable, "add_generated_column"):
_assert_exact_public_signature(AsyncTable.add_generated_column)
for keyword in _FORBIDDEN_METHOD_KWARGS:
assert (
keyword
not in inspect.signature(AsyncTable.add_generated_column).parameters
)
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,672 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python ``table.generated_column_status`` (B3d2).
Public user shape under test:
status = table.generated_column_status("complete_col") # "complete" | "incomplete"
These tests exercise the live worktree PyO3 extension and public sync/async
wrappers. While the public methods and hidden native bridge are absent they
fail against that extension; once present they freeze the public contract
below. They must not fake success paths.
"""
from __future__ import annotations
import contextlib
import http.server
import inspect
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable, Literal, get_type_hints
import pytest
import lancedb
import lancedb.table
from lancedb import _lancedb as _native
from lancedb.remote.table import RemoteTable
from lancedb.table import AsyncTable, LanceTable, Table
_TABLE_NAME = "articles"
_DESCRIBE_PATH = f"/v1/table/{_TABLE_NAME}/describe/"
_ORDINARY_FIELD_ID = 1
_COMPLETE_FIELD_ID = 5
_INCOMPLETE_FIELD_ID = 7
_STABLE_FIELD_IDS = [_ORDINARY_FIELD_ID, _COMPLETE_FIELD_ID, _INCOMPLETE_FIELD_ID]
_STATUS_FUNCTION_ID = "fn.exact.status.projection"
_METADATA_KEY = "lancedb::generated_column"
_RAW_METADATA_MARKER = "SENSITIVE_STATUS_METADATA_MARKER_b3d2_py_9f2e"
# Pinned Rust-canonical schema-only Utf8 type IPC (base64), shared with FF-028.
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_EXPECTED_RETURN = Literal["complete", "incomplete"]
_FORBIDDEN_PUBLIC_NAMES = (
"GeneratedColumnStatus",
"GeneratedColumnDefinition",
"GeneratedColumnBindingSnapshot",
"GeneratedColumnBindingEntry",
)
_FORBIDDEN_BRIDGE_KWARGS = (
"epoch",
"dependency_epoch",
"materialized_epoch",
"function_id",
"field_id",
"field_ids",
"version",
"branch",
"wait",
"job",
"request",
"backend",
)
def _definition_metadata_json(
output_field_id: int,
dependency_epoch: int,
materialized_epoch: int,
*,
text_field_id: int = _ORDINARY_FIELD_ID,
) -> str:
"""Exact JSON stored under Arrow field metadata ``lancedb::generated_column``."""
return json.dumps(
{
"format_version": 1,
"output_field_id": output_field_id,
"function_call": {
"function_id": _STATUS_FUNCTION_ID,
"arguments": [
{
"parameter": "text",
"value": {
"kind": "field",
"field_id": text_field_id,
"data_type_ipc": _UTF8_TYPE_IPC_B64,
},
}
],
},
"dependency_epoch": dependency_epoch,
"materialized_epoch": materialized_epoch,
},
separators=(",", ":"),
)
def _field(
name: str,
*,
arrow_type: str = "string",
nullable: bool = True,
metadata: dict[str, str] | None = None,
) -> dict[str, Any]:
body: dict[str, Any] = {
"name": name,
"type": {"type": arrow_type},
"nullable": nullable,
}
if metadata is not None:
body["metadata"] = metadata
return body
def _status_schema_fields(
*,
complete_meta: str | None = None,
incomplete_meta: str | None = None,
bad_name: str | None = None,
bad_meta: str | None = None,
) -> dict[str, Any]:
fields = [
_field("ordinary", arrow_type="string"),
_field(
"complete_col",
arrow_type="int32",
metadata={
_METADATA_KEY: complete_meta
if complete_meta is not None
else _definition_metadata_json(_COMPLETE_FIELD_ID, 3, 3)
},
),
_field(
"incomplete_col",
arrow_type="int32",
metadata={
_METADATA_KEY: incomplete_meta
if incomplete_meta is not None
else _definition_metadata_json(_INCOMPLETE_FIELD_ID, 4, 1)
},
),
]
if bad_name is not None and bad_meta is not None:
fields.append(
_field(
bad_name,
arrow_type="int32",
metadata={_METADATA_KEY: bad_meta},
)
)
return {"fields": fields}
def _describe_body(
*,
version: int = 11,
field_ids: list[int] | None = _STABLE_FIELD_IDS,
schema: dict[str, Any] | None = None,
) -> dict[str, Any]:
body: dict[str, Any] = {
"version": version,
"schema": schema if schema is not None else _status_schema_fields(),
}
if field_ids is not None:
body["field_ids"] = field_ids
return body
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _exception_text(exc: BaseException) -> str:
parts = [str(exc), repr(exc)]
current: BaseException | None = exc
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(f"{type(current).__name__}: {current}")
current = current.__cause__ or current.__context__
return "\n".join(parts)
def _json_response(
request: http.server.BaseHTTPRequestHandler, body: dict[str, Any]
) -> None:
payload = json.dumps(body).encode("utf-8")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(payload)
class _RequestLog:
"""Track post-open describe and any non-describe operation traffic."""
def __init__(self) -> None:
self.describe: list[dict[str, Any]] = []
self.other: list[str] = []
self.recording = False
def start(self) -> None:
self.describe.clear()
self.other.clear()
self.recording = True
def note(self, path: str, body: dict[str, Any] | None = None) -> None:
if not self.recording:
return
if path == _DESCRIBE_PATH:
self.describe.append(body or {})
else:
self.other.append(path)
def _assert_no_operation_traffic(log: _RequestLog) -> None:
assert log.describe == []
assert log.other == []
def _assert_one_status_describe(log: _RequestLog) -> None:
assert len(log.describe) == 1, f"expected one status describe, got {log.describe!r}"
assert log.other == [], f"unexpected non-describe traffic: {log.other!r}"
def _assert_exact_public_signature(method: Any) -> None:
"""Freeze ``(self, column_name)`` with no varargs/kwargs/keyword-only escape."""
params = list(inspect.signature(method).parameters.values())
assert [p.name for p in params] == ["self", "column_name"]
for param in params:
assert param.kind in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
)
assert param.default is inspect.Parameter.empty
assert param.kind is not inspect.Parameter.VAR_POSITIONAL
assert param.kind is not inspect.Parameter.VAR_KEYWORD
assert param.kind is not inspect.Parameter.KEYWORD_ONLY
def _assert_status_string(value: Any, expected: str) -> None:
assert value == expected
assert type(value) is str
assert value in ("complete", "incomplete")
def _open_remote_table(
*,
status_describe: dict[str, Any] | None = None,
):
"""Open sync RemoteTable; return (table, log, cm)."""
log = _RequestLog()
binding = status_describe if status_describe is not None else _describe_body()
open_describe = {
"version": 1,
"schema": {"fields": [_field("ordinary", arrow_type="string")]},
}
state = {"opened": False}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _DESCRIBE_PATH:
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding)
return
log.note(request.path, body)
request.send_response(404)
request.end_headers()
request.wfile.write(b"unexpected path")
cm = _mock_remote_db(handler)
db = cm.__enter__()
table = db.open_table(_TABLE_NAME)
assert isinstance(table, RemoteTable)
log.start()
return table, log, cm
async def _open_remote_table_async(
*,
status_describe: dict[str, Any] | None = None,
):
"""Open async table under a live mock server; return (table, log, cm)."""
log = _RequestLog()
binding = status_describe if status_describe is not None else _describe_body()
open_describe = {
"version": 1,
"schema": {"fields": [_field("ordinary", arrow_type="string")]},
}
state = {"opened": False}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _DESCRIBE_PATH:
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding)
return
log.note(request.path, body)
request.send_response(404)
request.end_headers()
request.wfile.write(b"unexpected path")
cm = _mock_remote_db_async(handler)
db = await cm.__aenter__()
table = await db.open_table(_TABLE_NAME)
assert isinstance(table, AsyncTable)
log.start()
return table, log, cm
def test_no_public_generated_column_status_resource_exported():
"""Baseline: no public status class/enum/resource is exported."""
for mod in (lancedb, lancedb.table, _native):
for name in _FORBIDDEN_PUBLIC_NAMES:
assert not hasattr(mod, name), f"{mod.__name__}.{name} must not be public"
def test_public_surface_signatures_annotations_and_hidden_bridge():
"""Four public methods + hidden native bridge must exist with frozen shape."""
assert hasattr(_native.Table, "_generated_column_status"), (
"native private bridge Table._generated_column_status is missing"
)
assert hasattr(Table, "generated_column_status"), (
"Table.generated_column_status is missing"
)
assert hasattr(LanceTable, "generated_column_status"), (
"LanceTable.generated_column_status is missing"
)
assert hasattr(RemoteTable, "generated_column_status"), (
"RemoteTable.generated_column_status is missing"
)
assert hasattr(AsyncTable, "generated_column_status"), (
"AsyncTable.generated_column_status is missing"
)
bridge = _native.Table._generated_column_status
_assert_exact_public_signature(bridge)
for keyword in _FORBIDDEN_BRIDGE_KWARGS:
assert keyword not in inspect.signature(bridge).parameters
for method in (
Table.generated_column_status,
LanceTable.generated_column_status,
RemoteTable.generated_column_status,
):
_assert_exact_public_signature(method)
assert not inspect.iscoroutinefunction(method)
assert get_type_hints(method)["return"] == _EXPECTED_RETURN
async_method = AsyncTable.generated_column_status
_assert_exact_public_signature(async_method)
assert inspect.iscoroutinefunction(async_method)
assert get_type_hints(async_method)["return"] == _EXPECTED_RETURN
def test_sync_remote_complete_and_incomplete_one_describe_each():
table, log, cm = _open_remote_table()
try:
complete = table.generated_column_status("complete_col")
_assert_status_string(complete, "complete")
_assert_one_status_describe(log)
log.start()
incomplete = table.generated_column_status("incomplete_col")
_assert_status_string(incomplete, "incomplete")
_assert_one_status_describe(log)
finally:
cm.__exit__(None, None, None)
@pytest.mark.asyncio
async def test_async_remote_complete_and_incomplete_one_describe_each():
table, log, cm = await _open_remote_table_async()
try:
complete = await table.generated_column_status("complete_col")
_assert_status_string(complete, "complete")
_assert_one_status_describe(log)
log.start()
incomplete = await table.generated_column_status("incomplete_col")
_assert_status_string(incomplete, "incomplete")
_assert_one_status_describe(log)
finally:
await cm.__aexit__(None, None, None)
@pytest.mark.parametrize(
("column_name", "status_describe", "expected_exc"),
[
(
"missing",
_describe_body(),
ValueError,
),
(
"Complete_Col",
_describe_body(),
ValueError,
),
(
"ordinary",
_describe_body(),
ValueError,
),
(
"complete_col",
_describe_body(
schema=_status_schema_fields(
complete_meta=_definition_metadata_json(
_COMPLETE_FIELD_ID + 1, 3, 3
)
)
),
ValueError,
),
(
"gen_bad",
_describe_body(
field_ids=[*_STABLE_FIELD_IDS, 9],
schema=_status_schema_fields(
bad_name="gen_bad",
bad_meta=(
'{"format_version":1,"output_field_id":9,'
f'"function_call":{_RAW_METADATA_MARKER},'
'"dependency_epoch":1,"materialized_epoch":1}'
),
),
),
ValueError,
),
(
"complete_col",
_describe_body(
schema=_status_schema_fields(
complete_meta=_definition_metadata_json(
_COMPLETE_FIELD_ID, 1, 1
).replace('"format_version":1', '"format_version":2')
)
),
ValueError,
),
(
"incomplete_col",
_describe_body(
schema=_status_schema_fields(
incomplete_meta=_definition_metadata_json(
_INCOMPLETE_FIELD_ID, 1, 2
)
)
),
ValueError,
),
(
"complete_col",
_describe_body(field_ids=None),
NotImplementedError,
),
],
ids=[
"missing",
"case_mismatch",
"ordinary",
"output_id_mismatch",
"malformed_metadata",
"unknown_format_version",
"reversed_epochs",
"old_server_missing_field_ids",
],
)
def test_remote_fail_closed_matrix_one_describe(
column_name: str,
status_describe: dict[str, Any],
expected_exc: type[BaseException],
):
table, log, cm = _open_remote_table(status_describe=status_describe)
try:
with pytest.raises(expected_exc) as raised:
table.generated_column_status(column_name)
text = _exception_text(raised.value)
assert _RAW_METADATA_MARKER not in text
_assert_one_status_describe(log)
finally:
cm.__exit__(None, None, None)
def test_sync_empty_name_zero_post_open_requests():
table, log, cm = _open_remote_table()
try:
with pytest.raises(ValueError):
table.generated_column_status("")
_assert_no_operation_traffic(log)
finally:
cm.__exit__(None, None, None)
@pytest.mark.asyncio
async def test_async_empty_name_zero_post_open_requests():
table, log, cm = await _open_remote_table_async()
try:
with pytest.raises(ValueError):
await table.generated_column_status("")
_assert_no_operation_traffic(log)
finally:
await cm.__aexit__(None, None, None)
@pytest.mark.asyncio
async def test_async_closed_status_empty_validation_wins_and_nonempty_closed():
"""Publicly closed AsyncTable: empty validates first; nonempty is closed."""
table, log, cm = await _open_remote_table_async()
try:
table.close()
log.start()
try:
await table.generated_column_status("complete_col")
except AttributeError:
raise
except Exception as exc:
text = _exception_text(exc)
assert "closed" in text.lower()
else:
pytest.fail("closed AsyncTable must fail before transport")
_assert_no_operation_traffic(log)
log.start()
with pytest.raises(ValueError) as raised:
await table.generated_column_status("")
text = _exception_text(raised.value)
assert "closed" not in text.lower()
_assert_no_operation_traffic(log)
finally:
await cm.__aexit__(None, None, None)
def test_local_sync_ordinary_column_fails_without_side_effects(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table(
"ordinary_only",
[{"ordinary": "alpha", "id": 1}, {"ordinary": "beta", "id": 2}],
)
assert isinstance(table, LanceTable)
version_before = table.version
schema_before = table.schema
data_before = table.to_arrow()
with pytest.raises(ValueError):
table.generated_column_status("ordinary")
assert table.version == version_before
assert table.schema == schema_before
assert table.to_arrow().equals(data_before)
@pytest.mark.asyncio
async def test_local_async_ordinary_column_fails_without_side_effects(tmp_path):
db = await lancedb.connect_async(tmp_path)
table = await db.create_table(
"ordinary_only_async",
[{"ordinary": "alpha", "id": 1}, {"ordinary": "beta", "id": 2}],
)
assert isinstance(table, AsyncTable)
version_before = await table.version()
schema_before = await table.schema()
data_before = await table.to_arrow()
with pytest.raises(ValueError):
await table.generated_column_status("ordinary")
assert await table.version() == version_before
assert await table.schema() == schema_before
assert (await table.to_arrow()).equals(data_before)
+291
View File
@@ -0,0 +1,291 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for the local @udf declaration surface."""
from __future__ import annotations
import importlib
import inspect
import types
import pyarrow as pa
import pytest
import lancedb
from lancedb import Function, Job, udf
from lancedb._udf import _get_udf_config
_REMOVED_AUTHORING_KNOBS = (
"user_version",
"idempotency_key",
"deterministic",
"null_handling",
"null_policy",
"on_error",
"error_policy",
"FunctionVersion",
"artifact",
"digest",
"geneva",
)
def _decorate(fn, **overrides):
kwargs = {
"inputs": {"x": pa.int32()},
"output": pa.int64(),
"python": "3.12",
}
kwargs.update(overrides)
return udf(**kwargs)(fn)
def test_udf_top_level_export_and_identity_metadata_behavior():
assert "udf" in lancedb.__all__
assert udf is lancedb.udf
assert isinstance(importlib.import_module("lancedb._udf"), types.ModuleType)
assert not isinstance(lancedb.udf, types.ModuleType)
def add(x, y=1):
"""Add locally."""
return x + y
original = add
decorated = _decorate(
add,
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
)
assert decorated is original
assert decorated.__name__ == "add"
assert decorated.__doc__ == "Add locally."
assert str(inspect.signature(decorated)) == "(x, y=1)"
assert decorated(2) == 3
assert decorated(2, 5) == 7
assert decorated(x=4, y=6) == 10
def test_udf_config_snapshot_order_defaults_and_immutability():
inputs = {"z": pa.string(), "a": pa.int32()}
packages = ["pkg-b==2", "pkg-a==1"]
def combine(z, a):
return f"{z}:{a}"
decorated = udf(
inputs=inputs,
output=pa.string(),
python="3.11",
packages=packages,
output_nullable=False,
)(combine)
config = _get_udf_config(decorated)
assert config.inputs == (("z", pa.string()), ("a", pa.int32()))
assert isinstance(config.inputs, tuple)
assert config.output == pa.string()
assert config.output_nullable is False
assert config.python == "3.11"
assert config.packages == ("pkg-b==2", "pkg-a==1")
assert isinstance(config.packages, tuple)
inputs["extra"] = pa.bool_()
del inputs["z"]
packages.append("pkg-c==3")
packages[0] = "mutated==0"
assert config.inputs == (("z", pa.string()), ("a", pa.int32()))
assert config.packages == ("pkg-b==2", "pkg-a==1")
for attr in ("inputs", "output", "output_nullable", "python", "packages"):
with pytest.raises(AttributeError):
setattr(config, attr, None)
def defaults_only(x):
return x
defaulted = udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
)(defaults_only)
default_config = _get_udf_config(defaulted)
assert default_config.packages == ()
assert default_config.output_nullable is True
def test_udf_accepts_lambda_and_closure_for_local_declaration():
ambient = "ambient-secret-value-xyz"
lam = udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
)(lambda n: n + 1)
assert lam(3) == 4
assert _get_udf_config(lam).inputs == (("n", pa.int32()),)
def factory(offset):
@udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
packages=["demo==0.1"],
)
def closed(n):
return n + offset + len(ambient)
return closed
closed = factory(10)
assert closed(2) == 12 + len(ambient)
assert _get_udf_config(closed).packages == ("demo==0.1",)
def test_udf_declaration_defers_signature_and_implementation_packaging():
"""Declaration must not validate callable signature or embed implementation."""
def local_add(left, right=1):
return left + right
decorated = udf(
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
python="3.12",
)(local_add)
assert decorated is local_add
assert str(inspect.signature(decorated)) == "(left, right=1)"
assert decorated(2) == 3
assert decorated(2, 5) == 7
config = _get_udf_config(decorated)
assert config.inputs == (("x", pa.int32()), ("y", pa.int32()))
for attr in (
"source",
"module",
"callable",
"function",
"implementation",
"bundle",
"artifact",
"digest",
):
assert not hasattr(config, attr)
def test_udf_lookup_double_decoration_and_non_function_target():
def plain(x):
return x
with pytest.raises((TypeError, ValueError)):
_get_udf_config(plain)
decorated = _decorate(plain)
with pytest.raises((TypeError, ValueError)):
_decorate(decorated)
with pytest.raises(TypeError):
udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)(object())
with pytest.raises(TypeError):
udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)(42)
def test_udf_config_validation_errors():
def target(x):
return x
with pytest.raises(TypeError):
udf({"x": pa.int32()}, pa.int32(), "3.12")(target)
with pytest.raises(TypeError):
_decorate(target, inputs=[("x", pa.int32())])
with pytest.raises(TypeError):
_decorate(target, inputs={1: pa.int32()})
with pytest.raises(ValueError):
_decorate(target, inputs={"": pa.int32()})
with pytest.raises(TypeError):
_decorate(target, inputs={"x": "int32"})
with pytest.raises(TypeError):
_decorate(target, output="int64")
with pytest.raises(TypeError):
_decorate(target, python=3.12)
with pytest.raises(ValueError):
_decorate(target, python="")
with pytest.raises(TypeError):
_decorate(target, packages="pkg==1")
with pytest.raises(ValueError):
_decorate(target, packages=["pkg==1", ""])
with pytest.raises(ValueError):
_decorate(target, packages=["pkg==1", "pkg==1"])
with pytest.raises(TypeError):
_decorate(target, packages=["pkg==1", 2])
with pytest.raises(TypeError):
_decorate(target, output_nullable=1)
with pytest.raises(TypeError):
_decorate(target, output_nullable="true")
def test_udf_rejects_removed_overdesign_and_has_no_durable_side_effects():
params = inspect.signature(udf).parameters
for name in _REMOVED_AUTHORING_KNOBS:
assert name not in params
def score(x):
"""score body marker unique-xyz."""
ambient = "ambient-secret-value-xyz"
return f"{ambient}:{x}"
decorated = _decorate(
score,
packages=["score==1.0"],
output_nullable=True,
)
config = _get_udf_config(decorated)
text = repr(config).lower()
assert "score body marker unique-xyz" not in text
assert "ambient-secret-value-xyz" not in text
for token in (
"user_version",
"idempotency_key",
"deterministic",
"null_policy",
"on_error",
"functionversion",
"artifact",
"digest",
"geneva",
):
assert token not in text
for attr in _REMOVED_AUTHORING_KNOBS:
assert not hasattr(config, attr)
assert not isinstance(decorated, Function)
assert not isinstance(decorated, Job)
for attr in ("id", "function_id", "job", "job_id", "registration"):
assert not hasattr(decorated, attr)
@@ -0,0 +1,490 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for local FunctionCapability authoring and @udf capabilities."""
from __future__ import annotations
import inspect
import pyarrow as pa
import pytest
import lancedb
from lancedb import Function, FunctionCapability, Job, udf
from lancedb._udf import _get_udf_config, _package_udf
_SECRET_REFERENCE = "secret://team/capability-redact-token-xyz"
_SECRET_ENV = "API_TOKEN"
_NETWORK_ORIGIN = "https://api.example.com"
_OVERDESIGN_ATTRS = (
"id",
"function_id",
"FunctionVersion",
"function_version",
"artifact",
"digest",
"authorization",
"authorized",
"value",
"plaintext",
"plaintext_secret",
"secret_value",
"job",
"job_id",
"catalog",
"retry_key",
"user_version",
"idempotency_key",
"deterministic",
"null_handling",
"null_policy",
"on_error",
"error_policy",
"geneva",
)
def _decorate(fn, **overrides):
kwargs = {
"inputs": {"x": pa.int32()},
"output": pa.int64(),
"python": "3.12",
}
kwargs.update(overrides)
return udf(**kwargs)(fn)
def _network(origin: str = _NETWORK_ORIGIN) -> FunctionCapability:
return FunctionCapability.network(origin)
def _secret(
reference: str = _SECRET_REFERENCE,
*,
environment_variable: str = _SECRET_ENV,
) -> FunctionCapability:
return FunctionCapability.secret(
reference,
environment_variable=environment_variable,
)
@udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=False,
)
def packable_without_capabilities(x):
return x + 1
@udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=False,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
],
)
def packable_with_capabilities(x):
return x + 1
def test_function_capability_export_factories_projection_equality_immutability():
assert "FunctionCapability" in lancedb.__all__
assert FunctionCapability is lancedb.FunctionCapability
network = _network()
secret = _secret()
assert network.kind == "network"
assert network.origin == _NETWORK_ORIGIN
assert network.reference is None
assert network.environment_variable is None
assert secret.kind == "secret"
assert secret.reference == _SECRET_REFERENCE
assert secret.environment_variable == _SECRET_ENV
assert secret.origin is None
assert network == FunctionCapability.network(_NETWORK_ORIGIN)
assert secret == FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
)
assert network != secret
assert network != FunctionCapability.network("https://other.example.com")
assert secret != FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable="OTHER_TOKEN",
)
public_attrs = ("kind", "origin", "reference", "environment_variable")
internal_slots = ("_kind", "_origin", "_reference", "_environment_variable")
immutable_attrs = public_attrs + internal_slots
for attr in public_attrs:
with pytest.raises(AttributeError):
setattr(network, attr, None)
with pytest.raises(AttributeError):
setattr(secret, attr, None)
for attr in immutable_attrs:
# Fresh instances per attempt so a RED slot mutation cannot corrupt
# shared fixtures used by later assertions in this test.
fresh_network = _network("https://fresh-immutability.example.com")
fresh_secret = _secret(
"secret://team/fresh-immutability-token",
environment_variable="FRESH_IMMUTABILITY_TOKEN",
)
with pytest.raises(AttributeError):
setattr(fresh_network, attr, None)
with pytest.raises(AttributeError):
setattr(fresh_secret, attr, None)
with pytest.raises(AttributeError):
delattr(fresh_network, attr)
with pytest.raises(AttributeError):
delattr(fresh_secret, attr)
retained_origin = "https://config-retain.example.com"
retained_reference = "secret://team/config-retain-token"
retained_env = "CONFIG_RETAIN_TOKEN"
retained_network = FunctionCapability.network(retained_origin)
retained_secret = FunctionCapability.secret(
retained_reference,
environment_variable=retained_env,
)
expected_capabilities = (
FunctionCapability.network(retained_origin),
FunctionCapability.secret(
retained_reference,
environment_variable=retained_env,
),
)
def retain_target(x):
return x
retained = udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
capabilities=[retained_network, retained_secret],
)(retain_target)
retained_config = _get_udf_config(retained)
assert retained_config.capabilities == expected_capabilities
for attr in immutable_attrs:
with pytest.raises(AttributeError):
setattr(retained_network, attr, "mutated")
with pytest.raises(AttributeError):
setattr(retained_secret, attr, "mutated")
with pytest.raises(AttributeError):
delattr(retained_network, attr)
with pytest.raises(AttributeError):
delattr(retained_secret, attr)
assert retained_config.capabilities == expected_capabilities
assert retained_config.capabilities[0] is retained_network
assert retained_config.capabilities[1] is retained_secret
assert retained_config.capabilities[0].kind == "network"
assert retained_config.capabilities[0].origin == retained_origin
assert retained_config.capabilities[0].reference is None
assert retained_config.capabilities[0].environment_variable is None
assert retained_config.capabilities[1].kind == "secret"
assert retained_config.capabilities[1].reference == retained_reference
assert retained_config.capabilities[1].environment_variable == retained_env
assert retained_config.capabilities[1].origin is None
with pytest.raises(TypeError):
FunctionCapability()
with pytest.raises(TypeError):
FunctionCapability( # type: ignore[call-arg]
kind="network",
origin=_NETWORK_ORIGIN,
)
assert not isinstance(network, Function)
assert not isinstance(secret, Function)
assert not isinstance(network, Job)
assert not isinstance(secret, Job)
for attr in _OVERDESIGN_ATTRS:
assert not hasattr(network, attr)
assert not hasattr(secret, attr)
def test_function_capability_validation_and_secret_redaction():
with pytest.raises(TypeError):
FunctionCapability.network(None) # type: ignore[arg-type]
with pytest.raises(TypeError):
FunctionCapability.network(123) # type: ignore[arg-type]
with pytest.raises(ValueError):
FunctionCapability.network("")
# Backend authorization owns URL/scheme policy; non-empty is enough here.
loose = FunctionCapability.network("example.com")
assert loose.kind == "network"
assert loose.origin == "example.com"
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[misc]
_SECRET_REFERENCE,
_SECRET_ENV,
)
with pytest.raises(TypeError):
FunctionCapability.secret(None, environment_variable=_SECRET_ENV) # type: ignore[arg-type]
with pytest.raises(TypeError):
FunctionCapability.secret(123, environment_variable=_SECRET_ENV) # type: ignore[arg-type]
with pytest.raises(TypeError):
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable=None) # type: ignore[arg-type]
with pytest.raises(TypeError):
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable=1) # type: ignore[arg-type]
with pytest.raises(ValueError) as empty_ref:
FunctionCapability.secret("", environment_variable=_SECRET_ENV)
assert _SECRET_REFERENCE not in str(empty_ref.value)
assert _SECRET_REFERENCE not in repr(empty_ref.value)
with pytest.raises(ValueError) as empty_env:
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable="")
assert _SECRET_REFERENCE not in str(empty_env.value)
assert _SECRET_REFERENCE not in repr(empty_env.value)
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[call-arg]
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
value="super-secret",
)
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[call-arg]
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
plaintext_secret="super-secret",
)
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[call-arg]
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
environment={_SECRET_ENV: "super-secret"},
)
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[call-arg]
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
headers={"Authorization": "Bearer super-secret"},
)
with pytest.raises(TypeError):
FunctionCapability.network( # type: ignore[call-arg]
_NETWORK_ORIGIN,
headers={"X-Trace": "1"},
)
secret = _secret()
assert not hasattr(secret, "value")
assert not hasattr(secret, "plaintext")
assert not hasattr(secret, "plaintext_secret")
assert not hasattr(secret, "secret_value")
secret_text = repr(secret)
assert "secret" in secret_text.lower()
assert _SECRET_ENV in secret_text
assert _SECRET_REFERENCE not in secret_text
assert "super-secret" not in secret_text
network_text = repr(_network())
assert "network" in network_text.lower()
assert _NETWORK_ORIGIN in network_text
def test_udf_capabilities_ordered_immutable_config_default_and_validation():
params = inspect.signature(udf).parameters
assert "capabilities" in params
assert params["capabilities"].kind is inspect.Parameter.KEYWORD_ONLY
assert params["capabilities"].default == ()
def identity_target(x):
"""capabilities identity marker."""
return x + 1
original = identity_target
decorated = _decorate(identity_target)
assert decorated is original
assert decorated.__name__ == "identity_target"
assert decorated.__doc__ == "capabilities identity marker."
assert decorated(2) == 3
assert _get_udf_config(decorated).capabilities == ()
first = _network("https://b.example.com")
second = _network("https://a.example.com")
third = _network("https://b.example.com")
secret = _secret()
capabilities = [first, second, third, secret]
def combine(x):
return x
with_caps = udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
packages=["pkg-b==2", "pkg-a==1"],
capabilities=capabilities,
)(combine)
config = _get_udf_config(with_caps)
assert config.capabilities == (first, second, third, secret)
assert isinstance(config.capabilities, tuple)
assert config.packages == ("pkg-b==2", "pkg-a==1")
assert config.inputs == (("x", pa.int32()),)
capabilities.append(_network("https://mutated.example.com"))
capabilities[0] = _network("https://replaced.example.com")
assert config.capabilities == (first, second, third, secret)
with pytest.raises(AttributeError):
setattr(config, "capabilities", ())
def target(x):
return x
with pytest.raises(TypeError):
_decorate(target, capabilities="https://api.example.com")
with pytest.raises(TypeError):
_decorate(target, capabilities=b"https://api.example.com")
class _BadCapability:
def __repr__(self) -> str:
return "unique-bad-capability-repr-xyz"
with pytest.raises(TypeError) as bad_item:
_decorate(target, capabilities=[_BadCapability()])
assert "unique-bad-capability-repr-xyz" not in str(bad_item.value)
assert "unique-bad-capability-repr-xyz" not in repr(bad_item.value)
with pytest.raises(TypeError) as bad_mixed:
_decorate(
target,
capabilities=[_network(), "unique-bad-capability-string-xyz"],
)
assert "unique-bad-capability-string-xyz" not in str(bad_mixed.value)
assert "unique-bad-capability-string-xyz" not in repr(bad_mixed.value)
def test_udf_capabilities_rejects_function_capability_subclass_before_property_access():
marker = "unique-hostile-capability-subclass-marker-xyz"
class _HostileFunctionCapability(FunctionCapability):
@property
def kind(self) -> str:
raise RuntimeError(marker)
@property
def origin(self) -> str | None:
raise RuntimeError(marker)
@property
def reference(self) -> str | None:
raise RuntimeError(marker)
@property
def environment_variable(self) -> str | None:
raise RuntimeError(marker)
hostile = object.__new__(_HostileFunctionCapability)
assert isinstance(hostile, FunctionCapability)
assert type(hostile) is not FunctionCapability
def target(x):
return x
with pytest.raises(TypeError) as exc_info:
_decorate(target, capabilities=[hostile])
assert marker not in str(exc_info.value)
assert marker not in repr(exc_info.value)
assert _SECRET_REFERENCE not in str(exc_info.value)
assert _SECRET_REFERENCE not in repr(exc_info.value)
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
def test_package_udf_preserves_capabilities_and_redacts_secret_reference():
packaged = _package_udf(packable_with_capabilities)
config = packaged.config
assert packaged.config is _get_udf_config(packable_with_capabilities)
assert config.capabilities == (
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
)
assert config.capabilities[0].kind == "network"
assert config.capabilities[0].origin == _NETWORK_ORIGIN
assert config.capabilities[1].kind == "secret"
assert config.capabilities[1].reference == _SECRET_REFERENCE
assert config.capabilities[1].environment_variable == _SECRET_ENV
assert config.packages == ("pkg-a==1",)
assert config.python == "3.12"
assert config.output_nullable is False
nested = (
f"{packaged!r}\n{config!r}\n{config.capabilities!r}\n{config.capabilities[1]!r}"
)
assert _SECRET_REFERENCE not in nested
assert _SECRET_ENV in repr(config.capabilities[1])
def test_capabilities_are_additive_to_existing_declaration_and_packaging():
def score(x):
return x
decorated = _decorate(
score,
packages=["score==1.0"],
output_nullable=True,
)
config = _get_udf_config(decorated)
assert config.inputs == (("x", pa.int32()),)
assert config.output == pa.int64()
assert config.output_nullable is True
assert config.python == "3.12"
assert config.packages == ("score==1.0",)
assert config.capabilities == ()
assert decorated is score
assert decorated(4) == 4
packaged = _package_udf(packable_without_capabilities)
assert packaged.config is _get_udf_config(packable_without_capabilities)
assert packaged.callable_name == "packable_without_capabilities"
assert packaged.config.capabilities == ()
assert packaged.config.packages == ("pkg-a==1",)
assert packaged.config.output_nullable is False
assert packable_without_capabilities(1) == 2
params = inspect.signature(udf).parameters
for name in (
"user_version",
"idempotency_key",
"deterministic",
"null_handling",
"null_policy",
"on_error",
"error_policy",
"FunctionVersion",
"artifact",
"digest",
"geneva",
):
assert name not in params
@@ -0,0 +1,506 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for the private UDF -> FunctionDefinition bridge."""
from __future__ import annotations
import base64
import io
import json
from pathlib import Path
import pyarrow as pa
import pytest
import lancedb
from lancedb import FunctionCapability, udf
from lancedb import _lancedb as _native
from lancedb import _udf as _udf_mod
_SOURCE_MARKER = "bridge-source-marker-unique-xyz"
_SECRET_REFERENCE = "secret://team/bridge-redact-token-xyz"
_SECRET_ENV = "BRIDGE_API_TOKEN"
_NETWORK_ORIGIN = "https://api.bridge-example.com"
_NETWORK_ORIGIN_B = "https://other.bridge-example.com"
_FORBIDDEN_WIRE_KEYS = (
"id",
"function_id",
"FunctionId",
"catalog",
"catalog_name",
"version",
"function_version",
"FunctionVersion",
"lineage",
"user_version",
"idempotency_key",
"digest",
"artifact",
"artifact_digest",
"storage",
"storage_location",
"location",
"deterministic",
"null_policy",
"nullPolicy",
"timestamp",
"created_at",
"updated_at",
"worker",
"scheduler",
"attempt",
"attempt_id",
"replica",
"placement",
"job",
"job_id",
"retry_key",
"registration",
)
_OVERDESIGN_ATTRS = (
"id",
"function_id",
"FunctionVersion",
"function_version",
"artifact",
"digest",
"job",
"job_id",
"catalog",
"retry_key",
"user_version",
"idempotency_key",
"deterministic",
"null_policy",
"null_handling",
)
@udf(
inputs={"text": pa.string(), "limit": pa.int32()},
output=pa.string(),
python="3.12",
packages=["pkg-b==2", "pkg-a==1"],
output_nullable=True,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
FunctionCapability.network(_NETWORK_ORIGIN_B),
],
)
def packable_bridge_normalize(text, limit):
"""bridge-source-marker-unique-xyz."""
return text[:limit]
def _build_function_definition(fn: object):
return _udf_mod._build_function_definition(fn)
def _function_definition_type():
return _native._FunctionDefinition
def _new_function_definition(**kwargs):
return _native._new_function_definition(**kwargs)
def _json_bytes(definition) -> bytes:
payload = definition._to_json()
if isinstance(payload, bytes):
return payload
assert isinstance(payload, str)
return payload.encode("utf-8")
def _decode_type_ipc(encoded: str) -> pa.DataType:
raw = base64.b64decode(encoded)
reader = pa.ipc.open_file(io.BytesIO(raw))
assert reader.num_record_batches == 0
assert len(reader.schema) == 1
return reader.schema.field(0).type
def _assert_exact_object_keys(value: dict, expected: set[str], *, context: str) -> None:
assert isinstance(value, dict), f"{context} must be an object"
assert set(value) == expected, f"{context} keys must match exactly: {set(value)!r}"
def _assert_forbidden_keys_absent(value: object, *, context: str) -> None:
if isinstance(value, dict):
for key in value:
assert key not in _FORBIDDEN_WIRE_KEYS, (
f"forbidden key {key!r} at {context}: {value!r}"
)
if key == "name" and context in {
"definition",
"signature",
"signature.output",
"implementation",
}:
raise AssertionError(
f"catalog/function identity key `name` must be absent at {context}"
)
child_context = f"{context}.{key}"
if key == "parameters" and context == "signature":
child_context = "signature.parameters"
_assert_forbidden_keys_absent(value[key], context=child_context)
elif isinstance(value, list):
for idx, item in enumerate(value):
item_context = (
f"signature.parameters[{idx}]"
if context == "signature.parameters"
else f"{context}[{idx}]"
)
if context == "signature.parameters":
assert isinstance(item, dict)
assert "name" in item
for key in item:
assert key not in _FORBIDDEN_WIRE_KEYS
assert key != "catalog_name"
_assert_forbidden_keys_absent(
{k: v for k, v in item.items() if k != "name"},
context=item_context,
)
else:
_assert_forbidden_keys_absent(item, context=item_context)
def _assert_sanitized_text(*parts: object) -> None:
combined = "\n".join(str(part) for part in parts)
lowered = combined.lower()
assert _SOURCE_MARKER.lower() not in lowered
assert _SECRET_REFERENCE.lower() not in lowered
assert str(Path(__file__).resolve()).lower() not in lowered
assert Path(__file__).resolve().as_posix().lower() not in lowered
def _assert_clean_validation_error(exc_info) -> None:
_assert_sanitized_text(exc_info.value, repr(exc_info.value))
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
def _valid_builder_kwargs(**overrides):
kwargs = {
"parameters": [("text", pa.string()), ("limit", pa.int32())],
"output_type": pa.string(),
"output_nullable": True,
"module": "bridge_mod",
"callable_name": "normalize",
"source": (
"def normalize(text, limit):\n"
f" # {_SOURCE_MARKER}\n"
" return text[:limit]\n"
),
"python": "3.12",
"packages": ["pkg-b==2", "pkg-a==1"],
"capabilities": [
("network", _NETWORK_ORIGIN, None),
("secret", _SECRET_REFERENCE, _SECRET_ENV),
("network", _NETWORK_ORIGIN_B, None),
],
}
kwargs.update(overrides)
return kwargs
def test_build_function_definition_private_native_immutability_and_export_surface():
assert "_build_function_definition" not in getattr(lancedb, "__all__", [])
assert "_FunctionDefinition" not in lancedb.__all__
assert not hasattr(lancedb, "_FunctionDefinition")
assert not hasattr(lancedb, "_build_function_definition")
assert not hasattr(lancedb, "_new_function_definition")
definition = _build_function_definition(packable_bridge_normalize)
definition_type = _function_definition_type()
assert type(definition) is definition_type
assert definition_type.__module__ == "lancedb._lancedb"
assert definition_type.__name__ == "_FunctionDefinition"
with pytest.raises(TypeError):
definition_type()
for attr in _OVERDESIGN_ATTRS:
assert not hasattr(definition, attr)
for attr in ("signature", "module", "source", "capabilities"):
with pytest.raises(AttributeError):
setattr(definition, attr, None)
def test_build_function_definition_json_wire_ordered_contract_without_identity():
definition = _build_function_definition(packable_bridge_normalize)
encoded_a = _json_bytes(definition)
encoded_b = _json_bytes(definition)
assert encoded_a == encoded_b
wire = json.loads(encoded_a.decode("utf-8"))
_assert_exact_object_keys(
wire,
{"format_version", "signature", "implementation", "capabilities"},
context="definition",
)
assert wire["format_version"] == 1
_assert_forbidden_keys_absent(wire, context="definition")
signature = wire["signature"]
_assert_exact_object_keys(signature, {"parameters", "output"}, context="signature")
parameters = signature["parameters"]
assert [parameter["name"] for parameter in parameters] == ["text", "limit"]
for parameter in parameters:
_assert_exact_object_keys(
parameter, {"name", "data_type_ipc"}, context="parameter"
)
assert isinstance(parameter["data_type_ipc"], str)
assert parameter["data_type_ipc"]
assert _decode_type_ipc(parameters[0]["data_type_ipc"]) == pa.string()
assert _decode_type_ipc(parameters[1]["data_type_ipc"]) == pa.int32()
output = signature["output"]
_assert_exact_object_keys(
output, {"data_type_ipc", "nullable"}, context="signature.output"
)
assert output["nullable"] is True
assert _decode_type_ipc(output["data_type_ipc"]) == pa.string()
implementation = wire["implementation"]
_assert_exact_object_keys(
implementation,
{"kind", "module", "callable", "source", "python", "packages"},
context="implementation",
)
assert implementation["kind"] == "python"
assert implementation["module"] == __name__
assert implementation["callable"] == "packable_bridge_normalize"
assert implementation["source"] == Path(__file__).read_text(encoding="utf-8")
assert _SOURCE_MARKER in implementation["source"]
assert implementation["python"] == "3.12"
assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"]
capabilities = wire["capabilities"]
assert capabilities == [
{"kind": "network", "origin": _NETWORK_ORIGIN},
{
"kind": "secret",
"reference": _SECRET_REFERENCE,
"environment_variable": _SECRET_ENV,
},
{"kind": "network", "origin": _NETWORK_ORIGIN_B},
]
for capability in capabilities:
assert "value" not in capability
assert "plaintext" not in capability
assert "plaintext_secret" not in capability
assert "secret_value" not in capability
def test_native_definition_repr_includes_safe_structure_and_redacts_sensitive_text():
definition = _build_function_definition(packable_bridge_normalize)
rendered = repr(definition)
assert "_FunctionDefinition" in rendered or "FunctionDefinition" in rendered
assert __name__ in rendered
assert "packable_bridge_normalize" in rendered
assert "3.12" in rendered
_assert_sanitized_text(rendered)
def test_new_function_definition_builder_preserves_normalized_wire():
definition = _new_function_definition(**_valid_builder_kwargs())
assert type(definition) is _function_definition_type()
encoded_a = _json_bytes(definition)
encoded_b = _json_bytes(definition)
assert encoded_a == encoded_b
wire = json.loads(encoded_a.decode("utf-8"))
assert wire["format_version"] == 1
assert [parameter["name"] for parameter in wire["signature"]["parameters"]] == [
"text",
"limit",
]
assert _decode_type_ipc(wire["signature"]["parameters"][0]["data_type_ipc"]) == (
pa.string()
)
assert _decode_type_ipc(wire["signature"]["parameters"][1]["data_type_ipc"]) == (
pa.int32()
)
assert wire["signature"]["output"]["nullable"] is True
assert _decode_type_ipc(wire["signature"]["output"]["data_type_ipc"]) == pa.string()
implementation = wire["implementation"]
assert implementation["kind"] == "python"
assert implementation["module"] == "bridge_mod"
assert implementation["callable"] == "normalize"
assert implementation["source"] == _valid_builder_kwargs()["source"]
assert _SOURCE_MARKER in implementation["source"]
assert implementation["python"] == "3.12"
assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"]
assert wire["capabilities"] == [
{"kind": "network", "origin": _NETWORK_ORIGIN},
{
"kind": "secret",
"reference": _SECRET_REFERENCE,
"environment_variable": _SECRET_ENV,
},
{"kind": "network", "origin": _NETWORK_ORIGIN_B},
]
_assert_forbidden_keys_absent(wire, context="definition")
@pytest.mark.parametrize(
("overrides",),
[
({"parameters": [("text", pa.string()), ("text", pa.int32())]},),
({"parameters": [("", pa.string())]},),
({"module": ""},),
({"callable_name": ""},),
({"source": ""},),
({"python": ""},),
({"packages": ["pkg-a==1", ""]},),
({"packages": ["pkg-a==1", "pkg-a==1"]},),
({"capabilities": [("filesystem", _NETWORK_ORIGIN, None)]},),
({"capabilities": [("network", _NETWORK_ORIGIN, _SECRET_ENV)]},),
({"capabilities": [("secret", _SECRET_REFERENCE, None)]},),
({"capabilities": [("secret", _SECRET_REFERENCE, "")]},),
({"capabilities": [("network", "", None)]},),
({"capabilities": [("secret", "", _SECRET_ENV)]},),
],
)
def test_new_function_definition_strict_validation_rejections(overrides):
kwargs = _valid_builder_kwargs(**overrides)
with pytest.raises(ValueError) as exc_info:
_new_function_definition(**kwargs)
_assert_clean_validation_error(exc_info)
def test_new_function_definition_validation_does_not_echo_secret_or_source_marker():
with pytest.raises(ValueError) as exc_info:
_new_function_definition(**_valid_builder_kwargs(module=""))
_assert_clean_validation_error(exc_info)
with pytest.raises(ValueError) as exc_info:
_new_function_definition(
**_valid_builder_kwargs(packages=["pkg-a==1", "pkg-a==1"])
)
_assert_clean_validation_error(exc_info)
with pytest.raises(ValueError) as exc_info:
_new_function_definition(
**_valid_builder_kwargs(
capabilities=[("secret", _SECRET_REFERENCE, None)],
)
)
_assert_clean_validation_error(exc_info)
@pytest.mark.parametrize(
("overrides",),
[
({"parameters": [("text", "not-a-datatype")]},),
({"parameters": [(123, pa.string())]},),
({"output_type": "not-a-datatype"},),
({"output_type": None},),
({"output_nullable": "yes"},),
({"packages": "pkg-a==1"},),
({"capabilities": "network"},),
({"capabilities": [("network", _NETWORK_ORIGIN)]},),
({"capabilities": [("network", _NETWORK_ORIGIN, None, "extra")]},),
],
)
def test_new_function_definition_wrong_pyarrow_and_shape_values_fail_closed(overrides):
kwargs = _valid_builder_kwargs(**overrides)
with pytest.raises((TypeError, ValueError)) as exc_info:
_new_function_definition(**kwargs)
_assert_clean_validation_error(exc_info)
class _HostileRaisingIterable:
def __iter__(self):
raise RuntimeError(f"{_SECRET_REFERENCE} {_SOURCE_MARKER}")
@pytest.mark.parametrize(
("overrides",),
[
({"parameters": _HostileRaisingIterable()},),
({"packages": _HostileRaisingIterable()},),
({"capabilities": _HostileRaisingIterable()},),
(
{
"capabilities": [
("network", _NETWORK_ORIGIN, None),
_HostileRaisingIterable(),
("network", _NETWORK_ORIGIN_B, None),
]
},
),
],
)
def test_new_function_definition_hostile_iterable_iter_raises_fail_closed(overrides):
kwargs = _valid_builder_kwargs(**overrides)
with pytest.raises((TypeError, ValueError)) as exc_info:
_new_function_definition(**kwargs)
_assert_clean_validation_error(exc_info)
@udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=False,
)
def packable_bridge_capability_exact_type(x):
return x + 1
def test_build_function_definition_rejects_forged_function_capability_subclass():
marker = f"{_SECRET_REFERENCE} {_SOURCE_MARKER}"
class _HostileFunctionCapability(FunctionCapability):
@property
def kind(self) -> str:
raise RuntimeError(marker)
@property
def origin(self) -> str | None:
raise RuntimeError(marker)
@property
def reference(self) -> str | None:
raise RuntimeError(marker)
@property
def environment_variable(self) -> str | None:
raise RuntimeError(marker)
hostile = object.__new__(_HostileFunctionCapability)
assert isinstance(hostile, FunctionCapability)
assert type(hostile) is not FunctionCapability
config_attr = _udf_mod._CONFIG_ATTR
original = getattr(packable_bridge_capability_exact_type, config_attr)
forged = _udf_mod._UdfConfig(
inputs=original.inputs,
output=original.output,
output_nullable=original.output_nullable,
python=original.python,
packages=original.packages,
capabilities=(hostile,),
)
setattr(packable_bridge_capability_exact_type, config_attr, forged)
try:
with pytest.raises((TypeError, ValueError)) as exc_info:
_build_function_definition(packable_bridge_capability_exact_type)
_assert_clean_validation_error(exc_info)
assert marker not in str(exc_info.value)
assert marker not in repr(exc_info.value)
finally:
setattr(packable_bridge_capability_exact_type, config_attr, original)
@@ -0,0 +1,486 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for private UDF packaging validation."""
from __future__ import annotations
import importlib
import inspect
import json
import sys
from collections.abc import Iterator
from contextlib import contextmanager
from pathlib import Path
import pyarrow as pa
import pytest
from lancedb import Function, Job, udf
from lancedb._udf import _get_udf_config, _package_udf
_BODY_MARKER = "packaging body marker unique-xyz"
_AMBIENT_SECRET = "ambient-secret-value-xyz"
_BUILTIN_SHADOW_SECRET = "builtin-shadow-secret-xyz"
_SOURCE_MISMATCH_SECRET = "source-mismatch-secret-xyz"
_INVALID_UTF8_SECRET = "invalid-utf8-secret-xyz"
_OVERDESIGN_ATTRS = (
"user_version",
"idempotency_key",
"deterministic",
"null_handling",
"null_policy",
"on_error",
"error_policy",
"FunctionVersion",
"function_version",
"artifact",
"digest",
"geneva",
"id",
"function_id",
"job",
"job_id",
"registration",
"catalog",
"retry_key",
"source_path",
"path",
"function",
)
_PACKAGING_CONSTANT = 41
def _packaging_helper(value: int) -> int:
return value + _PACKAGING_CONSTANT
@udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=False,
)
def packable_add(x):
"""packaging body marker unique-xyz."""
return _packaging_helper(x) + len(json.dumps({"k": 1}))
@udf(
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
python="3.12",
)
def packable_kwonly(x, *, y=2):
return x + y
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def packable_rebind_target(x):
return x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def uses_injected_ambient(x):
return x + len(INJECTED_AMBIENT_GLOBAL) # noqa: F821
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def uses_shadowed_builtin_len(x):
return x + len((1, 2, 3))
@udf(
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
python="3.12",
)
def mismatch_names(left, right):
return left + right
@udf(
inputs={"y": pa.int32(), "x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def mismatch_order(x, y):
return x + y
@udf(
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
python="3.12",
)
def positional_only(x, /, y):
return x + y
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def varargs_fn(x, *args):
return x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def kwargs_fn(x, **kwargs):
return x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
async def async_fn(x):
return x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
async def async_gen_fn(x):
yield x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def generator_fn(x):
yield x
def _assert_sanitized_text(*parts: object, secret: str = _AMBIENT_SECRET) -> None:
combined = "\n".join(str(part) for part in parts)
lowered = combined.lower()
assert _BODY_MARKER.lower() not in lowered
assert secret.lower() not in lowered
assert str(Path(__file__).resolve()).lower() not in lowered
assert Path(__file__).resolve().as_posix().lower() not in lowered
def _assert_packaging_rejection(exc_info, *, secret: str = _AMBIENT_SECRET) -> None:
_assert_sanitized_text(exc_info.value, repr(exc_info.value), secret=secret)
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
@contextmanager
def _temporary_imported_module(
directory: Path, module_name: str, source: str
) -> Iterator[tuple[Path, object]]:
path = directory / f"{module_name}.py"
path.write_text(source, encoding="utf-8")
inserted = str(directory)
sys.path.insert(0, inserted)
try:
sys.modules.pop(module_name, None)
module = importlib.import_module(module_name)
yield path, module
finally:
sys.modules.pop(module_name, None)
try:
sys.path.remove(inserted)
except ValueError:
pass
def _temp_udf_module_source(*, body: str, secret: str | None = None) -> str:
secret_line = f"_SECRET = {secret!r}\n" if secret is not None else ""
return (
"import pyarrow as pa\n"
"from lancedb import udf\n"
f"{secret_line}\n"
"@udf(\n"
' inputs={"x": pa.int32()},\n'
" output=pa.int32(),\n"
' python="3.12",\n'
")\n"
"def temp_pack_target(x):\n"
f" {body}\n"
)
def test_package_udf_success_snapshot_source_module_callable_config_and_repr():
packaged = _package_udf(packable_add)
source = Path(__file__).read_text(encoding="utf-8")
assert packaged.source == source
assert packaged.module == __name__
assert packaged.module != "__main__"
assert packaged.callable_name == "packable_add"
assert packable_add.__qualname__ == "packable_add"
assert packaged.config is _get_udf_config(packable_add)
assert packaged.config.inputs == (("x", pa.int32()),)
assert packaged.config.output == pa.int64()
assert packaged.config.output_nullable is False
assert packaged.config.python == "3.12"
assert packaged.config.packages == ("pkg-a==1",)
for attr in ("source", "module", "callable_name", "config"):
with pytest.raises(AttributeError):
setattr(packaged, attr, None)
text = repr(packaged)
_assert_sanitized_text(text)
assert _BODY_MARKER not in text
def test_package_udf_allows_source_bound_import_constant_and_helper():
packaged = _package_udf(packable_add)
assert packaged.callable_name == "packable_add"
assert "import json" in packaged.source
assert "_PACKAGING_CONSTANT" in packaged.source
assert "_packaging_helper" in packaged.source
assert packable_add(1) == _packaging_helper(1) + len(json.dumps({"k": 1}))
def test_package_udf_accepts_positional_or_keyword_and_keyword_only_defaults():
packaged = _package_udf(packable_kwonly)
assert packaged.callable_name == "packable_kwonly"
assert packaged.config.inputs == (("x", pa.int32()), ("y", pa.int32()))
assert str(inspect.signature(packable_kwonly)) == "(x, *, y=2)"
assert packable_kwonly(3) == 5
assert packable_kwonly(3, y=7) == 10
def test_package_udf_rejects_lambda_and_closure():
lam = udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
)(lambda n: n + 1)
with pytest.raises(ValueError) as exc_info:
_package_udf(lam)
_assert_packaging_rejection(exc_info)
ambient = _AMBIENT_SECRET
def factory(offset):
@udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
)
def closed(n):
return n + offset + len(ambient)
return closed
closed = factory(10)
with pytest.raises(ValueError) as exc_info:
_package_udf(closed)
_assert_packaging_rejection(exc_info)
def outer():
total = 0
@udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
)
def nested(n):
nonlocal total
total += n
return total
return nested
with pytest.raises(ValueError) as exc_info:
_package_udf(outer())
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_signature_mismatches_and_unsupported_parameter_kinds():
for target in (
mismatch_names,
mismatch_order,
positional_only,
varargs_fn,
kwargs_fn,
):
with pytest.raises(ValueError) as exc_info:
_package_udf(target)
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_async_and_generator_functions():
for target in (async_fn, async_gen_fn, generator_fn):
with pytest.raises(ValueError) as exc_info:
_package_udf(target)
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_dynamic_exec_source():
namespace: dict[str, object] = {}
exec(
"def dynamic_pack_target(x):\n return x + 1\n",
namespace,
)
dynamic = udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)(namespace["dynamic_pack_target"])
with pytest.raises(ValueError) as exc_info:
_package_udf(dynamic)
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_undecorated_and_wrong_input_types():
def plain(x):
return x
with pytest.raises(TypeError) as exc_info:
_package_udf(plain)
_assert_packaging_rejection(exc_info)
with pytest.raises(TypeError) as exc_info:
_package_udf(object())
_assert_packaging_rejection(exc_info)
with pytest.raises(TypeError) as exc_info:
_package_udf(42)
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_rebound_module_attribute():
module = sys.modules[__name__]
original = module.packable_rebind_target
replacement = udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)(lambda x: x)
module.packable_rebind_target = replacement
try:
with pytest.raises(ValueError) as exc_info:
_package_udf(original)
_assert_packaging_rejection(exc_info)
finally:
module.packable_rebind_target = original
def test_package_udf_rejects_injected_ambient_global():
module = sys.modules[__name__]
secret = _AMBIENT_SECRET
module.INJECTED_AMBIENT_GLOBAL = secret
try:
assert uses_injected_ambient(3) == 3 + len(secret)
with pytest.raises(ValueError) as exc_info:
_package_udf(uses_injected_ambient)
_assert_packaging_rejection(exc_info, secret=secret)
finally:
delattr(module, "INJECTED_AMBIENT_GLOBAL")
def test_package_udf_rejects_builtin_shadow_injection():
module = sys.modules[__name__]
secret = _BUILTIN_SHADOW_SECRET
assert not hasattr(module, "len")
module.len = secret
try:
with pytest.raises(ValueError) as exc_info:
_package_udf(uses_shadowed_builtin_len)
_assert_packaging_rejection(exc_info, secret=secret)
finally:
delattr(module, "len")
def test_package_udf_rejects_loaded_code_source_mismatch(tmp_path: Path):
secret = _SOURCE_MISMATCH_SECRET
module_name = "udf_pkg_source_mismatch_mod"
original = _temp_udf_module_source(body="return x + 1")
replacement = _temp_udf_module_source(
body=f"return x + 99 # {secret}",
secret=secret,
)
with _temporary_imported_module(tmp_path, module_name, original) as (
path,
module,
):
target = module.temp_pack_target
assert target(1) == 2
path.write_text(replacement, encoding="utf-8")
with pytest.raises(ValueError) as exc_info:
_package_udf(target)
_assert_packaging_rejection(exc_info, secret=secret)
err_text = f"{exc_info.value}\n{exc_info.value!r}"
assert str(path.resolve()) not in err_text
assert path.resolve().as_posix() not in err_text
def test_package_udf_rejects_invalid_utf8_after_import(tmp_path: Path):
secret = _INVALID_UTF8_SECRET
module_name = "udf_pkg_invalid_utf8_mod"
original = _temp_udf_module_source(body="return x + 1")
with _temporary_imported_module(tmp_path, module_name, original) as (
path,
module,
):
target = module.temp_pack_target
assert target(1) == 2
path.write_bytes(secret.encode("utf-8") + b"\xff\xfe invalid-bytes")
with pytest.raises(ValueError) as exc_info:
_package_udf(target)
assert type(exc_info.value) is ValueError
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
err_text = f"{exc_info.value}\n{exc_info.value!r}"
assert secret not in err_text
assert "b'" not in err_text
assert r"\xff" not in err_text
assert str(path.resolve()) not in err_text
assert path.resolve().as_posix() not in err_text
def test_package_udf_snapshot_has_no_durable_overdesign_and_is_not_function_or_job():
packaged = _package_udf(packable_add)
assert not isinstance(packaged, Function)
assert not isinstance(packaged, Job)
for attr in _OVERDESIGN_ATTRS:
assert not hasattr(packaged, attr)
text = repr(packaged).lower()
for token in (
"user_version",
"idempotency_key",
"deterministic",
"null_policy",
"on_error",
"functionversion",
"artifact",
"digest",
"geneva",
"retry_key",
):
assert token not in text
_assert_sanitized_text(text)
+225
View File
@@ -2306,3 +2306,228 @@ def test_remote_connection_jobs_surface():
assert job.status() == "failed"
with pytest.raises(JobFailedError, match="worker died"):
job.wait(timeout=timedelta(seconds=5))
# Pinned Rust-canonical schema-only type IPC (base64). PyArrow's schema-only
# FileWriter bytes are not byte-identical to the Arrow Rust FileWriter used by
# the strict Function decoder, so these fixtures are derived from Rust serde.
_FIRST_CLASS_FUNCTION_JOB_RESULT_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_FIRST_CLASS_FUNCTION_JOB_RESULT_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID = "fn.exact.python-job-result"
_FIRST_CLASS_FUNCTION_JOB_RESULT_ABSENT = object()
_FIRST_CLASS_FUNCTION_JOB_RESULT_NULL = object()
def _first_class_function_job_result_function_wire():
int32_ipc = _FIRST_CLASS_FUNCTION_JOB_RESULT_INT32_TYPE_IPC_B64
utf8_ipc = _FIRST_CLASS_FUNCTION_JOB_RESULT_UTF8_TYPE_IPC_B64
return {
"kind": "function",
"format_version": 1,
"function": {
"format_version": 1,
"id": _FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "x", "data_type_ipc": int32_ipc},
{"name": "label", "data_type_ipc": utf8_ipc},
],
"output": {
"data_type_ipc": int32_ipc,
"nullable": True,
},
},
},
}
def _first_class_function_job_result_none_wire():
return {"kind": "none", "format_version": 1}
def _first_class_function_job_result_describe_body(
job_id, job_type, result=_FIRST_CLASS_FUNCTION_JOB_RESULT_ABSENT
):
body = {
"job_id": job_id,
"job_state": "DONE",
"job_type": job_type,
"creation_ms": 1,
"spec": {},
}
if result is _FIRST_CLASS_FUNCTION_JOB_RESULT_NULL:
body["result"] = None
elif result is not _FIRST_CLASS_FUNCTION_JOB_RESULT_ABSENT:
body["result"] = result
return body
def _first_class_function_job_result_describe_handler(bodies_by_job_id):
def handler(request):
content_len = int(request.headers.get("Content-Length", 0))
body = request.rfile.read(content_len) if content_len > 0 else b""
payload = json.loads(body) if body else {}
if request.path != "/v1/jobs/describe":
request.send_response(404)
request.end_headers()
return
job_id = payload["job_id"]
if job_id not in bodies_by_job_id:
request.send_response(404)
request.end_headers()
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(bodies_by_job_id[job_id]).encode())
return handler
def _assert_exact_first_class_function_job_result(function):
assert isinstance(function, lancedb.Function)
assert function is not None
assert not isinstance(function, dict)
assert function.id == _FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID
assert function.parameters == (("x", pa.int32()), ("label", pa.utf8()))
assert function.output_type == pa.int32()
assert function.output_nullable is True
text = repr(function)
assert "Function" in text
assert _FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID in text
for token in ("definition", "source", "packages", "artifact", "digest", "secret"):
assert token not in text.lower()
def test_first_class_function_job_result_sync_wait_returns_exact_function():
bodies = {
"job-register": _first_class_function_job_result_describe_body(
"job-register",
"register_function",
_first_class_function_job_result_function_wire(),
)
}
with mock_lancedb_connection(
_first_class_function_job_result_describe_handler(bodies)
) as db:
result = db.job("job-register").wait()
_assert_exact_first_class_function_job_result(result)
timed_out = db.job("job-register").wait(timeout=timedelta(seconds=5))
_assert_exact_first_class_function_job_result(timed_out)
with pytest.raises(TypeError):
lancedb.Function()
with pytest.raises(AttributeError):
result.id = "mutated"
with pytest.raises(AttributeError):
result.parameters = ()
with pytest.raises(AttributeError):
result.output_type = pa.int64()
with pytest.raises(AttributeError):
result.output_nullable = False
@pytest.mark.asyncio
async def test_first_class_function_job_result_async_wait_returns_exact_function():
bodies = {
"job-register": _first_class_function_job_result_describe_body(
"job-register",
"register_function",
_first_class_function_job_result_function_wire(),
)
}
async with mock_lancedb_connection_async(
_first_class_function_job_result_describe_handler(bodies)
) as db:
result = await db.job("job-register").wait()
_assert_exact_first_class_function_job_result(result)
timed_out = await db.job("job-register").wait(timeout=timedelta(seconds=5))
_assert_exact_first_class_function_job_result(timed_out)
def test_first_class_function_job_result_no_result_wait_returns_none():
bodies = {
"job-index-absent": _first_class_function_job_result_describe_body(
"job-index-absent", "create_index"
),
"job-index-explicit": _first_class_function_job_result_describe_body(
"job-index-explicit",
"create_index",
_first_class_function_job_result_none_wire(),
),
}
with mock_lancedb_connection(
_first_class_function_job_result_describe_handler(bodies)
) as db:
assert db.job("job-index-absent").wait() is None
assert db.job("job-index-explicit").wait(timeout=timedelta(seconds=5)) is None
@pytest.mark.asyncio
async def test_first_class_function_job_result_async_no_result_wait_returns_none():
bodies = {
"job-index-absent": _first_class_function_job_result_describe_body(
"job-index-absent", "create_index"
),
"job-index-explicit": _first_class_function_job_result_describe_body(
"job-index-explicit",
"create_index",
_first_class_function_job_result_none_wire(),
),
}
async with mock_lancedb_connection_async(
_first_class_function_job_result_describe_handler(bodies)
) as db:
assert await db.job("job-index-absent").wait() is None
assert (
await db.job("job-index-explicit").wait(timeout=timedelta(seconds=5))
is None
)
def test_first_class_function_job_result_get_job_result_projection():
bodies = {
"job-register": _first_class_function_job_result_describe_body(
"job-register",
"register_function",
_first_class_function_job_result_function_wire(),
),
"job-absent": _first_class_function_job_result_describe_body(
"job-absent", "create_index"
),
"job-null": _first_class_function_job_result_describe_body(
"job-null",
"create_index",
_FIRST_CLASS_FUNCTION_JOB_RESULT_NULL,
),
"job-explicit-none": _first_class_function_job_result_describe_body(
"job-explicit-none",
"create_index",
_first_class_function_job_result_none_wire(),
),
}
with mock_lancedb_connection(
_first_class_function_job_result_describe_handler(bodies)
) as db:
register_description = db.get_job("job-register")
_assert_exact_first_class_function_job_result(register_description.result)
assert db.get_job("job-absent").result is None
assert db.get_job("job-null").result is None
assert db.get_job("job-explicit-none").result is None
+116
View File
@@ -23,6 +23,7 @@ use lancedb::{
connection::NamespaceClientPushdownOperation,
database::namespace::LanceNamespaceDatabase,
database::{CreateTableMode, Database, ReadConsistency},
function::{FunctionId, RegisterFunctionJobSpec},
};
use pyo3::{
Bound, FromPyObject, Py, PyAny, PyRef, PyResult, Python,
@@ -589,6 +590,121 @@ impl Connection {
})
})
}
/// Submit a first-class Function registration job.
///
/// Accepts the exact private [`crate::function::PyFunctionDefinition`] and
/// builds [`RegisterFunctionJobSpec`] with `expected_current_function_id =
/// None` (create-if-absent). Does not JSON round-trip the definition.
pub fn _register_function<'py>(
self_: PyRef<'py, Self>,
name: String,
definition: Bound<'_, crate::function::PyFunctionDefinition>,
) -> PyResult<Bound<'py, PyAny>> {
let inner = self_.get_inner()?.clone();
let definition = definition.get().inner().clone();
future_into_py(self_.py(), async move {
let spec = RegisterFunctionJobSpec::try_new(name, definition, None).infer_error()?;
let job = inner.register_function(spec).await.infer_error()?;
Ok(crate::job::Job::new(job))
})
}
/// Submit a first-class Function conditional replace job.
///
/// Accepts the observed native [`crate::function::Function`] handle and the
/// exact private [`crate::function::PyFunctionDefinition`], then builds
/// [`RegisterFunctionJobSpec`] with `expected_current_function_id =
/// Some(current.id)`. Reads only `current.inner().id().clone()`. Does not
/// JSON round-trip the definition.
pub fn _replace_function<'py>(
self_: PyRef<'py, Self>,
name: String,
current: Bound<'_, crate::function::Function>,
definition: Bound<'_, crate::function::PyFunctionDefinition>,
) -> PyResult<Bound<'py, PyAny>> {
let inner = self_.get_inner()?.clone();
let definition = definition.get().inner().clone();
let current_id = current.get().inner().id().clone();
future_into_py(self_.py(), async move {
let spec = RegisterFunctionJobSpec::try_new(name, definition, Some(current_id))
.infer_error()?;
let job = inner.register_function(spec).await.infer_error()?;
Ok(crate::job::Job::new(job))
})
}
/// Look up the Function currently bound to a database-scoped name.
///
/// Wraps the exact Rust [`lancedb::function::Function`] once. Empty names
/// fail as [`PyValueError`] before transport via the Rust connection.
pub fn _lookup_function_by_name<'py>(
self_: PyRef<'py, Self>,
name: String,
) -> PyResult<Bound<'py, PyAny>> {
let inner = self_.get_inner()?.clone();
future_into_py(self_.py(), async move {
let function = inner.lookup_function_by_name(&name).await.infer_error()?;
Ok(crate::function::Function::new(function))
})
}
/// Look up an immutable Function by exact opaque Function ID string.
///
/// Constructs [`FunctionId`] with [`FunctionId::try_new`] before dispatch so
/// empty IDs fail as [`PyValueError`] before transport.
pub fn _lookup_function_by_id<'py>(
self_: PyRef<'py, Self>,
function_id: String,
) -> PyResult<Bound<'py, PyAny>> {
let inner = self_.get_inner()?.clone();
future_into_py(self_.py(), async move {
let id = FunctionId::try_new(function_id).infer_error()?;
let function = inner.lookup_function_by_id(&id).await.infer_error()?;
Ok(crate::function::Function::new(function))
})
}
/// Conditionally remove a database-scoped Function catalog name.
///
/// Clones the observed native [`crate::function::Function`] once and
/// delegates to Rust [`lancedb::Connection::remove_function_name`]. Empty
/// names fail as [`PyValueError`] before transport via the Rust connection.
pub fn _remove_function_name<'py>(
self_: PyRef<'py, Self>,
name: String,
current: Bound<'_, crate::function::Function>,
) -> PyResult<Bound<'py, PyAny>> {
let inner = self_.get_inner()?.clone();
let current = current.get().inner().clone();
future_into_py(self_.py(), async move {
inner
.remove_function_name(&name, &current)
.await
.infer_error()?;
// `()` maps to an empty Python tuple via IntoPyObject; return Option
// so the async bridge yields exact Python None.
Ok(None::<()>)
})
}
/// Revoke an exact immutable Function by administrator set-bit.
///
/// Clones the observed native [`crate::function::Function`] once and
/// delegates to Rust [`lancedb::Connection::revoke_function`].
pub fn _revoke_function<'py>(
self_: PyRef<'py, Self>,
function: Bound<'_, crate::function::Function>,
) -> PyResult<Bound<'py, PyAny>> {
let inner = self_.get_inner()?.clone();
let function = function.get().inner().clone();
future_into_py(self_.py(), async move {
inner.revoke_function(&function).await.infer_error()?;
// `()` maps to an empty Python tuple via IntoPyObject; return Option
// so the async bridge yields exact Python None.
Ok(None::<()>)
})
}
}
#[pyfunction]
+15 -2
View File
@@ -102,11 +102,14 @@ impl<T> PythonErrorExt<T> for std::result::Result<T, LanceError> {
err.setattr(intern!(py, "__cause__"), cause_err)?;
Err(PyErr::from_value(err))
}),
LanceError::JobFailed { .. } => Python::attach(|py| {
LanceError::JobFailed { failure, .. } => Python::attach(|py| {
let cls = py
.import(intern!(py, "lancedb.exceptions"))?
.getattr(intern!(py, "JobFailedError"))?;
Err(PyErr::from_value(cls.call1((err.to_string(),))?))
// Structural projection only: failure.error_code.as_str().
// Never infer a code from message, phase, retryable, or source.
let error_code = failure.error_code.as_ref().map(|code| code.as_str());
Err(PyErr::from_value(cls.call1((err.to_string(), error_code))?))
}),
LanceError::JobCancelled { .. } => Python::attach(|py| {
let cls = py
@@ -114,6 +117,16 @@ impl<T> PythonErrorExt<T> for std::result::Result<T, LanceError> {
.getattr(intern!(py, "JobCancelledError"))?;
Err(PyErr::from_value(cls.call1((err.to_string(),))?))
}),
LanceError::Function { code, message } => Python::attach(|py| {
let cls = py
.import(intern!(py, "lancedb.exceptions"))?
.getattr(intern!(py, "FunctionError"))?;
// Structural projection only: code.as_str() + sanitized message.
// Never infer a code from HTTP status or diagnostic text.
Err(PyErr::from_value(
cls.call1((message.as_str(), code.as_str()))?,
))
}),
_ => self.runtime_error(),
},
}
+28 -1
View File
@@ -10,7 +10,7 @@
use std::ops::{Add, Div, Mul, Not, Sub};
use arrow::{datatypes::DataType, pyarrow::PyArrowType};
use datafusion_common::ScalarValue;
use datafusion_common::{Column, ScalarValue};
use lancedb::expr::{
DfExpr, col as ldb_col, contains, expr_cast, is_in, lit as df_lit, lower, upper,
};
@@ -27,6 +27,33 @@ use pyo3::{Bound, PyAny, PyResult, exceptions::PyValueError, prelude::*, pyfunct
#[derive(Clone)]
pub struct PyExpr(pub DfExpr);
/// Crate-private inspection result for Function call authoring (FF-028).
#[derive(Debug, Clone)]
pub(crate) enum DirectExprView<'a> {
/// Direct unqualified DataFusion Column; name is case-sensitive.
UnqualifiedColumn(&'a str),
/// Direct Literal scalar; Arrow type is owned by the scalar value.
Literal(&'a ScalarValue),
}
impl PyExpr {
/// Inspect a direct Column/Literal node for Function call authoring.
///
/// Returns `None` for every other expression shape (arithmetic, cast,
/// scalar function, predicate, alias, qualified column, etc.).
pub(crate) fn as_direct_column_or_literal(&self) -> Option<DirectExprView<'_>> {
match &self.0 {
DfExpr::Column(Column {
relation: None,
name,
..
}) => Some(DirectExprView::UnqualifiedColumn(name.as_str())),
DfExpr::Literal(value, _) => Some(DirectExprView::Literal(value)),
_ => None,
}
}
}
#[pymethods]
impl PyExpr {
// ── comparisons ──────────────────────────────────────────────────────────
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -289,7 +289,7 @@ struct IvfHnswFlatParams {
target_partition_size: Option<u32>,
}
#[pyclass(get_all)]
#[pyclass(module = "lancedb._lancedb", get_all)]
/// A description of an index currently configured on a column
pub struct IndexConfig {
/// The type of the index
+31 -4
View File
@@ -3,6 +3,7 @@
use std::sync::Arc;
use crate::function::Function;
use crate::runtime::future_into_py;
use pyo3::{Bound, PyAny, PyRef, PyResult, pyclass, pymethods};
@@ -21,6 +22,23 @@ impl Job {
}
}
/// Project a Rust [`lancedb::JobResult`] onto the Python success surface.
///
/// Delegates variant interpretation to [`lancedb::JobResult::into_function`]:
/// no nested Function collapses to Python `None`; an exact Function becomes
/// the corresponding [`Function`] handle.
fn project_wait_result(result: lancedb::JobResult) -> Option<Function> {
result.into_function().map(Function::new)
}
/// Project a describe `result` onto Python `Optional[Function]`.
///
/// Rust `None`, `Some(JobResult::None)`, and JSON null all become Python
/// `None`. Only `Some(JobResult::Function)` becomes a [`Function`] handle.
fn project_description_result(result: Option<lancedb::JobResult>) -> Option<Function> {
result.and_then(project_wait_result)
}
#[pymethods]
impl Job {
#[getter]
@@ -39,8 +57,8 @@ impl Job {
pub fn wait(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.inner.clone();
future_into_py(self_.py(), async move {
inner.wait().await.infer_error()?;
Ok(())
let result = inner.wait().await.infer_error()?;
Ok(project_wait_result(result))
})
}
@@ -93,14 +111,16 @@ pub struct JobFailureInfo {
phase: Option<String>,
message: Option<String>,
retryable: Option<bool>,
/// Exact wire `error_code` string when Rust decoded one; never inferred.
error_code: Option<String>,
}
#[pymethods]
impl JobFailureInfo {
fn __repr__(&self) -> String {
format!(
"JobFailureInfo(phase={:?}, message={:?}, retryable={:?})",
self.phase, self.message, self.retryable
"JobFailureInfo(phase={:?}, message={:?}, retryable={:?}, error_code={:?})",
self.phase, self.message, self.retryable, self.error_code
)
}
}
@@ -115,6 +135,7 @@ pub struct JobDescription {
creation_ms: i64,
spec_json: Option<String>,
failure: Option<JobFailureInfo>,
result: Option<Function>,
}
#[pymethods]
@@ -139,7 +160,13 @@ impl From<lancedb::database::JobDescription> for JobDescription {
phase: failure.phase,
message: failure.message,
retryable: failure.retryable,
// Structural projection only: exact as_str(); never infer.
error_code: failure
.error_code
.as_ref()
.map(|code| code.as_str().to_string()),
}),
result: project_description_result(description.result),
}
}
}
+8
View File
@@ -23,6 +23,7 @@ pub mod arrow;
pub mod connection;
pub mod error;
pub mod expr;
pub mod function;
pub mod header;
pub mod index;
pub mod job;
@@ -45,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>()?;
@@ -88,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(())
}
+1 -1
View File
@@ -11,7 +11,7 @@ use pyo3::{PyResult, pyclass, pymethods};
/// Sessions allow you to configure cache sizes for index and metadata caches,
/// which can significantly impact memory use and performance. They can
/// also be re-used across multiple connections to share the same cache state.
#[pyclass(from_py_object)]
#[pyclass(module = "lancedb._lancedb", from_py_object)]
#[derive(Clone)]
pub struct Session {
pub(crate) inner: Arc<LanceSession>,
+142 -2
View File
@@ -26,7 +26,7 @@ use lancedb::table::{
use lancedb::tokenize as lancedb_tokenize;
use pyo3::{
Bound, FromPyObject, Py, PyAny, PyRef, PyResult, Python,
exceptions::{PyRuntimeError, PyValueError},
exceptions::{PyNotImplementedError, PyRuntimeError, PyValueError},
pyclass, pyfunction, pymethods,
types::{IntoPyDict, PyAnyMethods, PyBytes, PyDict, PyDictMethods, PyList, PyListMethods},
};
@@ -579,7 +579,7 @@ impl PyBlobFile {
}
}
#[pyclass(get_all, from_py_object)]
#[pyclass(module = "lancedb._lancedb", get_all, from_py_object)]
#[derive(Clone, Debug)]
pub struct FtsToken {
pub text: String,
@@ -930,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 {
+1 -3
View File
@@ -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 }
@@ -49,8 +50,6 @@ lance-namespace = { workspace = true }
lance-namespace-impls = { workspace = true }
metrics = { workspace = true, optional = true }
metrics-util = { workspace = true, optional = true }
# Pin the GooseFS SDK to the version required by Lance's OpenDAL dependency.
goosefs-sdk = { version = "=0.1.9", optional = true }
moka = { workspace = true }
pin-project = { workspace = true }
tokio = { version = "1.23", features = ["rt-multi-thread", "sync"] }
@@ -136,7 +135,6 @@ azure = [
]
cos = ["lance/tencent", "lance-io/tencent"]
goosefs = [
"dep:goosefs-sdk",
"lance/goosefs",
"lance-io/goosefs",
"lance-namespace-impls/dir-goosefs",
+14 -6
View File
@@ -17,7 +17,7 @@ use arrow_array::builder::LargeBinaryBuilder;
use arrow_schema::{DataType, Field, Schema};
use lance::dataset::{BlobRangeRequest as LanceBlobRangeRequest, Dataset, WriteParams};
use lance_arrow::FieldExt;
use lance_file::version::LanceFileVersion;
use lance_file::version::{ConcreteFileVersion, LanceFileVersion};
use lance_io::object_store::ObjectStore;
use object_store::path::Path;
@@ -333,8 +333,13 @@ pub(crate) fn ensure_blob_storage_version(schema: &Schema, params: &mut WritePar
.data_storage_version
.unwrap_or(LanceFileVersion::Stable)
.resolve();
if resolved < LanceFileVersion::V2_2 {
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 => {}
}
}
@@ -499,7 +504,7 @@ mod tests {
ensure_blob_storage_version(&blob_schema(), &mut params);
assert_eq!(
params.data_storage_version.unwrap().resolve(),
LanceFileVersion::V2_2
LanceFileVersion::V2_2.resolve()
);
}
@@ -512,7 +517,7 @@ mod tests {
ensure_blob_storage_version(&blob_schema(), &mut params);
assert_eq!(
params.data_storage_version.unwrap().resolve(),
LanceFileVersion::V2_2
LanceFileVersion::V2_2.resolve()
);
}
@@ -523,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]
+83
View File
@@ -28,6 +28,7 @@ use crate::database::{
};
use crate::embeddings::{EmbeddingRegistry, MemoryRegistry};
use crate::error::{Error, Result};
use crate::function::{Function, FunctionId, RegisterFunctionJobSpec};
#[cfg(feature = "remote")]
use crate::remote::{
client::ClientConfig,
@@ -550,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
+3 -4
View File
@@ -438,10 +438,9 @@ mod tests {
.await
.unwrap()
.data_storage_format
.lance_file_version()
.unwrap();
// Compare resolved versions since Stable/Next are aliases that resolve at storage time
assert_eq!(storage_format.resolve(), data_storage_version.resolve());
.lance_file_format();
// Compare concrete stored format to the resolved requested alias.
assert_eq!(storage_format, data_storage_version.resolve());
}
#[tokio::test]
+68
View File
@@ -30,12 +30,16 @@ use lance_namespace::models::{
use crate::data::scannable::Scannable;
use crate::error::Result;
use crate::function::{Function, FunctionId, RegisterFunctionJobSpec};
use crate::table::{BaseTable, WriteOptions};
pub mod listing;
pub mod namespace;
pub(crate) mod read_freshness;
#[cfg(test)]
mod create_table_generated_column_schema_admission_contract;
pub trait DatabaseOptions {
fn serialize_into_map(&self, map: &mut HashMap<String, String>);
}
@@ -230,6 +234,12 @@ pub struct JobDescription {
pub creation_ms: i64,
/// The job-type-specific specification. Null when the server omits it.
pub spec: serde_json::Value,
/// Explicit success result from the describe envelope, when present.
///
/// Missing or JSON `null` wire `result` is [`None`]. An explicit
/// [`crate::JobResult::None`] object is `Some(JobResult::None)`. An exact
/// Function result is `Some(JobResult::Function(...))`.
pub result: Option<crate::job::JobResult>,
/// Why the job failed, when the job is failed and the server reports a
/// reason.
pub failure: Option<crate::error::JobFailure>,
@@ -311,6 +321,64 @@ pub trait Database:
async fn job_history(&self, _job_id: Option<&str>) -> Result<Vec<RecordBatch>> {
job_op_not_supported("job_history")
}
/// Submit a first-class Function registration job.
///
/// Returns a [`crate::job::Job`] handle for the accepted server-side job.
/// Local databases do not support registration.
async fn register_function(&self, _spec: RegisterFunctionJobSpec) -> Result<crate::job::Job> {
job_op_not_supported("register_function")
}
/// Look up the Function currently bound to a database-scoped name.
///
/// The name is lookup indirection only and is never part of the returned
/// [`Function`]. Empty names return [`crate::Error::InvalidInput`] before
/// the unsupported fallback so local and remote backends agree. Nonempty
/// names on databases without enterprise catalog lookup return
/// [`crate::Error::NotSupported`].
async fn lookup_function_by_name(&self, name: &str) -> Result<Function> {
// Public nonempty invariant on the Database trait seam itself:
// Connection::database() exposes Arc<dyn Database>, so empty-name
// rejection cannot rely solely on Connection prevalidation.
if name.is_empty() {
return Err(crate::error::Error::InvalidInput {
message: "function lookup name must be non-empty".into(),
});
}
job_op_not_supported("lookup_function_by_name")
}
/// Look up an immutable Function by exact opaque [`FunctionId`].
///
/// Exact-ID lookup is independent of later catalog name changes. Local
/// databases do not support enterprise catalog lookup.
async fn lookup_function_by_id(&self, _function_id: &FunctionId) -> Result<Function> {
job_op_not_supported("lookup_function_by_id")
}
/// Conditionally remove a database-scoped Function catalog name.
///
/// Direct synchronous catalog CAS, not a Job and not physical Function
/// deletion. Empty names return [`crate::Error::InvalidInput`] before the
/// unsupported fallback so local and remote backends agree. Nonempty names
/// on databases without enterprise catalog mutation return
/// [`crate::Error::NotSupported`].
async fn remove_function_name(&self, name: &str, _current: &Function) -> Result<()> {
// Public nonempty invariant on the Database trait seam itself:
// Connection::database() exposes Arc<dyn Database>, so empty-name
// rejection cannot rely solely on Connection prevalidation.
if name.is_empty() {
return Err(crate::error::Error::InvalidInput {
message: "function name removal name must be non-empty".into(),
});
}
job_op_not_supported("remove_function_name")
}
/// Revoke an exact immutable [`Function`] by opaque id.
///
/// Direct synchronous administrator catalog set-bit, not a Job, not name
/// removal, and not physical Function deletion. Databases without
/// enterprise catalog mutation return [`crate::Error::NotSupported`].
async fn revoke_function(&self, _function: &Function) -> Result<()> {
job_op_not_supported("revoke_function")
}
/// Open a table in the database
async fn open_table(&self, request: OpenTableRequest) -> Result<Arc<dyn BaseTable>>;
/// Rename a table in the database
@@ -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");
}
}
+5
View File
@@ -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;
@@ -1038,6 +1039,10 @@ impl Database for ListingDatabase {
}
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;
}
+5
View File
@@ -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
View File
@@ -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()
@@ -92,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)
}
}
+427
View File
@@ -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)
}
}
+154
View File
@@ -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(())
}
}
+293 -13
View File
@@ -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,16 +249,16 @@ 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 {
message: format!("index job task failed: {err}"),
message: format!("job task failed: {err}"),
})),
};
let _ = tx.send(Some(outcome));
@@ -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:?}"),
}
}
}
+2 -1
View File
@@ -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
+1345 -3
View File
File diff suppressed because it is too large Load Diff
+2
View File
@@ -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";
+400 -22
View File
@@ -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"
);
}
}
File diff suppressed because it is too large Load Diff
+270
View File
@@ -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,
}),
}
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+247
View File
@@ -0,0 +1,247 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! Shared Remote HTTP transport primitives for explicit `error_code` classification.
//!
//! These helpers implement the same retry-bucket policy as
//! [`RestfulLanceDbClient::send_with_retry`]: connect → connect budget; timeout /
//! body / decode (including non-success body reads) → read budget; configured
//! retryable status without an explicit nonempty `error_code` → request budget;
//! other client/header errors are immediate. An explicit nonempty top-level
//! `error_code` is terminal [`Error::Function`] and wins over HTTP status.
//! Request/response payload bytes never enter error chains on the classified
//! path.
use reqwest::{RequestBuilder, Response, StatusCode};
use serde_json::Value;
use crate::error::{Error, FunctionErrorCode, Result};
use super::client::{HttpSend, RestfulLanceDbClient};
use super::retry::RetryCounter;
/// Prepare a [`RetryCounter`] with one SDK-generated request id taken from the
/// request builder. The same id is reused across every attempt.
fn prepare_transport_retry<'a, S: HttpSend>(
client: &'a RestfulLanceDbClient<S>,
req_builder: &RequestBuilder,
) -> Result<RetryCounter<'a>> {
let tmp_req = req_builder.try_clone().ok_or_else(|| Error::Runtime {
message: "Attempted to retry a request that cannot be cloned".to_string(),
})?;
let (_, built) = tmp_req.build_split();
let mut built = built.map_err(|e| Error::Runtime {
message: format!("Failed to build request: {}", e),
})?;
let request_id = client.extract_request_id(&mut built);
Ok(RetryCounter::new(&client.retry_config, request_id))
}
/// Classify a send-attempt error using the same buckets as `send_with_retry`.
///
/// Returns `Ok(())` when the caller should sleep and retry. Returns `Err` for
/// nonretryable failures or when a retry budget is exhausted (no extra attempt).
fn classify_transport_send_error(retry_counter: &mut RetryCounter<'_>, err: Error) -> Result<()> {
match err {
Error::Http {
source,
request_id,
status_code,
} => match source.downcast::<reqwest::Error>() {
Ok(reqwest_err) if reqwest_err.is_connect() => {
retry_counter.increment_connect_failures(*reqwest_err)
}
Ok(reqwest_err)
if reqwest_err.is_timeout() || reqwest_err.is_body() || reqwest_err.is_decode() =>
{
retry_counter.increment_read_failures(*reqwest_err)
}
Ok(reqwest_err) => Err(Error::Http {
source: Box::new(*reqwest_err),
request_id,
status_code,
}),
Err(source) => Err(Error::Http {
source,
request_id,
status_code,
}),
},
// Header-provider / client failures are not transport retries.
other => Err(other),
}
}
/// Decode a stable category only from an explicit nonempty string `error_code`.
/// Missing, empty, wrong-type, nested-only, or non-JSON bodies yield [`None`].
pub(super) fn explicit_error_code(bytes: &[u8]) -> Option<FunctionErrorCode> {
let value: Value = serde_json::from_slice(bytes).ok()?;
let code = value.get("error_code")?;
match code {
Value::String(raw) if !raw.is_empty() => {
serde_json::from_value(Value::String(raw.clone())).ok()
}
_ => None,
}
}
/// Send a clonable request with explicit-`error_code` classification.
///
/// Uses the sensitive attempt path (no request body/header logging). Successful
/// 2xx responses are returned **unconsumed** so callers can run their own
/// success decoders. Non-success bodies are inspected for an explicit nonempty
/// `error_code` before status-based retry. Payload bytes never enter
/// [`Error::Function`], [`Error::Http`], or exhausted [`Error::Retry`] chains.
pub(super) async fn send_with_explicit_error_code<S: HttpSend>(
client: &RestfulLanceDbClient<S>,
req_builder: RequestBuilder,
http_error_message: &'static str,
function_error_message: &'static str,
) -> Result<(String, Response)> {
let mut retry_counter = prepare_transport_retry(client, &req_builder)?;
loop {
let attempt = req_builder.try_clone().ok_or_else(|| Error::Runtime {
message: "Attempted to retry a request that cannot be cloned".to_string(),
})?;
let rsp = match client
.send_attempt_with_request_id(attempt, &retry_counter.request_id)
.await
{
Ok(rsp) => rsp,
Err(err) => {
classify_transport_send_error(&mut retry_counter, err)?;
tokio::time::sleep(retry_counter.next_sleep_time()).await;
continue;
}
};
let status = rsp.status();
if status.is_success() {
// Leave the body unconsumed for Arrow / plan success decoders.
return Ok((retry_counter.request_id.clone(), rsp));
}
// Inspect the body before deciding whether the status is retryable.
let bytes = match rsp.bytes().await {
Ok(bytes) => bytes,
Err(err) => {
// Response body/decode failures share the read budget with
// send-time timeout/body/decode errors.
retry_counter.increment_read_failures(err)?;
tokio::time::sleep(retry_counter.next_sleep_time()).await;
continue;
}
};
if let Some(code) = explicit_error_code(&bytes) {
return Err(Error::Function {
code,
message: function_error_message.to_string(),
});
}
if client.retry_config.statuses.contains(&status) {
let source = Error::Http {
source: http_error_message.into(),
request_id: retry_counter.request_id.clone(),
status_code: Some(status),
};
retry_counter.increment_request_failures(source)?;
tokio::time::sleep(retry_counter.next_sleep_time()).await;
continue;
}
return Err(Error::Http {
source: http_error_message.into(),
request_id: retry_counter.request_id.clone(),
status_code: Some(status),
});
}
}
/// Decision before reading response bytes (catalog mutations / lookup).
pub(super) enum BeforeBody<T> {
/// Finish without reading or interpreting any body (HTTP 204 mutations).
Done(Result<T>),
/// Read bytes and continue classification.
ReadBody,
}
/// How a classified body is treated after a successful read.
pub(super) enum BodyAction<T> {
/// Terminal success or failure for this attempt.
Done(Result<T>),
/// Configured retryable status without an explicit `error_code`: consume
/// the request budget and retry with the same request id and body.
RetryRequest,
}
/// Shared POST retry loop used by Function-catalog helpers.
///
/// Sensitive attempt sending logs no body/header selectors. One SDK-generated
/// request id and the exact cloned JSON body are reused across attempts.
/// Callers supply before/after body classifiers; this loop owns budgets only.
pub(super) async fn post_with_body_classification<S, Before, After, T>(
client: &RestfulLanceDbClient<S>,
req_builder: RequestBuilder,
http_error_message: &'static str,
mut before_body: Before,
mut after_body: After,
) -> Result<T>
where
S: HttpSend,
Before: FnMut(StatusCode, &str) -> BeforeBody<T>,
After: FnMut(StatusCode, &[u8], String) -> BodyAction<T>,
{
let mut retry_counter = prepare_transport_retry(client, &req_builder)?;
loop {
let attempt = req_builder.try_clone().ok_or_else(|| Error::Runtime {
message: "Attempted to retry a request that cannot be cloned".to_string(),
})?;
let rsp = match client
.send_attempt_with_request_id(attempt, &retry_counter.request_id)
.await
{
Ok(rsp) => rsp,
Err(err) => {
classify_transport_send_error(&mut retry_counter, err)?;
tokio::time::sleep(retry_counter.next_sleep_time()).await;
continue;
}
};
let status = rsp.status();
match before_body(status, &retry_counter.request_id) {
BeforeBody::Done(result) => return result,
BeforeBody::ReadBody => {}
}
let bytes = match rsp.bytes().await {
Ok(bytes) => bytes,
Err(err) => {
// Response body/decode failures share the read budget with
// send-time timeout/body/decode errors.
retry_counter.increment_read_failures(err)?;
tokio::time::sleep(retry_counter.next_sleep_time()).await;
continue;
}
};
match after_body(status, &bytes, retry_counter.request_id.clone()) {
BodyAction::Done(result) => return result,
BodyAction::RetryRequest => {
let source = Error::Http {
source: http_error_message.into(),
request_id: retry_counter.request_id.clone(),
status_code: Some(status),
};
retry_counter.increment_request_failures(source)?;
tokio::time::sleep(retry_counter.next_sleep_time()).await;
}
}
}
}
+1243 -8
View File
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,441 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! RED runtime contract tests for Native add-columns schema admission (B4h).
//!
//! Caller-authored Arrow field metadata under
//! [`crate::function::GENERATED_COLUMN_METADATA_KEY`] must not enter table
//! schema state through general-purpose Native `add_columns`. Generated
//! definitions are Job-owned. Schema-bearing transforms (`BatchUDF`, `Stream`,
//! `Reader`, `AllNulls`) currently accept and persist reserved top-level field
//! metadata; these tests pin the missing pre-consumption admission guard.
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use arrow_array::{Int32Array, RecordBatch, RecordBatchIterator, RecordBatchReader, StringArray};
use arrow_schema::{ArrowError, DataType, Field, Schema, SchemaRef};
use datafusion_physical_plan::stream::RecordBatchStreamAdapter;
use futures::{TryStreamExt, stream};
use lance::dataset::{BatchUDF, NewColumnTransform};
use tempfile::TempDir;
use crate::connection::ConnectBuilder;
use crate::error::Error;
use crate::function::{
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition,
};
use crate::query::{ExecutableQuery, QueryBase, Select};
use crate::table::Table;
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.b4h.add_columns.literal";
const MALFORMED_MARKER: &str = "SENSITIVE_B4H_ADD_COLUMNS_METADATA_MARKER_4f8a_c3e2";
struct Fixture {
_tmp: TempDir,
table: Table,
}
/// Counts [`RecordBatchReader::next`] calls. [`RecordBatchReader::schema`] is free.
struct ObservableReader {
inner: Box<dyn RecordBatchReader + Send>,
next_calls: Arc<AtomicUsize>,
}
impl ObservableReader {
fn wrap(
inner: Box<dyn RecordBatchReader + Send>,
next_calls: Arc<AtomicUsize>,
) -> Box<dyn RecordBatchReader + Send> {
Box::new(Self { inner, next_calls })
}
}
impl Iterator for ObservableReader {
type Item = Result<RecordBatch, ArrowError>;
fn next(&mut self) -> Option<Self::Item> {
self.next_calls.fetch_add(1, Ordering::SeqCst);
self.inner.next()
}
}
impl RecordBatchReader for ObservableReader {
fn schema(&self) -> SchemaRef {
self.inner.schema()
}
}
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 seed_batch() -> RecordBatch {
let schema = Arc::new(Schema::new(vec![
Field::new(ID, DataType::Int32, false),
Field::new(ORDINARY, DataType::Utf8, true),
]));
RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from(vec![1, 2])),
Arc::new(StringArray::from(vec![Some("a"), Some("b")])),
],
)
.unwrap()
}
fn field_with_metadata(metadata: HashMap<String, String>) -> Field {
Field::new(GEN_OUT, DataType::Int32, true).with_metadata(metadata)
}
fn reserved_field(payload: &str) -> Field {
field_with_metadata(
[(
GENERATED_COLUMN_METADATA_KEY.to_string(),
payload.to_string(),
)]
.into(),
)
}
fn ordinary_metadata_field() -> Field {
field_with_metadata(
[(
ORDINARY_META_KEY.to_string(),
ORDINARY_META_VALUE.to_string(),
)]
.into(),
)
}
fn values_batch(schema: SchemaRef, values: Vec<i32>) -> RecordBatch {
RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(values))]).unwrap()
}
fn boxed_reader(batch: RecordBatch) -> Box<dyn RecordBatchReader + Send> {
let schema = batch.schema();
Box::new(RecordBatchIterator::new(
vec![Ok(batch)].into_iter(),
schema,
))
}
fn observable_stream(
batch: RecordBatch,
yield_calls: Arc<AtomicUsize>,
) -> datafusion_physical_plan::SendableRecordBatchStream {
let schema = batch.schema();
let counter = yield_calls.clone();
Box::pin(RecordBatchStreamAdapter::new(
schema,
stream::once(async move {
counter.fetch_add(1, Ordering::SeqCst);
Ok(batch)
}),
))
}
fn assert_not_supported_redacted(err: &Error, label: &str, payload: &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(payload),
"{label}: leaked raw payload: {rendered}"
);
assert!(
!rendered.contains(FN_ID),
"{label}: leaked Function ID: {rendered}"
);
assert!(
!rendered.contains(GEN_OUT),
"{label}: leaked output field name: {rendered}"
);
assert!(
!rendered.contains(MALFORMED_MARKER),
"{label}: leaked malformed marker: {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 create_table(name: &str) -> Fixture {
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap().to_string();
let conn = ConnectBuilder::new(&uri).execute().await.unwrap();
let table = conn
.create_table(name, seed_batch())
.execute()
.await
.unwrap();
Fixture { _tmp: tmp, table }
}
async fn snapshot_rows(table: &Table) -> Vec<(i32, String)> {
let batches: Vec<RecordBatch> = table
.query()
.select(Select::columns(&[ID, ORDINARY]))
.execute()
.await
.unwrap()
.try_collect()
.await
.unwrap();
let mut rows = Vec::new();
for batch in batches {
let ids = batch
.column(0)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap();
let ordinary = batch
.column(1)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
for i in 0..batch.num_rows() {
rows.push((ids.value(i), ordinary.value(i).to_string()));
}
}
rows.sort_by_key(|(id, _)| *id);
rows
}
async fn assert_table_unchanged(
table: &Table,
version_before: u64,
schema_before: &Schema,
rows_before: &[(i32, String)],
) {
assert_eq!(table.version().await.unwrap(), version_before);
let schema_after = table.schema().await.unwrap();
assert_eq!(schema_after.as_ref(), schema_before);
assert!(
schema_after.field_with_name(GEN_OUT).is_err(),
"rejected add_columns must leave column `{GEN_OUT}` absent"
);
assert_eq!(snapshot_rows(table).await, rows_before);
}
#[tokio::test]
async fn batch_udf_rejects_valid_reserved_before_mapper() {
let fixture = create_table("b4h_batch_udf").await;
let table = &fixture.table;
let version_before = table.version().await.unwrap();
let schema_before = table.schema().await.unwrap();
let rows_before = snapshot_rows(table).await;
let payload = valid_reserved_payload();
let output_schema = Arc::new(Schema::new(vec![reserved_field(&payload)]));
let mapper_schema = output_schema.clone();
let mapper_calls = Arc::new(AtomicUsize::new(0));
let calls = mapper_calls.clone();
let udf = BatchUDF {
mapper: Box::new(move |batch: &RecordBatch| {
calls.fetch_add(1, Ordering::SeqCst);
let values = Int32Array::from(vec![Some(10); batch.num_rows()]);
Ok(RecordBatch::try_new(
mapper_schema.clone(),
vec![Arc::new(values)],
)?)
}),
output_schema,
result_checkpoint: None,
};
// Public Table::add_columns builder path.
let err = table
.add_columns()
.transform(NewColumnTransform::BatchUDF(udf))
.execute()
.await
.expect_err("BatchUDF must reject reserved generated-column metadata");
assert_not_supported_redacted(&err, "BatchUDF reserved admission", &payload);
assert_eq!(
mapper_calls.load(Ordering::SeqCst),
0,
"rejection must occur before invoking the BatchUDF mapper"
);
assert_table_unchanged(table, version_before, schema_before.as_ref(), &rows_before).await;
}
#[tokio::test]
async fn stream_rejects_malformed_reserved_before_yield() {
let fixture = create_table("b4h_stream").await;
let table = &fixture.table;
let version_before = table.version().await.unwrap();
let schema_before = table.schema().await.unwrap();
let rows_before = snapshot_rows(table).await;
let payload = malformed_reserved_payload();
assert!(payload.contains(MALFORMED_MARKER));
let output_schema = Arc::new(Schema::new(vec![reserved_field(&payload)]));
let yield_calls = Arc::new(AtomicUsize::new(0));
let stream = observable_stream(
values_batch(output_schema, vec![10, 20]),
yield_calls.clone(),
);
// Direct experimental BaseTable::add_columns path.
let err = table
.base_table()
.add_columns(NewColumnTransform::Stream(stream), None)
.await
.expect_err("Stream must reject reserved generated-column metadata");
assert_not_supported_redacted(&err, "Stream reserved admission", &payload);
assert_eq!(
yield_calls.load(Ordering::SeqCst),
0,
"rejection must occur before polling/yielding the user Stream"
);
assert_table_unchanged(table, version_before, schema_before.as_ref(), &rows_before).await;
}
#[tokio::test]
async fn reader_rejects_valid_reserved_before_next() {
let fixture = create_table("b4h_reader").await;
let table = &fixture.table;
let version_before = table.version().await.unwrap();
let schema_before = table.schema().await.unwrap();
let rows_before = snapshot_rows(table).await;
let payload = valid_reserved_payload();
let output_schema = Arc::new(Schema::new(vec![reserved_field(&payload)]));
let next_calls = Arc::new(AtomicUsize::new(0));
let reader = ObservableReader::wrap(
boxed_reader(values_batch(output_schema, vec![10, 20])),
next_calls.clone(),
);
let err = table
.base_table()
.add_columns(NewColumnTransform::Reader(reader), None)
.await
.expect_err("Reader must reject reserved generated-column metadata");
assert_not_supported_redacted(&err, "Reader reserved admission", &payload);
assert_eq!(
next_calls.load(Ordering::SeqCst),
0,
"rejection must occur before RecordBatchReader::next"
);
assert_table_unchanged(table, version_before, schema_before.as_ref(), &rows_before).await;
}
#[tokio::test]
async fn all_nulls_rejects_malformed_reserved_before_commit() {
let fixture = create_table("b4h_all_nulls").await;
let table = &fixture.table;
let version_before = table.version().await.unwrap();
let schema_before = table.schema().await.unwrap();
let rows_before = snapshot_rows(table).await;
let payload = malformed_reserved_payload();
let output_schema = Arc::new(Schema::new(vec![reserved_field(&payload)]));
let err = table
.add_columns()
.transform(NewColumnTransform::AllNulls(output_schema))
.execute()
.await
.expect_err("AllNulls must reject reserved generated-column metadata");
assert_not_supported_redacted(&err, "AllNulls reserved admission", &payload);
assert_table_unchanged(table, version_before, schema_before.as_ref(), &rows_before).await;
}
#[tokio::test]
async fn sql_expressions_add_columns_still_succeeds() {
let fixture = create_table("b4h_sql_control").await;
let table = &fixture.table;
table
.add_columns()
.transform(NewColumnTransform::SqlExpressions(vec![(
"doubled".into(),
"id * 2".into(),
)]))
.execute()
.await
.expect("ordinary SqlExpressions add_columns must remain supported");
let schema = table.schema().await.unwrap();
assert!(schema.field_with_name("doubled").is_ok());
assert!(schema.field_with_name(GEN_OUT).is_err());
assert!(
!schema
.field_with_name("doubled")
.unwrap()
.metadata()
.contains_key(GENERATED_COLUMN_METADATA_KEY)
);
}
#[tokio::test]
async fn schema_bearing_ordinary_metadata_is_preserved() {
let fixture = create_table("b4h_ordinary_meta").await;
let table = &fixture.table;
let output_schema = Arc::new(Schema::new(vec![ordinary_metadata_field()]));
// AllNulls is schema-bearing and metadata-only; proves ordinary metadata
// remains accepted so a later guard cannot reject every field metadata map.
table
.add_columns()
.transform(NewColumnTransform::AllNulls(output_schema))
.execute()
.await
.expect("ordinary non-reserved field metadata must remain accepted");
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));
}
+23 -6
View File
@@ -10,6 +10,7 @@ use serde::{Deserialize, Serialize};
use crate::data::scannable::Scannable;
use crate::data::scannable::scannable_with_embeddings;
use crate::embeddings::EmbeddingRegistry;
use crate::function::schema_admission::reject_caller_authored_generated_column_schema_on_overwrite;
use crate::table::datafusion::cast::cast_to_table_schema;
use crate::table::datafusion::reject_nan::reject_nan_vectors;
use crate::table::datafusion::scannable_exec::ScannableExec;
@@ -151,6 +152,27 @@ impl AddDataBuilder {
self.parent.clone().add(self).await
}
/// Effective overwrite for schema-replacement admission and planning.
///
/// True when either `WriteOptions.lance_write_params.mode` is
/// [`WriteMode::Overwrite`] or [`AddDataMode::Overwrite`] is set.
pub(crate) fn is_effective_overwrite(&self) -> bool {
self.write_options
.lance_write_params
.as_ref()
.is_some_and(|p| matches!(p.mode, WriteMode::Overwrite))
|| matches!(self.mode, AddDataMode::Overwrite)
}
/// Borrowed preflight: reject reserved generated-column metadata on
/// effective overwrite before source scan or table/network work.
pub(crate) fn admit_input_schema(&self) -> Result<()> {
reject_caller_authored_generated_column_schema_on_overwrite(
self.data.schema().as_ref(),
self.is_effective_overwrite(),
)
}
/// Build a DataFusion execution plan that applies embeddings, casts data to
/// the table schema, and optionally rejects NaN vectors.
///
@@ -161,12 +183,7 @@ impl AddDataBuilder {
table_schema: &Schema,
table_def: &TableDefinition,
) -> Result<PreprocessingOutput> {
let overwrite = self
.write_options
.lance_write_params
.as_ref()
.is_some_and(|p| matches!(p.mode, WriteMode::Overwrite))
|| matches!(self.mode, AddDataMode::Overwrite);
let overwrite = self.is_effective_overwrite();
if !overwrite {
validate_schema(&self.data.schema(), table_schema)?;
@@ -0,0 +1,797 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! RED runtime contract tests for Native append invalidation (B4b).
//!
//! These tests pin Native Table API and DataFusion SQL insert behavior for
//! generated-column dependency-epoch invalidation. They use real local Native
//! tables and existing public/internal APIs; Lance commits and query guards are
//! not mocked.
use std::collections::HashSet;
use std::sync::Arc;
use std::time::Duration;
use arrow_array::{Array, Int32Array, RecordBatch, RecordBatchIterator, StringArray};
use arrow_schema::{DataType, Field, Schema};
use datafusion::prelude::SessionContext;
use futures::TryStreamExt;
use lance::dataset::{WriteMode, WriteParams};
use tempfile::TempDir;
use crate::connection::ConnectBuilder;
use crate::error::{Error, FunctionErrorCode};
use crate::function::{
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition,
GeneratedColumnStatus,
};
use crate::query::{ExecutableQuery, QueryBase, Select};
use crate::table::datafusion::BaseTableAdapter;
use crate::table::{AddDataMode, Table, WriteOptions};
const GEN_OUT: &str = "gen_out";
const ORDINARY: &str = "ordinary";
const INITIAL_DEPENDENCY_EPOCH: u64 = 3;
const INITIAL_MATERIALIZED_EPOCH: u64 = 3;
const MALFORMED_MARKER: &str = "SENSITIVE_B4B_APPEND_METADATA_MARKER_9f2c_a81d";
struct Fixture {
_tmp: TempDir,
table: Table,
table_name: String,
uri: String,
}
fn literal_only_function() -> Function {
Function::new(
FunctionId::try_new("fn.exact.b4b.append.literal").unwrap(),
FunctionSignature::try_new(
vec![FunctionParameter::new("label", DataType::Utf8)],
FunctionOutput::new(DataType::Int32, true),
)
.unwrap(),
)
}
fn literal_only_definition(
output_field_id: i32,
dependency_epoch: u64,
materialized_epoch: u64,
) -> GeneratedColumnDefinition {
let function = literal_only_function();
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();
assert!(
call.arguments()
.iter()
.all(|(_, argument)| argument.field_id().is_none()),
"fixture must be literal-only so row-set coverage, not field dependency, drives invalidation"
);
GeneratedColumnDefinition::try_new(output_field_id, call, dependency_epoch, materialized_epoch)
.unwrap()
}
async fn create_table_with_complete_literal_generated(name: &str) -> Fixture {
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap().to_string();
let conn = ConnectBuilder::new(&uri).execute().await.unwrap();
let schema = Arc::new(Schema::new(vec![
Field::new(GEN_OUT, DataType::Int32, true),
Field::new(ORDINARY, DataType::Utf8, true),
]));
let batch = RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from(vec![1])),
Arc::new(StringArray::from(vec![Some("seed")])),
],
)
.unwrap();
let table = conn.create_table(name, batch).execute().await.unwrap();
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
let field_id = snapshot.field(GEN_OUT).expect(GEN_OUT).field_id();
let definition = literal_only_definition(
field_id,
INITIAL_DEPENDENCY_EPOCH,
INITIAL_MATERIALIZED_EPOCH,
);
let json = definition.to_metadata_json().unwrap();
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
table
.as_native()
.expect("generated-column fixture planting requires a Native table"),
GEN_OUT,
json,
)
.await
.unwrap();
assert_eq!(
table.generated_column_status(GEN_OUT).await.unwrap(),
GeneratedColumnStatus::Complete
);
let planted = read_generated_definition(&table).await;
assert!(
planted
.function_call()
.arguments()
.iter()
.all(|(_, argument)| argument.field_id().is_none()),
"planted metadata must remain literal-only"
);
Fixture {
_tmp: tmp,
table,
table_name: name.to_string(),
uri,
}
}
async fn create_ordinary_table(name: &str) -> Fixture {
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap().to_string();
let conn = ConnectBuilder::new(&uri).execute().await.unwrap();
let schema = Arc::new(Schema::new(vec![
Field::new(GEN_OUT, DataType::Int32, true),
Field::new(ORDINARY, DataType::Utf8, true),
]));
let batch = RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from(vec![1])),
Arc::new(StringArray::from(vec![Some("seed")])),
],
)
.unwrap();
let table = conn.create_table(name, batch).execute().await.unwrap();
Fixture {
_tmp: tmp,
table,
table_name: name.to_string(),
uri,
}
}
async fn read_generated_definition(table: &Table) -> GeneratedColumnDefinition {
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
snapshot
.field(GEN_OUT)
.expect(GEN_OUT)
.generated_column_definition()
.expect("generated metadata must decode")
.expect("generated metadata must be present")
}
fn ordinary_rows_batch(values: &[&str]) -> RecordBatch {
let schema = Arc::new(Schema::new(vec![Field::new(
ORDINARY,
DataType::Utf8,
true,
)]));
RecordBatch::try_new(
schema,
vec![Arc::new(StringArray::from(
values.iter().map(|value| Some(*value)).collect::<Vec<_>>(),
))],
)
.unwrap()
}
fn full_rows_batch(gen_values: &[Option<i32>], ordinary_values: &[&str]) -> RecordBatch {
assert_eq!(gen_values.len(), ordinary_values.len());
let schema = Arc::new(Schema::new(vec![
Field::new(GEN_OUT, DataType::Int32, true),
Field::new(ORDINARY, DataType::Utf8, true),
]));
RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from(gen_values.to_vec())),
Arc::new(StringArray::from(
ordinary_values
.iter()
.map(|value| Some(*value))
.collect::<Vec<_>>(),
)),
],
)
.unwrap()
}
async fn ordinary_values(table: &Table) -> HashSet<String> {
let batches = table
.query()
.select(Select::columns(&[ORDINARY]))
.execute()
.await
.unwrap()
.try_collect::<Vec<_>>()
.await
.unwrap();
let mut values = HashSet::new();
for batch in batches {
let column = batch
.column(0)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
for index in 0..column.len() {
if !column.is_null(index) {
values.insert(column.value(index).to_string());
}
}
}
values
}
fn assert_generated_column_incomplete(err: &Error, label: &str) {
match err {
Error::Function {
code: FunctionErrorCode::GeneratedColumnIncomplete,
..
} => {}
other => panic!("{label}: expected generated_column_incomplete, got {other:?}"),
}
}
fn assert_not_supported(err: &Error, label: &str) {
assert!(
matches!(err, Error::NotSupported { .. }),
"{label}: expected NotSupported, got {err:?}"
);
}
fn assert_invalid_input_redacted(err: &Error, label: &str) {
assert!(
matches!(err, Error::InvalidInput { .. }),
"{label}: expected InvalidInput, got {err:?}"
);
let rendered = format!("{err}\n{err:?}");
assert!(
!rendered.contains(MALFORMED_MARKER),
"{label}: diagnostic echoed unique metadata marker: {rendered}"
);
assert!(
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
"{label}: diagnostic echoed metadata wire key: {rendered}"
);
}
fn assert_conflict_error(err: &Error, label: &str) {
match err {
Error::Lance { source } => {
assert!(
matches!(
source,
lance::Error::IncompatibleTransaction { .. }
| lance::Error::RetryableCommitConflict { .. }
| lance::Error::CommitConflict { .. }
),
"{label}: expected Lance commit conflict category, got {source:?}"
);
}
Error::Function {
code: FunctionErrorCode::StaleOrConflictingInput,
..
} => {}
other => panic!("{label}: expected conflict error category, got {other:?}"),
}
}
fn from_datafusion_error(err: datafusion_common::DataFusionError) -> Error {
Error::from(err)
}
async fn sql_ctx_for(table: &Table, name: &str) -> SessionContext {
let ctx = SessionContext::new();
let provider = BaseTableAdapter::try_new(table.base_table().clone())
.await
.unwrap();
ctx.register_table(name, Arc::new(provider)).unwrap();
ctx
}
async fn run_sql(ctx: &SessionContext, sql: &str) -> Result<(), Error> {
match ctx.sql(sql).await {
Err(err) => Err(from_datafusion_error(err)),
Ok(df) => match df.collect().await {
Ok(_) => Ok(()),
Err(err) => Err(from_datafusion_error(err)),
},
}
}
#[tokio::test]
async fn nonempty_table_api_append_invalidates_literal_only_generated_column() {
let fixture = create_table_with_complete_literal_generated("b4b_table_append").await;
let before = read_generated_definition(&fixture.table).await;
fixture
.table
.add(ordinary_rows_batch(&["appended"]))
.execute()
.await
.expect("non-empty Table API append must commit");
let values = ordinary_values(&fixture.table).await;
assert!(values.contains("seed"));
assert!(values.contains("appended"));
assert_eq!(
fixture
.table
.generated_column_status(GEN_OUT)
.await
.unwrap(),
GeneratedColumnStatus::Incomplete
);
let after = read_generated_definition(&fixture.table).await;
assert_eq!(after.dependency_epoch(), before.dependency_epoch() + 1);
assert_eq!(after.materialized_epoch(), before.materialized_epoch());
assert_eq!(after.output_field_id(), before.output_field_id());
assert_eq!(after.function_call(), before.function_call());
let Err(err) = fixture
.table
.query()
.select(Select::columns(&[GEN_OUT]))
.execute()
.await
else {
panic!("incomplete generated column query must fail");
};
assert_generated_column_incomplete(&err, "table api append query");
}
#[tokio::test]
async fn nonempty_table_api_append_atomic_version_visibility() {
let fixture = create_table_with_complete_literal_generated("b4b_atomic_visibility").await;
let previous_version = fixture.table.version().await.unwrap();
let previous_rows = ordinary_values(&fixture.table).await;
let previous_definition = read_generated_definition(&fixture.table).await;
assert_eq!(
previous_definition.status(),
GeneratedColumnStatus::Complete
);
fixture
.table
.add(ordinary_rows_batch(&["atomic-new"]))
.execute()
.await
.expect("non-empty append must commit");
let new_version = fixture.table.version().await.unwrap();
assert_ne!(new_version, previous_version);
// Exact new version: new rows + incomplete metadata together.
let new_rows = ordinary_values(&fixture.table).await;
assert!(new_rows.contains("atomic-new"));
assert_eq!(
fixture
.table
.generated_column_status(GEN_OUT)
.await
.unwrap(),
GeneratedColumnStatus::Incomplete
);
let new_definition = read_generated_definition(&fixture.table).await;
assert_eq!(
new_definition.dependency_epoch(),
previous_definition.dependency_epoch() + 1
);
// Immediately previous version: neither new rows nor incomplete metadata.
fixture.table.checkout(previous_version).await.unwrap();
assert_eq!(ordinary_values(&fixture.table).await, previous_rows);
assert!(!ordinary_values(&fixture.table).await.contains("atomic-new"));
assert_eq!(
fixture
.table
.generated_column_status(GEN_OUT)
.await
.unwrap(),
GeneratedColumnStatus::Complete
);
let checked_out = read_generated_definition(&fixture.table).await;
assert_eq!(checked_out, previous_definition);
fixture
.table
.query()
.select(Select::columns(&[GEN_OUT]))
.execute()
.await
.expect("previous complete version must remain readable");
}
#[tokio::test]
async fn empty_table_api_append_leaves_complete_generated_column() {
let fixture = create_table_with_complete_literal_generated("b4b_empty_table_append").await;
let before = read_generated_definition(&fixture.table).await;
let rows_before = ordinary_values(&fixture.table).await;
fixture
.table
.add(RecordBatch::new_empty(Arc::new(Schema::new(vec![
Field::new(ORDINARY, DataType::Utf8, true),
]))))
.execute()
.await
.expect("empty Table API append is a supported path");
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
assert_eq!(
fixture
.table
.generated_column_status(GEN_OUT)
.await
.unwrap(),
GeneratedColumnStatus::Complete
);
let after = read_generated_definition(&fixture.table).await;
assert_eq!(after, before);
}
#[tokio::test]
async fn multipartition_table_api_append_advances_dependency_epoch_once() {
let fixture = create_table_with_complete_literal_generated("b4b_multipartition").await;
let before = read_generated_definition(&fixture.table).await;
fixture
.table
.add(ordinary_rows_batch(&["p0", "p1", "p2", "p3"]))
.write_parallelism(2)
.execute()
.await
.expect("multi-partition append must commit");
let values = ordinary_values(&fixture.table).await;
assert!(values.contains("p0"));
assert!(values.contains("p3"));
assert_eq!(
fixture
.table
.generated_column_status(GEN_OUT)
.await
.unwrap(),
GeneratedColumnStatus::Incomplete
);
let after = read_generated_definition(&fixture.table).await;
assert_eq!(
after.dependency_epoch(),
before.dependency_epoch() + 1,
"multi-partition append must attach one whole-transaction patch"
);
assert_eq!(after.materialized_epoch(), before.materialized_epoch());
assert_eq!(after.function_call(), before.function_call());
}
#[tokio::test]
async fn table_api_overwrite_rejects_before_mutation_when_generated_column_present() {
let fixture = create_table_with_complete_literal_generated("b4b_table_overwrite").await;
let version_before = fixture.table.version().await.unwrap();
let rows_before = ordinary_values(&fixture.table).await;
let definition_before = read_generated_definition(&fixture.table).await;
let err = fixture
.table
.add(full_rows_batch(&[Some(9)], &["overwrite"]))
.mode(AddDataMode::Overwrite)
.execute()
.await
.expect_err("overwrite must reject when any generated column is present");
assert_not_supported(&err, "table api overwrite");
assert_eq!(fixture.table.version().await.unwrap(), version_before);
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
assert_eq!(
read_generated_definition(&fixture.table).await,
definition_before
);
assert_eq!(
fixture
.table
.generated_column_status(GEN_OUT)
.await
.unwrap(),
GeneratedColumnStatus::Complete
);
}
#[tokio::test]
async fn table_api_effective_overwrite_from_add_data_mode_rejects_when_lance_params_append() {
let fixture =
create_table_with_complete_literal_generated("b4b_table_effective_overwrite").await;
let version_before = fixture.table.version().await.unwrap();
let rows_before = ordinary_values(&fixture.table).await;
let definition_before = read_generated_definition(&fixture.table).await;
assert_eq!(definition_before.status(), GeneratedColumnStatus::Complete);
let err = fixture
.table
.add(full_rows_batch(&[Some(9)], &["effective-overwrite"]))
.mode(AddDataMode::Overwrite)
.write_options(WriteOptions {
lance_write_params: Some(WriteParams {
mode: WriteMode::Append,
..Default::default()
}),
})
.execute()
.await
.expect_err(
"AddDataMode::Overwrite must reject generated-table writes even when \
explicit lance WriteParams.mode is Append",
);
assert_not_supported(&err, "table api effective overwrite");
assert_eq!(fixture.table.version().await.unwrap(), version_before);
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
assert_eq!(
read_generated_definition(&fixture.table).await,
definition_before
);
assert_eq!(
fixture
.table
.generated_column_status(GEN_OUT)
.await
.unwrap(),
GeneratedColumnStatus::Complete
);
}
#[tokio::test]
async fn ordinary_table_api_overwrite_still_supported() {
let fixture = create_ordinary_table("b4b_ordinary_overwrite_control").await;
fixture
.table
.add(full_rows_batch(&[Some(42)], &["replaced"]))
.mode(AddDataMode::Overwrite)
.execute()
.await
.expect("ordinary tables must keep overwrite support");
let values = ordinary_values(&fixture.table).await;
assert_eq!(values, HashSet::from(["replaced".to_string()]));
assert_eq!(fixture.table.count_rows(None).await.unwrap(), 1);
}
#[tokio::test]
async fn nonempty_sql_insert_invalidates_generated_column() {
let fixture = create_table_with_complete_literal_generated("b4b_sql_insert").await;
let before = read_generated_definition(&fixture.table).await;
let ctx = sql_ctx_for(&fixture.table, &fixture.table_name).await;
run_sql(
&ctx,
&format!(
"INSERT INTO {} VALUES (CAST(NULL AS INT), 'sql-appended')",
fixture.table_name
),
)
.await
.expect("non-empty SQL INSERT must commit");
fixture.table.checkout_latest().await.unwrap();
let values = ordinary_values(&fixture.table).await;
assert!(values.contains("sql-appended"));
assert_eq!(
fixture
.table
.generated_column_status(GEN_OUT)
.await
.unwrap(),
GeneratedColumnStatus::Incomplete
);
let after = read_generated_definition(&fixture.table).await;
assert_eq!(after.dependency_epoch(), before.dependency_epoch() + 1);
assert_eq!(after.materialized_epoch(), before.materialized_epoch());
assert_eq!(after.function_call(), before.function_call());
let Err(err) = fixture
.table
.query()
.select(Select::columns(&[GEN_OUT]))
.execute()
.await
else {
panic!("SQL INSERT invalidation must trip generated query guard");
};
assert_generated_column_incomplete(&err, "sql insert query");
}
#[tokio::test]
async fn empty_sql_insert_leaves_complete_generated_column() {
let fixture = create_table_with_complete_literal_generated("b4b_empty_sql_insert").await;
let before = read_generated_definition(&fixture.table).await;
let rows_before = ordinary_values(&fixture.table).await;
let conn = ConnectBuilder::new(&fixture.uri).execute().await.unwrap();
let source_schema = Arc::new(Schema::new(vec![
Field::new(GEN_OUT, DataType::Int32, true),
Field::new(ORDINARY, DataType::Utf8, true),
]));
let empty_reader: Box<dyn arrow_array::RecordBatchReader + Send> =
Box::new(RecordBatchIterator::new(
std::iter::empty::<Result<RecordBatch, arrow_schema::ArrowError>>(),
source_schema,
));
let source = conn
.create_table("empty_source", empty_reader)
.execute()
.await
.unwrap();
let ctx = sql_ctx_for(&fixture.table, &fixture.table_name).await;
let source_provider = BaseTableAdapter::try_new(source.base_table().clone())
.await
.unwrap();
ctx.register_table("empty_source", Arc::new(source_provider))
.unwrap();
run_sql(
&ctx,
&format!(
"INSERT INTO {} SELECT * FROM empty_source",
fixture.table_name
),
)
.await
.expect("empty SQL INSERT is a supported path");
fixture.table.checkout_latest().await.unwrap();
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
assert_eq!(
fixture
.table
.generated_column_status(GEN_OUT)
.await
.unwrap(),
GeneratedColumnStatus::Complete
);
assert_eq!(read_generated_definition(&fixture.table).await, before);
}
#[tokio::test]
async fn sql_insert_overwrite_rejects_before_mutation_when_generated_column_present() {
let fixture = create_table_with_complete_literal_generated("b4b_sql_overwrite").await;
let version_before = fixture.table.version().await.unwrap();
let rows_before = ordinary_values(&fixture.table).await;
let definition_before = read_generated_definition(&fixture.table).await;
let ctx = sql_ctx_for(&fixture.table, &fixture.table_name).await;
let err = run_sql(
&ctx,
&format!(
"INSERT OVERWRITE INTO {} VALUES (10, 'sql-overwrite')",
fixture.table_name
),
)
.await
.expect_err("SQL INSERT OVERWRITE must reject when any generated column is present");
assert_not_supported(&err, "sql insert overwrite");
fixture.table.checkout_latest().await.unwrap();
assert_eq!(fixture.table.version().await.unwrap(), version_before);
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
assert_eq!(
read_generated_definition(&fixture.table).await,
definition_before
);
}
#[tokio::test]
async fn malformed_generated_metadata_rejects_append_before_mutation_and_redacts_marker() {
let fixture = create_ordinary_table("b4b_malformed_preflight").await;
let snapshot = fixture
.table
.generated_column_binding_snapshot()
.await
.unwrap();
let field_id = snapshot.field(GEN_OUT).expect(GEN_OUT).field_id();
let raw = format!(
r#"{{"format_version":1,"output_field_id":{field_id},"function_call":{MALFORMED_MARKER},"dependency_epoch":1,"materialized_epoch":1}}"#
);
assert!(raw.contains(MALFORMED_MARKER));
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
fixture
.table
.as_native()
.expect("generated-column fixture planting requires a Native table"),
GEN_OUT,
raw.clone(),
)
.await
.unwrap();
let version_before = fixture.table.version().await.unwrap();
let rows_before = ordinary_values(&fixture.table).await;
let err = fixture
.table
.add(ordinary_rows_batch(&["must-not-land"]))
.execute()
.await
.expect_err("malformed generated metadata must fail closed before append visibility");
assert_invalid_input_redacted(&err, "malformed append preflight");
assert_eq!(fixture.table.version().await.unwrap(), version_before);
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
assert!(
!ordinary_values(&fixture.table)
.await
.contains("must-not-land")
);
}
#[tokio::test]
async fn concurrent_same_field_append_one_winner_one_conflict() {
let fixture = create_table_with_complete_literal_generated("b4b_concurrent_append").await;
let conn = ConnectBuilder::new(&fixture.uri)
.read_consistency_interval(Duration::from_secs(3600))
.execute()
.await
.unwrap();
let table_a = conn
.open_table(&fixture.table_name)
.execute()
.await
.unwrap();
let table_b = conn
.open_table(&fixture.table_name)
.execute()
.await
.unwrap();
let basis_version = table_a.version().await.unwrap();
assert_eq!(table_b.version().await.unwrap(), basis_version);
let (result_a, result_b) = tokio::join!(
table_a.add(ordinary_rows_batch(&["winner-a"])).execute(),
table_b.add(ordinary_rows_batch(&["winner-b"])).execute(),
);
let outcomes = [result_a, result_b];
let wins = outcomes.iter().filter(|result| result.is_ok()).count();
let losses = outcomes.iter().filter(|result| result.is_err()).count();
assert_eq!(wins, 1, "exactly one same-basis append may publish");
assert_eq!(losses, 1, "exactly one same-basis append must conflict");
for result in &outcomes {
if let Err(err) = result {
assert_conflict_error(err, "concurrent same-field append loser");
}
}
let fresh = conn
.open_table(&fixture.table_name)
.execute()
.await
.unwrap();
let values = ordinary_values(&fresh).await;
assert!(values.contains("seed"));
let has_a = values.contains("winner-a");
let has_b = values.contains("winner-b");
assert!(
has_a ^ has_b,
"only winner rows may be visible, got {values:?}"
);
assert_eq!(
fresh.generated_column_status(GEN_OUT).await.unwrap(),
GeneratedColumnStatus::Incomplete
);
let definition = read_generated_definition(&fresh).await;
assert_eq!(definition.dependency_epoch(), INITIAL_DEPENDENCY_EPOCH + 1);
assert_eq!(definition.materialized_epoch(), INITIAL_MATERIALIZED_EPOCH);
}
+8 -4
View File
@@ -423,6 +423,7 @@ mod tests {
use futures::TryStreamExt;
use tempfile::tempdir;
use crate::JobResult;
use crate::connect;
use crate::connection::ConnectBuilder;
use crate::index::Index;
@@ -538,7 +539,8 @@ mod tests {
assert_eq!(job.id(), None);
// The build runs as a task, so the index need not exist yet; it must
// once the job resolves.
job.wait().await.unwrap();
let result = job.wait().await.unwrap();
assert_eq!(result, JobResult::None);
assert_eq!(table.list_indices().await.unwrap().len(), 1);
// Cancelling a finished job is a no-op.
job.cancel().await.unwrap();
@@ -570,10 +572,12 @@ mod tests {
})
.collect::<Vec<_>>();
for waiter in waiters {
waiter.await.unwrap().unwrap();
let result = waiter.await.unwrap().unwrap();
assert_eq!(result, JobResult::None);
}
// A wait after the job settled still reports the same outcome.
job.wait().await.unwrap();
let late = job.wait().await.unwrap();
assert_eq!(late, JobResult::None);
assert_eq!(table.list_indices().await.unwrap().len(), 1);
}
@@ -716,7 +720,7 @@ mod tests {
match job.wait().await {
Err(crate::Error::JobCancelled { .. }) => {}
// The build may finish before the abort lands.
Ok(()) => {}
Ok(JobResult::None) => {}
other => panic!("unexpected job outcome: {other:?}"),
}
}
+23 -4
View File
@@ -19,7 +19,7 @@ use datafusion_physical_plan::{
};
use futures::TryStreamExt;
use lance::Dataset;
use lance::dataset::transaction::{Operation, Transaction};
use lance::dataset::transaction::{Operation, SchemaMetadataUpdates, Transaction};
use lance::dataset::{CommitBuilder, InsertBuilder, WriteParams, WriteProgressFn};
use lance::io::exec::utils::InstrumentedRecordBatchStreamAdapter;
use lance_table::format::Fragment;
@@ -74,7 +74,9 @@ fn merge_transactions(mut transactions: Vec<Transaction>) -> Option<Transaction>
///
/// This plan executes inserts by:
/// 1. Each partition writes data independently using InsertBuilder::execute_uncommitted_stream
/// 2. The last partition to complete commits all transactions atomically
/// 2. The last partition to complete merges transactions, optionally attaches one
/// precomputed generated-column metadata patch when the merged write has rows,
/// then commits once
/// 3. Returns the count of inserted rows per partition
#[derive(Debug)]
pub struct InsertExec {
@@ -83,6 +85,10 @@ pub struct InsertExec {
input: Arc<dyn ExecutionPlan>,
write_params: WriteParams,
tracker: Option<Arc<WriteProgressTracker>>,
/// Optional whole-transaction field-metadata patch for generated-column
/// invalidation. Attached once after partition merge, and only when the
/// merged operation contains at least one written row.
schema_metadata_updates: Option<SchemaMetadataUpdates>,
properties: Arc<PlanProperties>,
partial_transactions: Arc<Mutex<Vec<Transaction>>>,
metrics: ExecutionPlanMetricsSet,
@@ -95,7 +101,7 @@ impl InsertExec {
input: Arc<dyn ExecutionPlan>,
write_params: WriteParams,
) -> Self {
Self::new_with_tracker(ds_wrapper, dataset, input, write_params, None)
Self::new_with_tracker(ds_wrapper, dataset, input, write_params, None, None)
}
pub(crate) fn new_with_tracker(
@@ -104,6 +110,7 @@ impl InsertExec {
input: Arc<dyn ExecutionPlan>,
write_params: WriteParams,
tracker: Option<Arc<WriteProgressTracker>>,
schema_metadata_updates: Option<SchemaMetadataUpdates>,
) -> Self {
let schema = COUNT_SCHEMA.clone();
let num_partitions = input.output_partitioning().partition_count();
@@ -120,6 +127,7 @@ impl InsertExec {
input,
write_params,
tracker,
schema_metadata_updates,
properties: Arc::new(properties),
partial_transactions: Arc::new(Mutex::new(Vec::with_capacity(num_partitions))),
metrics: ExecutionPlanMetricsSet::new(),
@@ -176,6 +184,7 @@ impl ExecutionPlan for InsertExec {
children[0].clone(),
self.write_params.clone(),
self.tracker.clone(),
self.schema_metadata_updates.clone(),
)))
}
@@ -191,6 +200,7 @@ impl ExecutionPlan for InsertExec {
let total_partitions = self.input.output_partitioning().partition_count();
let ds_wrapper = self.ds_wrapper.clone();
let tracker = self.tracker.clone();
let schema_metadata_updates = self.schema_metadata_updates.clone();
let output_bytes = MetricBuilder::new(&self.metrics).output_bytes(partition);
let input_schema = input_stream.schema();
@@ -220,6 +230,8 @@ impl ExecutionPlan for InsertExec {
}));
}
// Each partition stages an uncommitted data-only transaction.
// Metadata invalidation is attached once on the merged commit.
let transaction = InsertBuilder::new(dataset.clone())
.with_params(&write_params)
.execute_uncommitted_stream(input_stream)
@@ -241,8 +253,15 @@ impl ExecutionPlan for InsertExec {
};
if let Some(transactions) = to_commit
&& let Some(merged_txn) = merge_transactions(transactions)
&& let Some(mut merged_txn) = merge_transactions(transactions)
{
// Attach the precomputed patch only for non-empty writes, and
// only once for the whole multi-partition transaction.
if count_rows_from_operation(&merged_txn.operation) > 0
&& let Some(updates) = schema_metadata_updates
{
merged_txn = merged_txn.with_schema_metadata_updates(updates)?;
}
let new_dataset = CommitBuilder::new(dataset.clone())
.execute(merged_txn)
.await?;
+37 -31
View File
@@ -1,9 +1,9 @@
use std::sync::Arc;
use futures::FutureExt;
use lance::dataset::DeleteBuilder;
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
use std::sync::Arc;
use lance::dataset::DeleteBuilder;
use serde::{Deserialize, Serialize};
use super::{NativeTable, Predicate};
@@ -29,34 +29,40 @@ pub(crate) async fn execute_delete(
predicate: Predicate<'_>,
) -> Result<DeleteResult> {
table.dataset.ensure_mutable()?;
match predicate {
Predicate::String(s) => {
let mut dataset = (*table.dataset.get().await?).clone();
let delete_result = dataset.delete(s).boxed().await?;
let num_deleted_rows = delete_result.num_deleted_rows;
let version = dataset.version().version;
table.dataset.update(dataset);
Ok(DeleteResult {
num_deleted_rows,
version,
})
}
Predicate::Expr(expr) => {
let dataset = table.dataset.get().await?;
let delete_result = DeleteBuilder::from_expr(Arc::clone(&dataset), expr.clone())
.execute()
.await?;
let num_deleted_rows = delete_result.num_deleted_rows;
let version = delete_result.new_dataset.version().version;
table.dataset.update(
Arc::try_unwrap(delete_result.new_dataset).unwrap_or_else(|arc| (*arc).clone()),
);
Ok(DeleteResult {
num_deleted_rows,
version,
})
}
// One exact dataset supplies binding-snapshot planning, the DeleteBuilder,
// and its transaction basis. Do not call table schema()/version() or another
// get(). Conflicts are not caught/replanned here.
let dataset = table.dataset.get().await?;
// String preserves the legacy Dataset::delete zero-retry baseline; Expr
// retains DeleteBuilder defaults until a generated patch is attached.
let mut builder = match predicate {
Predicate::String(s) => DeleteBuilder::new(Arc::clone(&dataset), s).conflict_retries(0),
Predicate::Expr(expr) => DeleteBuilder::from_expr(Arc::clone(&dataset), expr.clone()),
};
if let Some(schema_metadata_updates) =
super::generated_column_invalidation::plan_native_delete_generated_column_invalidation(
dataset.as_ref(),
)?
{
// Exact-basis fence: never retry an old generated patch on latest.
builder = builder
.with_schema_metadata_updates(schema_metadata_updates)?
.conflict_retries(0);
}
let delete_result = builder.execute().await?;
let num_deleted_rows = delete_result.num_deleted_rows;
let version = delete_result.new_dataset.version().version;
table
.dataset
.update(Arc::try_unwrap(delete_result.new_dataset).unwrap_or_else(|arc| (*arc).clone()));
Ok(DeleteResult {
num_deleted_rows,
version,
})
}
#[cfg(test)]
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,242 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! Crate-private Native wiring for generated-column invalidation (B4b / B4c / B4d / B4e).
//!
//! Converts the B4a pure planner into one Lance field-metadata patch for Native
//! append, update, and delete commits. Planning is strict-decode/validate;
//! overwrite of a table with any generated-column definition, direct writes of
//! generated outputs via Update, and Native merge-insert (standard and LSM)
//! fail closed as [`Error::NotSupported`].
use std::collections::{BTreeSet, HashMap};
use lance::Dataset;
use lance::dataset::transaction::{SchemaMetadataUpdates, UpdateMap, UpdateMapEntry};
use crate::Result;
use crate::error::Error;
use crate::function::GENERATED_COLUMN_METADATA_KEY;
use crate::function::plan_generated_column_invalidation::{
GeneratedColumnMutationImpact, PlannedGeneratedColumnMetadataUpdate,
plan_generated_column_invalidation,
};
use super::generated_column_binding_snapshot_from_dataset;
/// Plan Native append invalidation against one exact dataset snapshot.
///
/// Strict-decodes and validates every present generated-column metadata value
/// through the B4a planner. When `is_overwrite` is true and any generated column
/// is present, returns [`Error::NotSupported`] before mutation. Otherwise returns
/// `Some(patch)` when at least one generated column would be invalidated, or
/// `None` when the table has no generated columns.
pub(super) fn plan_native_append_generated_column_invalidation(
dataset: &Dataset,
is_overwrite: bool,
) -> Result<Option<SchemaMetadataUpdates>> {
let snapshot = generated_column_binding_snapshot_from_dataset(dataset)?;
let plan = plan_generated_column_invalidation(
&snapshot,
&GeneratedColumnMutationImpact::RowSetChanged,
)?;
if plan.is_empty() {
return Ok(None);
}
if is_overwrite {
return Err(Error::NotSupported {
message: "Overwrite is not supported on tables with generated columns".to_string(),
});
}
Ok(Some(planned_invalidation_to_schema_metadata_updates(plan)))
}
/// Plan Native update invalidation against one exact dataset snapshot.
///
/// Strict-decodes and validates every present generated-column definition before
/// impact calculation, even when `updated_field_ids` does not affect any
/// generated output. After the global planner succeeds, a target whose snapshot
/// entry contains generated metadata is rejected as a direct generated-output
/// write ([`Error::NotSupported`]) before any Update file write. Returns
/// `Some(patch)` when the impact closure is non-empty, otherwise `None`.
pub(super) fn plan_native_update_generated_column_invalidation(
dataset: &Dataset,
updated_field_ids: BTreeSet<i32>,
) -> Result<Option<SchemaMetadataUpdates>> {
let snapshot = generated_column_binding_snapshot_from_dataset(dataset)?;
let plan = plan_generated_column_invalidation(
&snapshot,
&GeneratedColumnMutationImpact::UpdatedFields(updated_field_ids.clone()),
)?;
for field_id in &updated_field_ids {
let Some(entry) = snapshot
.entries()
.iter()
.find(|entry| entry.field_id() == *field_id)
else {
return Err(Error::InvalidInput {
message: format!("updated field id {field_id} was not found in the table schema"),
});
};
if entry
.field()
.metadata()
.contains_key(GENERATED_COLUMN_METADATA_KEY)
{
return Err(Error::NotSupported {
message: "Updating generated columns is not supported".to_string(),
});
}
}
if plan.is_empty() {
return Ok(None);
}
Ok(Some(planned_invalidation_to_schema_metadata_updates(plan)))
}
/// Plan Native delete invalidation against one exact dataset snapshot.
///
/// Strict-decodes and validates every present generated-column metadata value
/// through the B4a `RowSetChanged` planner before any Delete scanner/file IO.
/// Returns `Some(patch)` when at least one generated column would be invalidated,
/// or `None` when the table has no generated columns. Actual zero-row Delete
/// suppression is owned by Lance A4d, not this planner.
pub(super) fn plan_native_delete_generated_column_invalidation(
dataset: &Dataset,
) -> Result<Option<SchemaMetadataUpdates>> {
let snapshot = generated_column_binding_snapshot_from_dataset(dataset)?;
let plan = plan_generated_column_invalidation(
&snapshot,
&GeneratedColumnMutationImpact::RowSetChanged,
)?;
if plan.is_empty() {
return Ok(None);
}
Ok(Some(planned_invalidation_to_schema_metadata_updates(plan)))
}
/// Fail closed before Native `merge_insert` when any generated column is present.
///
/// Strict-decodes and validates every present generated-column metadata value
/// through the B4a `RowSetChanged` planner against one exact dataset snapshot.
/// Malformed metadata returns the existing [`Error::InvalidInput`] validation
/// category. When at least one valid generated column is present, returns
/// [`Error::NotSupported`] before LSM dispatch or source iteration. Ordinary
/// tables (no generated metadata) return `Ok(())`.
pub(super) fn reject_native_merge_insert_if_generated_columns_present(
dataset: &Dataset,
) -> Result<()> {
let snapshot = generated_column_binding_snapshot_from_dataset(dataset)?;
let plan = plan_generated_column_invalidation(
&snapshot,
&GeneratedColumnMutationImpact::RowSetChanged,
)?;
if plan.is_empty() {
return Ok(());
}
Err(Error::NotSupported {
message: "Merge insert is not supported on tables with generated columns".to_string(),
})
}
/// Convert planner replacements into one non-empty Lance field-metadata patch.
///
/// Each entry is keyed by stable output field ID, uses `replace: false`, and
/// replaces only [`GENERATED_COLUMN_METADATA_KEY`].
fn planned_invalidation_to_schema_metadata_updates(
plan: Vec<PlannedGeneratedColumnMetadataUpdate>,
) -> SchemaMetadataUpdates {
SchemaMetadataUpdates {
schema_metadata_updates: None,
field_metadata_updates: plan
.into_iter()
.map(|update| {
(
update.output_field_id(),
UpdateMap {
update_entries: vec![UpdateMapEntry {
key: GENERATED_COLUMN_METADATA_KEY.to_string(),
value: Some(update.metadata_json().to_string()),
}],
replace: false,
},
)
})
.collect::<HashMap<_, _>>(),
}
}
#[cfg(test)]
mod tests {
use super::*;
/// Construct a planned update through the public accessors by planning a
/// minimal in-memory snapshot, then assert the Lance patch shape.
#[test]
fn planned_replacements_become_non_replace_field_patch() {
use crate::function::{
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput,
FunctionParameter, FunctionSignature, GeneratedColumnBindingSnapshot,
GeneratedColumnDefinition,
};
use arrow_array::{ArrayRef, StringArray};
use arrow_schema::{DataType, Field};
use std::sync::Arc;
let field_id = 11;
let function = Function::new(
FunctionId::try_new("fn.exact.b4b.helper.patch").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("x")])) as ArrayRef
)
.unwrap(),
)],
)
.unwrap();
let definition = GeneratedColumnDefinition::try_new(field_id, call, 3, 3).unwrap();
let json = definition.to_metadata_json().unwrap();
let snap = GeneratedColumnBindingSnapshot::try_new(
1,
vec![Arc::new(
Field::new("gen_out", DataType::Int32, true)
.with_metadata([(GENERATED_COLUMN_METADATA_KEY.to_string(), json)].into()),
)],
vec![field_id],
)
.unwrap();
let plan = plan_generated_column_invalidation(
&snap,
&GeneratedColumnMutationImpact::RowSetChanged,
)
.unwrap();
let patch = planned_invalidation_to_schema_metadata_updates(plan);
assert!(!patch.is_empty());
assert!(patch.schema_metadata_updates.is_none());
let map = patch
.field_metadata_updates
.get(&field_id)
.expect("stable field id must be present");
assert!(!map.replace);
assert_eq!(map.update_entries.len(), 1);
assert_eq!(map.update_entries[0].key, GENERATED_COLUMN_METADATA_KEY);
let decoded = GeneratedColumnDefinition::from_metadata_json(
map.update_entries[0].value.as_deref().unwrap(),
field_id,
)
.unwrap();
assert_eq!(decoded.dependency_epoch(), 4);
assert_eq!(decoded.materialized_epoch(), 3);
}
}
+15 -4
View File
@@ -233,10 +233,22 @@ pub(crate) async fn execute_merge_insert(
params: MergeInsertBuilder,
new_data: Box<dyn RecordBatchReader + Send>,
) -> Result<MergeResult> {
match lsm::lsm_dispatch_decision(table, &params).await? {
// One exact dataset supplies the generated-column fail-closed guard and
// downstream standard/LSM routing/execution. Do not refetch after the guard.
let dataset = table.dataset.get().await?;
super::generated_column_invalidation::reject_native_merge_insert_if_generated_columns_present(
dataset.as_ref(),
)?;
match lsm::lsm_dispatch_decision(&params, dataset.as_ref()).await? {
lsm::LsmDispatch::Lsm(plan) => {
let future =
lsm::execute_lsm_merge_insert(table, plan, params.validate_single_shard, new_data);
let future = lsm::execute_lsm_merge_insert(
table,
plan,
params.validate_single_shard,
new_data,
dataset,
);
return match params.timeout {
Some(timeout) => match tokio::time::timeout(timeout, future).await {
Ok(result) => result,
@@ -250,7 +262,6 @@ pub(crate) async fn execute_merge_insert(
lsm::LsmDispatch::Standard => {}
}
let dataset = table.dataset.get().await?;
let mut builder = LanceMergeInsertBuilder::try_new(dataset.clone(), params.on)?;
match (
params.when_matched_update_all,
+7 -4
View File
@@ -531,18 +531,18 @@ pub(crate) enum LsmDispatch {
}
/// Decide whether a `merge_insert` should be routed through the MemWAL write
/// path, validating the builder against the installed spec.
/// path, validating the builder against the installed spec on the exact
/// caller-supplied dataset snapshot.
#[allow(clippy::redundant_pub_crate)]
pub(crate) async fn lsm_dispatch_decision(
table: &NativeTable,
params: &MergeInsertBuilder,
dataset: &Dataset,
) -> Result<LsmDispatch> {
// Explicit opt-out: use the standard path regardless of any installed spec.
if params.use_lsm == Some(false) {
return Ok(LsmDispatch::Standard);
}
let dataset = table.dataset.get().await?;
let Some(details) = dataset.mem_wal_index_details().await? else {
// No write spec installed. `use_lsm(true)` demanded MemWAL routing, so
// that is an error; otherwise fall back to the standard path.
@@ -646,14 +646,17 @@ fn resolve_lsm_mode(details: &MemWalIndexDetails) -> Result<LsmMode> {
/// a validation failure (e.g. input spanning shards) never leaves a partial
/// write behind. When `validate_single_shard` is set, every row is checked to
/// route to one shard; when disabled, only the first row of the whole input is.
///
/// `dataset` must be the same exact snapshot used for the generated-column
/// guard and [`lsm_dispatch_decision`].
#[allow(clippy::redundant_pub_crate)]
pub(crate) async fn execute_lsm_merge_insert(
table: &NativeTable,
plan: LsmPlan,
validate_single_shard: bool,
new_data: Box<dyn RecordBatchReader + Send>,
dataset: Arc<Dataset>,
) -> Result<MergeResult> {
let dataset = table.dataset.get().await?;
let target_schema: SchemaRef = Arc::new(ArrowSchema::from(dataset.schema()));
// Collect, align and shard-validate the whole input before writing
@@ -0,0 +1,531 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! RED runtime contract tests for Native merge-insert fail-closed guard (B4e).
//!
//! Tables with generated-column definitions cannot carry dependency-epoch
//! metadata updates through Native merge-insert in this slice. Both the
//! standard and MemWAL/LSM routes must reject before consuming source input or
//! mutating the table. Ordinary tables keep existing merge-insert semantics.
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use arrow_array::{Int32Array, RecordBatch, RecordBatchReader, StringArray};
use arrow_schema::{ArrowError, DataType, Field, Schema, SchemaRef};
use futures::TryStreamExt;
use tempfile::TempDir;
use crate::connection::ConnectBuilder;
use crate::error::Error;
use crate::function::{
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition,
GeneratedColumnStatus,
};
use crate::query::{ExecutableQuery, QueryBase, Select};
use crate::table::Table;
const ID: &str = "id";
const ORDINARY: &str = "ordinary";
const GEN_OUT: &str = "gen_out";
const INITIAL_DEPENDENCY_EPOCH: u64 = 3;
const INITIAL_MATERIALIZED_EPOCH: u64 = 3;
const FN_ID: &str = "fn.exact.b4e.merge.literal";
const MALFORMED_MARKER: &str = "SENSITIVE_B4E_MERGE_METADATA_MARKER_7c91_e2ab";
struct Fixture {
_tmp: TempDir,
table: Table,
table_name: String,
uri: String,
}
/// RecordBatchReader that counts how many times [`Self::next`] is called.
struct ObservableReader {
inner: Box<dyn RecordBatchReader + Send>,
next_calls: Arc<AtomicUsize>,
}
impl ObservableReader {
fn wrap(
inner: Box<dyn RecordBatchReader + Send>,
next_calls: Arc<AtomicUsize>,
) -> Box<dyn RecordBatchReader + Send> {
Box::new(Self { inner, next_calls })
}
}
impl Iterator for ObservableReader {
type Item = Result<RecordBatch, ArrowError>;
fn next(&mut self) -> Option<Self::Item> {
self.next_calls.fetch_add(1, Ordering::SeqCst);
self.inner.next()
}
}
impl RecordBatchReader for ObservableReader {
fn schema(&self) -> SchemaRef {
self.inner.schema()
}
}
fn literal_only_function() -> Function {
Function::new(
FunctionId::try_new(FN_ID).unwrap(),
FunctionSignature::try_new(
vec![FunctionParameter::new("label", DataType::Utf8)],
FunctionOutput::new(DataType::Int32, true),
)
.unwrap(),
)
}
fn literal_only_definition(output_field_id: i32) -> GeneratedColumnDefinition {
let function = literal_only_function();
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,
INITIAL_DEPENDENCY_EPOCH,
INITIAL_MATERIALIZED_EPOCH,
)
.unwrap()
}
fn seed_batch() -> RecordBatch {
let schema = Arc::new(Schema::new(vec![
Field::new(ID, DataType::Int32, false),
Field::new(ORDINARY, DataType::Utf8, true),
Field::new(GEN_OUT, DataType::Int32, true),
]));
RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from(vec![1, 2])),
Arc::new(StringArray::from(vec![Some("a"), Some("b")])),
Arc::new(Int32Array::from(vec![10, 20])),
],
)
.unwrap()
}
fn source_batch(ids: &[i32], ordinary: &[&str], gen_values: &[i32]) -> RecordBatch {
assert_eq!(ids.len(), ordinary.len());
assert_eq!(ids.len(), gen_values.len());
let schema = Arc::new(Schema::new(vec![
Field::new(ID, DataType::Int32, false),
Field::new(ORDINARY, DataType::Utf8, true),
Field::new(GEN_OUT, DataType::Int32, true),
]));
RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from(ids.to_vec())),
Arc::new(StringArray::from(
ordinary
.iter()
.map(|value| Some(*value))
.collect::<Vec<_>>(),
)),
Arc::new(Int32Array::from(gen_values.to_vec())),
],
)
.unwrap()
}
fn boxed_reader(batch: RecordBatch) -> Box<dyn RecordBatchReader + Send> {
let schema = batch.schema();
Box::new(arrow_array::RecordBatchIterator::new(
vec![Ok(batch)].into_iter(),
schema,
))
}
fn empty_reader() -> Box<dyn RecordBatchReader + Send> {
let schema = Arc::new(Schema::new(vec![
Field::new(ID, DataType::Int32, false),
Field::new(ORDINARY, DataType::Utf8, true),
Field::new(GEN_OUT, DataType::Int32, true),
]));
Box::new(arrow_array::RecordBatchIterator::new(
std::iter::empty::<Result<RecordBatch, ArrowError>>(),
schema,
))
}
async fn create_ordinary_table(name: &str) -> Fixture {
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap().to_string();
let conn = ConnectBuilder::new(&uri).execute().await.unwrap();
let table = conn
.create_table(name, seed_batch())
.execute()
.await
.unwrap();
Fixture {
_tmp: tmp,
table,
table_name: name.to_string(),
uri,
}
}
async fn create_table_with_complete_literal_generated(name: &str) -> Fixture {
let fixture = create_ordinary_table(name).await;
let snapshot = fixture
.table
.generated_column_binding_snapshot()
.await
.unwrap();
let field_id = snapshot.field(GEN_OUT).expect(GEN_OUT).field_id();
let definition = literal_only_definition(field_id);
let json = definition.to_metadata_json().unwrap();
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
fixture
.table
.as_native()
.expect("generated-column fixture planting requires a Native table"),
GEN_OUT,
json,
)
.await
.unwrap();
assert_eq!(
fixture
.table
.generated_column_status(GEN_OUT)
.await
.unwrap(),
GeneratedColumnStatus::Complete
);
fixture
}
async fn read_generated_definition(table: &Table) -> GeneratedColumnDefinition {
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
snapshot
.field(GEN_OUT)
.expect(GEN_OUT)
.generated_column_definition()
.expect("generated metadata must decode")
.expect("generated metadata must be present")
}
async fn read_raw_generated_metadata(table: &Table) -> String {
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
snapshot
.field(GEN_OUT)
.expect(GEN_OUT)
.field()
.metadata()
.get(GENERATED_COLUMN_METADATA_KEY)
.expect("generated metadata key must be present")
.clone()
}
async fn ordinary_rows(table: &Table) -> Vec<(i32, String)> {
let batches = table
.query()
.select(Select::columns(&[ID, ORDINARY]))
.execute()
.await
.unwrap()
.try_collect::<Vec<_>>()
.await
.unwrap();
let mut rows = Vec::new();
for batch in batches {
let ids = batch
.column(0)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap();
let ordinary = batch
.column(1)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
for index in 0..batch.num_rows() {
rows.push((ids.value(index), ordinary.value(index).to_string()));
}
}
rows.sort_by_key(|(id, _)| *id);
rows
}
fn assert_not_supported(err: &Error, label: &str) {
assert!(
matches!(err, Error::NotSupported { .. }),
"{label}: expected NotSupported, got {err:?}"
);
}
fn assert_invalid_input_redacted(err: &Error, planted_raw: &str, label: &str) {
assert!(
matches!(err, Error::InvalidInput { .. }),
"{label}: expected InvalidInput, got {err:?}"
);
let rendered = format!("{err}\n{err:?}");
assert!(
!rendered.contains(MALFORMED_MARKER),
"{label}: diagnostic echoed unique metadata marker: {rendered}"
);
assert!(
!rendered.contains(FN_ID),
"{label}: diagnostic echoed Function ID: {rendered}"
);
assert!(
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
"{label}: diagnostic echoed metadata wire key: {rendered}"
);
assert!(
!rendered.contains(planted_raw),
"{label}: diagnostic echoed raw metadata JSON: {rendered}"
);
}
fn configure_standard_merge(builder: &mut crate::table::merge::MergeInsertBuilder) {
builder
.when_matched_update_all(None)
.when_not_matched_insert_all()
.when_not_matched_by_source_delete(None);
}
#[tokio::test]
async fn standard_merge_insert_rejects_when_generated_column_present_before_input_consumption() {
let fixture = create_table_with_complete_literal_generated("b4e_standard_reject").await;
let version_before = fixture.table.version().await.unwrap();
let rows_before = ordinary_rows(&fixture.table).await;
let definition_before = read_generated_definition(&fixture.table).await;
let raw_before = read_raw_generated_metadata(&fixture.table).await;
assert_eq!(
definition_before.function_call().function_id().as_str(),
FN_ID
);
let next_calls = Arc::new(AtomicUsize::new(0));
let reader = ObservableReader::wrap(
boxed_reader(source_batch(&[1, 3], &["updated", "inserted"], &[11, 30])),
next_calls.clone(),
);
let mut builder = fixture.table.merge_insert(&[ID]);
configure_standard_merge(&mut builder);
let err = builder
.execute(reader)
.await
.expect_err("generated-column table must reject standard merge_insert");
assert_not_supported(&err, "standard merge_insert generated reject");
assert_eq!(
next_calls.load(Ordering::SeqCst),
0,
"rejection must occur before consuming the RecordBatchReader"
);
assert_eq!(fixture.table.version().await.unwrap(), version_before);
assert_eq!(ordinary_rows(&fixture.table).await, rows_before);
assert_eq!(
read_generated_definition(&fixture.table).await,
definition_before
);
assert_eq!(
read_raw_generated_metadata(&fixture.table).await,
raw_before
);
assert_eq!(
fixture
.table
.generated_column_status(GEN_OUT)
.await
.unwrap(),
GeneratedColumnStatus::Complete
);
}
#[tokio::test]
async fn empty_standard_merge_insert_rejects_when_generated_column_present() {
let fixture = create_table_with_complete_literal_generated("b4e_empty_reject").await;
let version_before = fixture.table.version().await.unwrap();
let rows_before = ordinary_rows(&fixture.table).await;
let raw_before = read_raw_generated_metadata(&fixture.table).await;
let next_calls = Arc::new(AtomicUsize::new(0));
let reader = ObservableReader::wrap(empty_reader(), next_calls.clone());
let mut builder = fixture.table.merge_insert(&[ID]);
configure_standard_merge(&mut builder);
let err = builder
.execute(reader)
.await
.expect_err("empty merge_insert must still reject on generated-column tables");
assert_not_supported(&err, "empty standard merge_insert generated reject");
assert_eq!(next_calls.load(Ordering::SeqCst), 0);
assert_eq!(fixture.table.version().await.unwrap(), version_before);
assert_eq!(ordinary_rows(&fixture.table).await, rows_before);
assert_eq!(
read_raw_generated_metadata(&fixture.table).await,
raw_before
);
}
#[tokio::test]
async fn forced_lsm_without_spec_rejects_generated_before_missing_spec_and_input() {
let fixture = create_table_with_complete_literal_generated("b4e_lsm_force_reject").await;
let version_before = fixture.table.version().await.unwrap();
let rows_before = ordinary_rows(&fixture.table).await;
let raw_before = read_raw_generated_metadata(&fixture.table).await;
let next_calls = Arc::new(AtomicUsize::new(0));
let reader = ObservableReader::wrap(
boxed_reader(source_batch(&[1], &["must-not-land"], &[11])),
next_calls.clone(),
);
let mut builder = fixture.table.merge_insert(&[ID]);
builder
.when_matched_update_all(None)
.when_not_matched_insert_all()
.use_lsm(true);
let err = builder
.execute(reader)
.await
.expect_err("generated-column guard must run before LSM missing-spec validation");
assert_not_supported(&err, "forced LSM generated reject");
assert_eq!(next_calls.load(Ordering::SeqCst), 0);
assert_eq!(fixture.table.version().await.unwrap(), version_before);
assert_eq!(ordinary_rows(&fixture.table).await, rows_before);
assert_eq!(
read_raw_generated_metadata(&fixture.table).await,
raw_before
);
}
#[tokio::test]
async fn malformed_generated_metadata_rejects_merge_insert_before_mutation_and_redacts() {
let fixture = create_ordinary_table("b4e_malformed_preflight").await;
let snapshot = fixture
.table
.generated_column_binding_snapshot()
.await
.unwrap();
let field_id = snapshot.field(GEN_OUT).expect(GEN_OUT).field_id();
let planted_raw = format!(
r#"{{"format_version":1,"output_field_id":{field_id},"function_call":{{"function_id":"{FN_ID}","marker":"{MALFORMED_MARKER}"}},"dependency_epoch":1,"materialized_epoch":1}}"#
);
assert!(planted_raw.contains(MALFORMED_MARKER));
assert!(planted_raw.contains(FN_ID));
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
fixture
.table
.as_native()
.expect("generated-column fixture planting requires a Native table"),
GEN_OUT,
planted_raw.clone(),
)
.await
.unwrap();
assert_eq!(
read_raw_generated_metadata(&fixture.table).await,
planted_raw,
"planted malformed raw metadata must round-trip byte-for-byte"
);
let version_before = fixture.table.version().await.unwrap();
let rows_before = ordinary_rows(&fixture.table).await;
let next_calls = Arc::new(AtomicUsize::new(0));
let reader = ObservableReader::wrap(
boxed_reader(source_batch(&[1], &["must-not-land"], &[99])),
next_calls.clone(),
);
let mut builder = fixture.table.merge_insert(&[ID]);
configure_standard_merge(&mut builder);
let err = builder
.execute(reader)
.await
.expect_err("malformed generated metadata must fail closed before merge_insert");
assert_invalid_input_redacted(&err, &planted_raw, "malformed merge_insert preflight");
assert_eq!(next_calls.load(Ordering::SeqCst), 0);
let fresh = ConnectBuilder::new(&fixture.uri)
.execute()
.await
.unwrap()
.open_table(&fixture.table_name)
.execute()
.await
.unwrap();
assert_eq!(fresh.version().await.unwrap(), version_before);
assert_eq!(ordinary_rows(&fresh).await, rows_before);
assert_eq!(read_raw_generated_metadata(&fresh).await, planted_raw);
}
#[tokio::test]
async fn ordinary_table_standard_merge_insert_preserves_result_semantics() {
let fixture = create_ordinary_table("b4e_ordinary_standard").await;
let mut builder = fixture.table.merge_insert(&[ID]);
configure_standard_merge(&mut builder);
let result = builder
.execute(boxed_reader(source_batch(
&[1, 3],
&["updated", "inserted"],
&[11, 30],
)))
.await
.expect("ordinary-table standard merge_insert must succeed");
assert_eq!(result.num_inserted_rows, 1);
assert_eq!(result.num_updated_rows, 1);
assert_eq!(result.num_deleted_rows, 1);
assert_eq!(result.num_attempts, 1);
assert_eq!(result.num_rows, 2);
assert!(result.version > 0);
assert_eq!(
ordinary_rows(&fixture.table).await,
vec![(1, "updated".to_string()), (3, "inserted".to_string()),]
);
}
#[tokio::test]
async fn ordinary_table_forced_lsm_without_spec_keeps_missing_spec_error() {
let fixture = create_ordinary_table("b4e_ordinary_lsm_missing_spec").await;
let version_before = fixture.table.version().await.unwrap();
let rows_before = ordinary_rows(&fixture.table).await;
let mut builder = fixture.table.merge_insert(&[ID]);
builder
.when_matched_update_all(None)
.when_not_matched_insert_all()
.use_lsm(true);
let err = builder
.execute(boxed_reader(source_batch(&[1], &["x"], &[1])))
.await
.expect_err("ordinary table without MemWAL spec must keep missing-spec InvalidInput");
match err {
Error::InvalidInput { message } => {
assert!(
message.contains("no MemWAL write spec"),
"expected missing-spec message, got {message}"
);
}
other => panic!("expected InvalidInput missing-spec, got {other:?}"),
}
assert_eq!(fixture.table.version().await.unwrap(), version_before);
assert_eq!(ordinary_rows(&fixture.table).await, rows_before);
}
File diff suppressed because it is too large Load Diff
+377 -19
View File
@@ -5,12 +5,13 @@ use std::sync::Arc;
mod lsm;
use super::NativeTable;
use super::{NativeTable, generated_column_binding_snapshot_from_dataset};
use crate::connection::NamespaceClientPushdownOperation;
use crate::error::{Error, Result};
use crate::expr::expr_to_sql_string;
use crate::query::{
DEFAULT_TOP_K, QueryExecutionOptions, QueryFilter, QueryRequest, Select, VectorQueryRequest,
validate_generated_column_query,
};
use crate::utils::{MaxBatchLengthStream, TimeoutStream, default_vector_column};
use arrow::array::{AsArray, FixedSizeListBuilder, Float32Builder};
@@ -22,6 +23,7 @@ use datafusion_physical_plan::projection::ProjectionExec;
use datafusion_physical_plan::repartition::RepartitionExec;
use datafusion_physical_plan::union::UnionExec;
use futures::future::try_join_all;
use lance::Dataset;
use lance::dataset::mem_wal::DatasetMemWalExt;
use lance::dataset::scanner::DatasetRecordBatchStream;
use lance::dataset::scanner::Scanner;
@@ -47,7 +49,7 @@ impl AnyQuery {
}
}
//Decide between namespace or local
// Decide between namespace or local.
pub async fn execute_query(
table: &NativeTable,
query: &AnyQuery,
@@ -55,16 +57,25 @@ pub async fn execute_query(
) -> Result<DatasetRecordBatchStream> {
// QueryTable pushdown runs the query server-side, but only on the main
// branch: the namespace request carries no branch yet, so a branch handle
// must fall through to local execution.
if can_execute_namespace_query(table, query).await?
&& let Some(ref namespace_client) = table.namespace_client
{
return execute_namespace_query(table, namespace_client.clone(), query, options).await;
// must fall through to local execution. Successful pushdown owns one
// Dataset Arc for MemWAL eligibility, generated-column guard, and version
// fencing. Obviously ineligible paths avoid an unused get(); MemWAL
// fallthrough and other local paths acquire/guard/plan independently.
if let Some(stream) = try_execute_namespace_query(table, query).await? {
return Ok(stream);
}
execute_generic_query(table, query, options).await
}
async fn can_execute_namespace_query(table: &NativeTable, query: &AnyQuery) -> Result<bool> {
/// Attempt QueryTable pushdown against one exact Dataset snapshot.
///
/// Returns `Ok(None)` when pushdown is ineligible so the caller can fall through
/// to the local exact-snapshot planner. Does not `dataset.get()` on paths that
/// are obviously ineligible before the MemWAL check.
async fn try_execute_namespace_query(
table: &NativeTable,
query: &AnyQuery,
) -> Result<Option<DatasetRecordBatchStream>> {
if !(table
.pushdown_operations
.contains(&NamespaceClientPushdownOperation::QueryTable)
@@ -72,17 +83,38 @@ async fn can_execute_namespace_query(table: &NativeTable, query: &AnyQuery) -> R
&& table.dataset.current_branch().is_none()
&& !requires_local_namespace_execution(query))
{
return Ok(false);
return Ok(None);
}
let Some(namespace_client) = table.namespace_client.clone() else {
return Ok(None);
};
// One Dataset Arc owns MemWAL eligibility, the generated-column guard, and
// the version fence sent on the request.
let dataset = table.dataset.get().await?;
// A MemWAL write spec means reads auto-route through the LSM scanner in
// `create_plan` even when `use_lsm` is unset. The namespace request has no
// use_lsm field, so pushing the default query down would silently omit
// un-compacted rows — force local execution whenever a spec is installed.
let dataset = table.dataset.get().await?;
// Do not guard here; the local planner acquires its own snapshot.
if dataset.mem_wal_index_details().await?.is_some() {
return Ok(false);
return Ok(None);
}
Ok(true)
let snapshot = generated_column_binding_snapshot_from_dataset(dataset.as_ref())?;
validate_generated_column_query(&snapshot, query)?;
let version = i64::try_from(dataset.version().version).map_err(|_| Error::InvalidInput {
message: format!(
"dataset version {} exceeds i64::MAX and cannot be sent on QueryTable",
dataset.version().version
),
})?;
Ok(Some(
execute_namespace_query(table, namespace_client, query, version).await?,
))
}
fn requires_local_namespace_execution(query: &AnyQuery) -> bool {
@@ -128,18 +160,38 @@ async fn execute_generic_query(
Ok(DatasetRecordBatchStream::new(inner))
}
/// Public/internal Native planner entry: one Dataset get, guard, then plan.
///
/// Acquires exactly one [`Arc<Dataset>`], builds/validates the generated-column
/// snapshot from that object, then plans entirely against the same Arc.
/// Multi-vector recursion clones the owned Arc and must not call this outer
/// entry (no additional `dataset.get()`).
pub async fn create_plan(
table: &NativeTable,
query: &AnyQuery,
options: QueryExecutionOptions,
) -> Result<Arc<dyn ExecutionPlan>> {
let ds_ref = table.dataset.get().await?;
let snapshot = generated_column_binding_snapshot_from_dataset(ds_ref.as_ref())?;
// Pass the original full AnyQuery before VectorQuery conversion/splitting so
// select/filter/order/FTS/vector references are all covered. check_filter
// precedence is preserved inside validate_generated_column_query.
validate_generated_column_query(&snapshot, query)?;
create_plan_with_dataset(table, query, options, ds_ref).await
}
/// Plan against an already-owned Dataset Arc. Used by the guarded outer entry
/// and by multi-vector recursion (Arc clones only).
async fn create_plan_with_dataset(
table: &NativeTable,
query: &AnyQuery,
options: QueryExecutionOptions,
ds_ref: Arc<Dataset>,
) -> Result<Arc<dyn ExecutionPlan>> {
let query = match query {
AnyQuery::VectorQuery(query) => query.clone(),
AnyQuery::Query(query) => VectorQueryRequest::from_plain_query(query.clone()),
};
query.base.check_filter()?;
let ds_ref = table.dataset.get().await?;
// MemWAL read routing driven by `use_lsm`:
// * unset — route through the LSM scanner iff the table carries a write spec
@@ -198,7 +250,9 @@ pub async fn create_plan(
}
query_vector = Some(Arc::new(fsl_builder.finish()));
} else {
// Multiple query vectors: create a plan for each and union them
// Multiple query vectors: create a plan for each and union them.
// Recurse with clones of the already-owned Arc — never the outer
// create_plan entry (which would perform another dataset.get()).
let query_vecs = query.query_vector.clone();
let plan_futures = query_vecs
.into_iter()
@@ -206,8 +260,15 @@ pub async fn create_plan(
let mut sub_query = query.clone();
sub_query.query_vector = vec![query_vector];
let options_ref = options.clone();
let ds_ref = ds_ref.clone();
async move {
create_plan(table, &AnyQuery::VectorQuery(sub_query), options_ref).await
create_plan_with_dataset(
table,
&AnyQuery::VectorQuery(sub_query),
options_ref,
ds_ref,
)
.await
}
})
.collect::<Vec<_>>();
@@ -381,11 +442,15 @@ pub(crate) fn create_multi_vector_plan(
}
/// Execute a query on the namespace server instead of locally.
///
/// Caller must already have validated the generated-column guard against the
/// exact Dataset snapshot whose `version` is passed here. The incomplete /
/// malformed guard runs before this dispatch.
async fn execute_namespace_query(
table: &NativeTable,
namespace_client: Arc<dyn LanceNamespace>,
query: &AnyQuery,
_options: QueryExecutionOptions,
version: i64,
) -> Result<DatasetRecordBatchStream> {
// Build table_id from namespace + table name
let mut table_id = table.namespace.clone();
@@ -393,8 +458,9 @@ async fn execute_namespace_query(
// Convert AnyQuery to namespace QueryTableRequest
let mut ns_request = convert_to_namespace_query(query)?;
// Set the table ID on the request
// Set the table ID and exact guarded Dataset version on the request.
ns_request.id = Some(table_id);
ns_request.version = Some(version);
// Call the namespace query_table API
let response_bytes = namespace_client
@@ -1163,4 +1229,296 @@ mod tests {
Some(ApproxMode::Accurate)
);
}
/// Records namespace `query_table` requests and returns a valid Arrow IPC file.
#[derive(Debug, Default)]
struct RecordingNamespaceClient {
requests: std::sync::Mutex<Vec<NsQueryTableRequest>>,
}
impl RecordingNamespaceClient {
fn ipc_response() -> bytes::Bytes {
use arrow_array::{Int32Array, RecordBatch, StringArray};
use arrow_ipc::writer::FileWriter;
use arrow_schema::{DataType, Field, Schema};
let schema = Arc::new(Schema::new(vec![
Field::new("gen_out", DataType::Int32, true),
Field::new("ordinary", DataType::Utf8, true),
]));
let batch = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Int32Array::from(vec![1])) as ArrayRef,
Arc::new(StringArray::from(vec![Some("x")])) as ArrayRef,
],
)
.unwrap();
let mut buf = Vec::new();
{
let mut writer = FileWriter::try_new(&mut buf, &schema).unwrap();
writer.write(&batch).unwrap();
writer.finish().unwrap();
}
bytes::Bytes::from(buf)
}
fn call_count(&self) -> usize {
self.requests.lock().unwrap().len()
}
fn requests(&self) -> Vec<NsQueryTableRequest> {
self.requests.lock().unwrap().clone()
}
}
#[async_trait::async_trait]
impl LanceNamespace for RecordingNamespaceClient {
fn namespace_id(&self) -> String {
"recording".to_string()
}
async fn query_table(&self, request: NsQueryTableRequest) -> lance::Result<bytes::Bytes> {
self.requests.lock().unwrap().push(request);
Ok(Self::ipc_response())
}
}
async fn runtime_namespace_table(
name: &str,
) -> (
crate::table::Table,
NativeTable,
Arc<RecordingNamespaceClient>,
) {
use crate::connect;
use crate::connection::NamespaceClientPushdownOperation;
use arrow_array::{Int32Array, RecordBatch, StringArray};
use arrow_schema::{DataType, Field, Schema};
let conn = connect("memory://").execute().await.unwrap();
let schema = Arc::new(Schema::new(vec![
Field::new("gen_out", DataType::Int32, true),
Field::new("ordinary", DataType::Utf8, true),
]));
let batch = RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from(vec![1])),
Arc::new(StringArray::from(vec![Some("x")])),
],
)
.unwrap();
let table = conn.create_table(name, batch).execute().await.unwrap();
let namespace_client = Arc::new(RecordingNamespaceClient::default());
let mut native_table = table.as_native().unwrap().clone();
native_table.namespace_client = Some(namespace_client.clone());
native_table
.pushdown_operations
.insert(NamespaceClientPushdownOperation::QueryTable);
(table, native_table, namespace_client)
}
async fn plant_runtime_generated_column_metadata(
table: &crate::table::Table,
column: &str,
dependency_epoch: u64,
materialized_epoch: u64,
) {
use crate::function::{
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput,
FunctionParameter, FunctionSignature, GeneratedColumnDefinition,
};
use arrow_array::StringArray;
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
let field_id = snapshot.field(column).expect(column).field_id();
let function = Function::new(
FunctionId::try_new("fn.exact.status.native").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("ok")])) as ArrayRef
)
.unwrap(),
)],
)
.unwrap();
let json = GeneratedColumnDefinition::try_new(
field_id,
call,
dependency_epoch,
materialized_epoch,
)
.unwrap()
.to_metadata_json()
.unwrap();
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
table
.as_native()
.expect("generated-column fixture planting requires a Native table"),
column,
json,
)
.await
.unwrap();
}
fn assert_incomplete_runtime_error(err: &Error, label: &str) {
use crate::error::FunctionErrorCode;
use crate::function::GENERATED_COLUMN_METADATA_KEY;
match err {
Error::Function {
code: FunctionErrorCode::GeneratedColumnIncomplete,
message,
} => {
let rendered = format!("{err}\n{err:?}\n{message}");
assert!(
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
"{label}: leaked metadata key: {rendered}"
);
assert!(
!rendered.contains("function_call"),
"{label}: leaked function_call: {rendered}"
);
assert!(
!rendered.contains("fn.exact.status.native"),
"{label}: leaked Function ID: {rendered}"
);
}
other => panic!(
"{label}: expected Error::Function(GeneratedColumnIncomplete), got {other:?}"
),
}
}
#[tokio::test]
async fn generated_column_query_runtime_namespace_pushdown_fenced() {
use crate::function::GeneratedColumnStatus;
use crate::query::Select;
let (table, native_table, namespace_client) =
runtime_namespace_table("runtime_ns_fence").await;
plant_runtime_generated_column_metadata(&table, "gen_out", 3, 3).await;
assert_eq!(
table.generated_column_status("gen_out").await.unwrap(),
GeneratedColumnStatus::Complete
);
let guarded_version = native_table.dataset.get().await.unwrap().version().version;
let complete_query = AnyQuery::Query(QueryRequest {
select: Select::Columns(vec!["gen_out".to_string()]),
limit: Some(10),
..Default::default()
});
let stream = execute_query(
&native_table,
&complete_query,
QueryExecutionOptions::default(),
)
.await
.expect("complete generated-column query must be eligible for QueryTable pushdown");
let batches = stream.try_collect::<Vec<_>>().await.unwrap();
assert_eq!(batches.iter().map(|b| b.num_rows()).sum::<usize>(), 1);
assert_eq!(namespace_client.call_count(), 1);
let requests = namespace_client.requests();
assert_eq!(
requests[0].version,
Some(guarded_version as i64),
"pushed-down query must fence to the exact guarded Dataset version; None races to latest"
);
// Incomplete referenced output must reject before any namespace dispatch.
plant_runtime_generated_column_metadata(&table, "gen_out", 8, 2).await;
assert_eq!(
table.generated_column_status("gen_out").await.unwrap(),
GeneratedColumnStatus::Incomplete
);
let before_incomplete = namespace_client.call_count();
let incomplete_query = AnyQuery::Query(QueryRequest {
select: Select::Columns(vec!["gen_out".to_string()]),
limit: Some(10),
..Default::default()
});
let Err(incomplete_err) = execute_query(
&native_table,
&incomplete_query,
QueryExecutionOptions::default(),
)
.await
else {
panic!("incomplete generated-column query must reject before dispatch");
};
assert_incomplete_runtime_error(&incomplete_err, "namespace_incomplete");
assert_eq!(
namespace_client.call_count(),
before_incomplete,
"incomplete guard must not dispatch query_table"
);
}
#[tokio::test]
async fn generated_column_query_runtime_malformed_rejects_before_namespace() {
use crate::function::GENERATED_COLUMN_METADATA_KEY;
use crate::query::Select;
let (table, native_table, namespace_client) =
runtime_namespace_table("runtime_ns_malformed").await;
const MARKER: &str = "SENSITIVE_RUNTIME_MALFORMED_b3e2a_9c1d";
let raw = format!(
r#"{{"format_version":1,"output_field_id":0,"function_call":{MARKER},"dependency_epoch":1,"materialized_epoch":1}}"#
);
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
table
.as_native()
.expect("generated-column fixture planting requires a Native table"),
"gen_out",
raw.clone(),
)
.await
.unwrap();
let query = AnyQuery::Query(QueryRequest {
select: Select::Columns(vec!["gen_out".to_string()]),
limit: Some(10),
..Default::default()
});
let Err(err) = execute_query(&native_table, &query, QueryExecutionOptions::default()).await
else {
panic!("malformed referenced metadata must fail closed before dispatch");
};
assert!(
matches!(err, Error::InvalidInput { .. }),
"expected InvalidInput, got {err:?}"
);
let rendered = format!("{err}\n{err:?}");
assert!(
!rendered.contains(MARKER),
"malformed diagnostics must not echo raw marker: {rendered}"
);
assert!(
!rendered.contains(&raw),
"malformed diagnostics must not echo raw metadata: {rendered}"
);
assert!(
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
"malformed diagnostics must not echo metadata wire key: {rendered}"
);
assert_eq!(
namespace_client.call_count(),
0,
"malformed guard must not dispatch query_table"
);
}
}
+113 -1
View File
@@ -8,12 +8,94 @@
//! - [`alter_columns`](execute_alter_columns): Rename columns, change types, or modify nullability
//! - [`drop_columns`](execute_drop_columns): Remove columns from the table
use arrow_array::RecordBatchReader;
use lance::dataset::{ColumnAlteration, NewColumnTransform};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use super::NativeTable;
use crate::Result;
use crate::function::GENERATED_COLUMN_METADATA_KEY;
use crate::function::schema_admission::reject_caller_authored_generated_column_schema;
use crate::{Error, Result};
/// Reject caller-authored schema-bearing `add_columns` transforms that carry
/// reserved generated-column top-level field metadata.
///
/// Borrows without consuming the transform: Stream is not polled, Reader is
/// not iterated, and BatchUDF mapper is not invoked. `SqlExpressions` cannot
/// carry an Arrow output schema and is accepted.
pub(crate) fn reject_caller_authored_generated_column_add_columns_transform(
transforms: &NewColumnTransform,
) -> Result<()> {
match transforms {
NewColumnTransform::BatchUDF(udf) => {
reject_caller_authored_generated_column_schema(udf.output_schema.as_ref())
}
NewColumnTransform::Stream(stream) => {
reject_caller_authored_generated_column_schema(stream.schema().as_ref())
}
NewColumnTransform::Reader(reader) => {
reject_caller_authored_generated_column_schema(reader.schema().as_ref())
}
NewColumnTransform::AllNulls(schema) => {
reject_caller_authored_generated_column_schema(schema.as_ref())
}
NewColumnTransform::SqlExpressions(_) => Ok(()),
}
}
/// Shared rejection for general-purpose field-metadata updates that name the
/// reserved generated-column definition key.
///
/// Generated-column definitions are table-schema state owned by
/// create/change/refresh Jobs. The public `update_field_metadata` API must not
/// create, replace, or remove that reserved key. Both Native and Remote
/// `BaseTable` implementations call this helper so direct trait calls cannot
/// bypass the syntax guard.
pub(crate) fn reject_reserved_generated_column_metadata_key_updates(
updates: &[FieldMetadataUpdate],
) -> Result<()> {
for update in updates {
if update.metadata.contains_key(GENERATED_COLUMN_METADATA_KEY) {
return Err(reserved_generated_column_metadata_not_supported());
}
}
Ok(())
}
fn reserved_generated_column_metadata_not_supported() -> Error {
Error::NotSupported {
message: "generated column definitions are owned by create/change/refresh Jobs \
and cannot be created, replaced, or removed through update_field_metadata"
.into(),
}
}
/// Native-only state-aware guard: whole-map `replace()` on a field whose exact
/// Dataset snapshot metadata already contains the reserved generated-column key
/// would wipe that Job-owned definition even when the replacement map omits the
/// key. Detects raw key presence without decoding the payload.
fn reject_replace_that_would_remove_generated_column_metadata(
dataset: &lance::Dataset,
updates: &[FieldMetadataUpdate],
) -> Result<()> {
let schema = dataset.schema();
for update in updates {
if !update.replace {
continue;
}
let Some(fields) = schema.resolve_case_insensitive(&update.path) else {
continue;
};
let field = fields
.last()
.expect("resolve_case_insensitive returns a non-empty path");
if field.metadata.contains_key(GENERATED_COLUMN_METADATA_KEY) {
return Err(reserved_generated_column_metadata_not_supported());
}
}
Ok(())
}
/// The result of an add columns operation.
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
@@ -99,6 +181,7 @@ pub(crate) async fn execute_add_columns(
transforms: NewColumnTransform,
read_columns: Option<Vec<String>>,
) -> Result<AddColumnsResult> {
reject_caller_authored_generated_column_add_columns_transform(&transforms)?;
table.dataset.ensure_mutable()?;
let mut dataset = (*table.dataset.get().await?).clone();
dataset.add_columns(transforms, read_columns, None).await?;
@@ -144,8 +227,10 @@ pub(crate) async fn execute_update_field_metadata(
table: &NativeTable,
updates: &[FieldMetadataUpdate],
) -> Result<UpdateFieldMetadataResult> {
reject_reserved_generated_column_metadata_key_updates(updates)?;
table.dataset.ensure_mutable()?;
let mut dataset = (*table.dataset.get().await?).clone();
reject_replace_that_would_remove_generated_column_metadata(&dataset, updates)?;
let mut builder = dataset.update_field_metadata();
for update in updates {
@@ -163,6 +248,33 @@ pub(crate) async fn execute_update_field_metadata(
Ok(UpdateFieldMetadataResult { version })
}
/// Test-only raw installer for generated-column field metadata on Native tables.
///
/// Uses the Lance metadata commit path directly so contract fixtures can plant
/// reserved-key bytes without going through the public `update_field_metadata`
/// guard. Absent from non-test builds.
#[cfg(test)]
pub(crate) async fn install_raw_generated_column_metadata_for_tests(
table: &NativeTable,
path: impl AsRef<str>,
raw: impl Into<String>,
) -> Result<UpdateFieldMetadataResult> {
table.dataset.ensure_mutable()?;
let mut dataset = (*table.dataset.get().await?).clone();
let path = path.as_ref();
let raw = raw.into();
dataset
.update_field_metadata()
.update(
path,
[(GENERATED_COLUMN_METADATA_KEY.to_string(), Some(raw))],
)?
.await?;
let version = dataset.version().version;
table.dataset.update(dataset);
Ok(UpdateFieldMetadataResult { version })
}
#[cfg(test)]
mod tests {
use arrow_array::{Int32Array, StringArray, record_batch};
@@ -0,0 +1,306 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! Dependency-contract test for Lance A4 / A4u / A4d schema metadata attachment (B4p).
//!
//! Pins the exact generic Lance API shape LanceDB B4 will consume:
//! [`SchemaMetadataUpdates`], [`UpdateMap`], [`UpdateMapEntry`],
//! [`Transaction::with_schema_metadata_updates`], and the public
//! `with_schema_metadata_updates` methods on insert/update/delete builders.
//!
//! Also pins:
//! - A4u Update no-op: an attached field metadata patch must accompany a real
//! data change; when a predicate matches zero rows, `rows_updated == 0` and
//! the patch must not be published.
//! - A4d Delete no-op: when a predicate scans but deletes zero rows,
//! `num_deleted_rows == 0` and the attached patch must not be published.
//!
//! Neutral metadata keys only. No Function / UDF / Job semantics.
use std::collections::HashMap;
use std::sync::Arc;
use arrow_array::{Int32Array, RecordBatch, RecordBatchIterator, StringArray};
use arrow_schema::{DataType, Field, Schema as ArrowSchema};
use lance::Result;
use lance::dataset::transaction::{
Operation, SchemaMetadataUpdates, Transaction, UpdateMap, UpdateMapEntry,
};
use lance::dataset::{Dataset, DeleteBuilder, InsertBuilder, UpdateBuilder};
use lance_table::format::Fragment;
const FIELD_ID: i32 = 7;
const META_KEY: &str = "b4p.dependency.meta";
const META_VALUE: &str = "neutral-value";
fn field_metadata_patch(field_id: i32) -> SchemaMetadataUpdates {
SchemaMetadataUpdates {
schema_metadata_updates: None,
field_metadata_updates: HashMap::from([(
field_id,
UpdateMap {
update_entries: vec![UpdateMapEntry {
key: META_KEY.to_string(),
value: Some(META_VALUE.to_string()),
}],
replace: false,
},
)]),
}
}
async fn write_neutral_fixture(uri: &str) -> Dataset {
let schema = Arc::new(ArrowSchema::new(vec![
Field::new("id", DataType::Int32, false),
Field::new("value", DataType::Utf8, false),
]));
let batch = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Int32Array::from(vec![1, 2, 3])),
Arc::new(StringArray::from(vec!["a", "b", "c"])),
],
)
.unwrap();
Dataset::write(RecordBatchIterator::new(vec![Ok(batch)], schema), uri, None)
.await
.expect("fixture dataset must write")
}
/// Compile-time proof that InsertBuilder exposes the A4 attachment method.
#[allow(dead_code)]
fn typecheck_insert_builder_attachment<'a>(
builder: InsertBuilder<'a>,
updates: SchemaMetadataUpdates,
) -> Result<InsertBuilder<'a>> {
builder.with_schema_metadata_updates(updates)
}
/// Compile-time proof that UpdateBuilder exposes the A4 attachment method.
#[allow(dead_code)]
fn typecheck_update_builder_attachment(
builder: UpdateBuilder,
updates: SchemaMetadataUpdates,
) -> Result<UpdateBuilder> {
builder.with_schema_metadata_updates(updates)
}
/// Compile-time proof that DeleteBuilder exposes the A4 attachment method.
#[allow(dead_code)]
fn typecheck_delete_builder_attachment(
builder: DeleteBuilder,
updates: SchemaMetadataUpdates,
) -> Result<DeleteBuilder> {
builder.with_schema_metadata_updates(updates)
}
#[test]
fn append_transaction_retains_schema_metadata_updates_patch() {
let updates = field_metadata_patch(FIELD_ID);
assert!(
!updates.is_empty(),
"fixture must be a substantive non-empty field metadata patch"
);
let transaction = Transaction::new(
0,
Operation::Append {
fragments: vec![Fragment::new(1)],
},
None,
)
.with_schema_metadata_updates(updates.clone())
.expect("non-empty field metadata patch must attach to Append");
assert_eq!(transaction.schema_metadata_updates.as_ref(), Some(&updates));
let field_map = transaction
.schema_metadata_updates
.as_ref()
.expect("attached patch must be present")
.field_metadata_updates
.get(&FIELD_ID)
.expect("stable field id 7 must be present");
assert!(!field_map.replace);
assert_eq!(field_map.update_entries.len(), 1);
assert_eq!(field_map.update_entries[0].key, META_KEY);
assert_eq!(
field_map.update_entries[0].value.as_deref(),
Some(META_VALUE)
);
}
/// A4u dependency: a no-op Update (predicate matches zero rows) must not
/// publish an attached field metadata patch. Manifest version advancement is
/// unconstrained.
#[tokio::test]
async fn noop_update_does_not_publish_attached_field_metadata() {
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().join("noop_update.lance");
let uri = uri.to_str().unwrap();
let dataset = write_neutral_fixture(uri).await;
let field = dataset
.schema()
.field("value")
.expect("value column must exist");
let field_id = field.id;
assert!(
!field.metadata.contains_key(META_KEY),
"{META_KEY} must be initially absent, got {:?}",
field.metadata
);
let updates = field_metadata_patch(field_id);
assert!(
!updates.is_empty(),
"fixture must be a substantive non-empty field metadata patch"
);
let before_count = dataset.count_rows(None).await.unwrap();
assert_eq!(before_count, 3);
let result = UpdateBuilder::new(Arc::new(dataset))
.update_where("id < 0")
.unwrap()
.set("value", "'changed'")
.unwrap()
.with_schema_metadata_updates(updates)
.expect("Update attachment must construct")
.build()
.unwrap()
.execute()
.await
.expect("no-op attached Update must complete");
assert_eq!(result.rows_updated, 0, "predicate must match zero rows");
assert_eq!(
result.new_dataset.count_rows(None).await.unwrap(),
before_count,
"row count must remain unchanged"
);
assert_eq!(
result
.new_dataset
.count_rows(Some("value = 'changed'".into()))
.await
.unwrap(),
0,
"SET expression must not rewrite any rows"
);
assert_eq!(
result
.new_dataset
.count_rows(Some("value IN ('a', 'b', 'c')".into()))
.await
.unwrap(),
before_count,
"original values must remain unchanged"
);
let reopened = Dataset::open(uri).await.unwrap();
assert_eq!(reopened.count_rows(None).await.unwrap(), before_count);
assert_eq!(
reopened
.count_rows(Some("value = 'changed'".into()))
.await
.unwrap(),
0
);
let reopened_field = reopened
.schema()
.field_by_id(field_id)
.expect("stable field id must still exist");
assert!(
!reopened_field.metadata.contains_key(META_KEY),
"no-op Update must not publish attached field metadata; got {:?}",
reopened_field.metadata.get(META_KEY)
);
}
/// A4d dependency: a no-op Delete (predicate scans but matches zero rows) must
/// not publish an attached field metadata patch. Manifest version advancement
/// is unconstrained.
#[tokio::test]
async fn noop_delete_does_not_publish_attached_field_metadata() {
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().join("noop_delete.lance");
let uri = uri.to_str().unwrap();
let dataset = write_neutral_fixture(uri).await;
let field = dataset
.schema()
.field("value")
.expect("value column must exist");
let field_id = field.id;
assert!(
!field.metadata.contains_key(META_KEY),
"{META_KEY} must be initially absent, got {:?}",
field.metadata
);
let updates = field_metadata_patch(field_id);
assert!(
!updates.is_empty(),
"fixture must be a substantive non-empty field metadata patch"
);
let before_count = dataset.count_rows(None).await.unwrap();
assert_eq!(before_count, 3);
let result = DeleteBuilder::new(Arc::new(dataset), "id < 0")
.with_schema_metadata_updates(updates)
.expect("Delete attachment must construct")
.execute()
.await
.expect("no-op attached Delete must complete");
assert_eq!(result.num_deleted_rows, 0, "predicate must match zero rows");
assert_eq!(
result.new_dataset.count_rows(None).await.unwrap(),
before_count,
"row count must remain unchanged"
);
assert_eq!(
result
.new_dataset
.count_rows(Some("value IN ('a', 'b', 'c')".into()))
.await
.unwrap(),
before_count,
"original values must remain unchanged"
);
let returned_field = result
.new_dataset
.schema()
.field_by_id(field_id)
.expect("stable field id must still exist on returned dataset");
assert!(
!returned_field.metadata.contains_key(META_KEY),
"no-op Delete must not publish attached field metadata on returned dataset; got {:?}",
returned_field.metadata.get(META_KEY)
);
let reopened = Dataset::open(uri).await.unwrap();
assert_eq!(reopened.count_rows(None).await.unwrap(), before_count);
assert_eq!(
reopened
.count_rows(Some("value IN ('a', 'b', 'c')".into()))
.await
.unwrap(),
before_count,
"fresh open must preserve all original rows"
);
let reopened_field = reopened
.schema()
.field_by_id(field_id)
.expect("stable field id must still exist");
assert!(
!reopened_field.metadata.contains_key(META_KEY),
"no-op Delete must not publish attached field metadata; got {:?}",
reopened_field.metadata.get(META_KEY)
);
}
+39 -9
View File
@@ -1,8 +1,10 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
use std::collections::BTreeSet;
use std::sync::Arc;
use lance::Dataset;
use lance::dataset::UpdateBuilder as LanceUpdateBuilder;
use serde::{Deserialize, Serialize};
@@ -80,27 +82,37 @@ pub(crate) async fn execute_update(
) -> Result<UpdateResult> {
table.dataset.ensure_mutable()?;
// 1. Snapshot the current dataset
// One exact dataset supplies SET/filter planning, stable target field IDs,
// generated-column invalidation planning, the Lance UpdateBuilder, and its
// transaction basis. Do not call table schema()/version() or another get().
let dataset = table.dataset.get().await?;
// 2. Initialize the Lance Core builder
let mut builder = LanceUpdateBuilder::new(dataset);
let mut builder = LanceUpdateBuilder::new(dataset.clone());
// 3. Apply the filter (WHERE clause)
if let Some(predicate) = update.filter {
builder = builder.update_where(&predicate)?;
}
// 4. Apply the columns (SET clause)
for (column, value) in update.columns {
builder = builder.set(column, &value)?;
let columns = update.columns;
for (column, value) in &columns {
builder = builder.set(column, value)?;
}
// After Lance SET validation, resolve stable field IDs from the same
// snapshot and plan invalidation before UpdateJob writes files.
let updated_field_ids = updated_stable_field_ids(dataset.as_ref(), &columns)?;
if let Some(schema_metadata_updates) =
super::generated_column_invalidation::plan_native_update_generated_column_invalidation(
dataset.as_ref(),
updated_field_ids,
)?
{
builder = builder.with_schema_metadata_updates(schema_metadata_updates)?;
}
// 5. Execute the operation (Write new files)
let operation = builder.build()?;
let res = operation.execute().await?;
// 6. Update the table's view of the latest version
table.dataset.update(res.new_dataset.as_ref().clone());
Ok(UpdateResult {
@@ -109,6 +121,24 @@ pub(crate) async fn execute_update(
})
}
/// Resolve exact top-level SET targets to a deterministic set of stable field IDs.
fn updated_stable_field_ids(
dataset: &Dataset,
columns: &[(String, String)],
) -> Result<BTreeSet<i32>> {
let mut ids = BTreeSet::new();
for (column, _) in columns {
let field = dataset
.schema()
.field(column)
.ok_or_else(|| Error::InvalidInput {
message: format!("Column '{column}' does not exist in dataset schema"),
})?;
ids.insert(field.id);
}
Ok(ids)
}
#[cfg(test)]
mod tests {
use crate::connect;
@@ -0,0 +1,627 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! Runtime contract tests for the B4f reserved generated-metadata update guard.
//!
//! Pins that the general-purpose [`crate::table::Table::update_field_metadata`]
//! API cannot create, replace, or remove `GENERATED_COLUMN_METADATA_KEY`, and
//! that Native `replace()` cannot wipe an existing generated definition by
//! omitting the reserved key. Remote explicit-key attempts must reject before
//! transport.
use std::sync::Arc;
use arrow_array::{Int32Array, RecordBatch, StringArray};
use arrow_schema::{DataType, Field, Schema};
use futures::TryStreamExt;
use tempfile::TempDir;
use crate::connection::ConnectBuilder;
use crate::error::Error;
use crate::function::{
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition,
};
use crate::query::{ExecutableQuery, QueryBase, Select};
use crate::table::Table;
use crate::table::schema_evolution::FieldMetadataUpdate;
const GEN_OUT: &str = "gen_out";
const ORDINARY: &str = "ordinary";
const CATEGORY: &str = "category";
const FN_ID: &str = "fn.exact.b4f.guard.literal";
const MALFORMED_MARKER: &str = "SENSITIVE_B4F_GUARD_METADATA_MARKER_7c1e_d04b";
struct Fixture {
_tmp: TempDir,
table: Table,
uri: String,
}
fn literal_definition(
output_field_id: i32,
dependency_epoch: u64,
materialized_epoch: u64,
) -> 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("b4f-guard")])) as arrow_array::ArrayRef
)
.unwrap(),
)],
)
.unwrap();
GeneratedColumnDefinition::try_new(output_field_id, call, dependency_epoch, materialized_epoch)
.unwrap()
}
async fn create_ordinary_table(name: &str) -> Fixture {
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap().to_string();
let conn = ConnectBuilder::new(&uri).execute().await.unwrap();
let schema = Arc::new(Schema::new(vec![
Field::new(GEN_OUT, DataType::Int32, true),
Field::new(ORDINARY, DataType::Utf8, true),
Field::new(CATEGORY, DataType::Utf8, true),
]));
let batch = RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from(vec![1])),
Arc::new(StringArray::from(vec![Some("seed")])),
Arc::new(StringArray::from(vec![Some("A")])),
],
)
.unwrap();
let table = conn.create_table(name, batch).execute().await.unwrap();
Fixture {
_tmp: tmp,
table,
uri,
}
}
async fn plant_generated_raw(table: &Table, column: &str, raw: String) {
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
table
.as_native()
.expect("generated-column fixture planting requires a Native table"),
column,
raw,
)
.await
.expect("fixture raw generated-column metadata install must succeed");
}
async fn plant_valid_generated(table: &Table, column: &str) -> String {
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
let field_id = snapshot.field(column).expect(column).field_id();
let raw = literal_definition(field_id, 3, 3)
.to_metadata_json()
.unwrap();
plant_generated_raw(table, column, raw.clone()).await;
raw
}
async fn read_raw_generated_metadata(table: &Table, column: &str) -> Option<String> {
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
snapshot
.field(column)
.expect(column)
.field()
.metadata()
.get(GENERATED_COLUMN_METADATA_KEY)
.cloned()
}
async fn ordinary_values(table: &Table) -> Vec<String> {
let batches = table
.query()
.select(Select::columns(&[ORDINARY]))
.execute()
.await
.unwrap()
.try_collect::<Vec<_>>()
.await
.unwrap();
batches
.iter()
.flat_map(|batch| {
batch
.column(0)
.as_any()
.downcast_ref::<StringArray>()
.unwrap()
.iter()
.map(|v| v.unwrap().to_string())
})
.collect()
}
async fn reopen(uri: &str, name: &str) -> Table {
ConnectBuilder::new(uri)
.execute()
.await
.unwrap()
.open_table(name)
.execute()
.await
.unwrap()
}
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:?}"),
}
}
#[tokio::test]
async fn native_ordinary_field_explicit_reserved_key_set_rejects_and_preserves_state() {
let fixture = create_ordinary_table("b4f_ordinary_set").await;
let version_before = fixture.table.version().await.unwrap();
let rows_before = ordinary_values(&fixture.table).await;
let schema_before = fixture.table.schema().await.unwrap();
let category_md_before = schema_before
.field_with_name(CATEGORY)
.unwrap()
.metadata()
.clone();
let snapshot = fixture
.table
.generated_column_binding_snapshot()
.await
.unwrap();
let field_id = snapshot.field(CATEGORY).expect(CATEGORY).field_id();
let payload = literal_definition(field_id, 1, 1)
.to_metadata_json()
.unwrap();
let err = fixture
.table
.update_field_metadata(&[
FieldMetadataUpdate::new(CATEGORY).set(GENERATED_COLUMN_METADATA_KEY, payload.clone())
])
.await
.expect_err("explicit reserved-key set on ordinary field must reject");
assert_not_supported_redacted(&err, "ordinary reserved set", &[&payload]);
assert_eq!(fixture.table.version().await.unwrap(), version_before);
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
let schema_after = fixture.table.schema().await.unwrap();
assert_eq!(
schema_after.field_with_name(CATEGORY).unwrap().metadata(),
&category_md_before
);
assert!(
read_raw_generated_metadata(&fixture.table, CATEGORY)
.await
.is_none()
);
}
#[tokio::test]
async fn native_generated_field_explicit_remove_rejects_and_preserves_raw() {
let fixture = create_ordinary_table("b4f_gen_remove").await;
let planted = plant_valid_generated(&fixture.table, GEN_OUT).await;
let version_before = fixture.table.version().await.unwrap();
let err = fixture
.table
.update_field_metadata(&[
FieldMetadataUpdate::new(GEN_OUT).remove(GENERATED_COLUMN_METADATA_KEY)
])
.await
.expect_err("explicit reserved-key remove must reject");
assert_not_supported_redacted(&err, "generated remove", &[&planted]);
assert_eq!(fixture.table.version().await.unwrap(), version_before);
assert_eq!(
read_raw_generated_metadata(&fixture.table, GEN_OUT)
.await
.as_deref(),
Some(planted.as_str())
);
let fresh = reopen(&fixture.uri, "b4f_gen_remove").await;
assert_eq!(
read_raw_generated_metadata(&fresh, GEN_OUT)
.await
.as_deref(),
Some(planted.as_str())
);
}
#[tokio::test]
async fn native_generated_field_explicit_replacement_rejects_and_preserves_raw() {
let fixture = create_ordinary_table("b4f_gen_replace_value").await;
let planted = plant_valid_generated(&fixture.table, GEN_OUT).await;
let version_before = fixture.table.version().await.unwrap();
let replacement = literal_definition(
fixture
.table
.generated_column_binding_snapshot()
.await
.unwrap()
.field(GEN_OUT)
.unwrap()
.field_id(),
9,
1,
)
.to_metadata_json()
.unwrap();
assert_ne!(replacement, planted);
let err = fixture
.table
.update_field_metadata(&[FieldMetadataUpdate::new(GEN_OUT)
.set(GENERATED_COLUMN_METADATA_KEY, replacement.clone())])
.await
.expect_err("explicit reserved-key replacement must reject");
assert_not_supported_redacted(&err, "generated replace value", &[&planted, &replacement]);
assert_eq!(fixture.table.version().await.unwrap(), version_before);
let fresh = reopen(&fixture.uri, "b4f_gen_replace_value").await;
assert_eq!(
read_raw_generated_metadata(&fresh, GEN_OUT)
.await
.as_deref(),
Some(planted.as_str())
);
}
#[tokio::test]
async fn native_generated_field_replace_with_ordinary_metadata_rejects_and_preserves_raw() {
let fixture = create_ordinary_table("b4f_gen_replace_map").await;
let planted = plant_valid_generated(&fixture.table, GEN_OUT).await;
let version_before = fixture.table.version().await.unwrap();
let err = fixture
.table
.update_field_metadata(&[FieldMetadataUpdate::new(GEN_OUT)
.replace()
.set("unit", "label")])
.await
.expect_err("replace() that would wipe generated definition must reject");
assert_not_supported_redacted(&err, "generated replace map", &[&planted]);
assert_eq!(fixture.table.version().await.unwrap(), version_before);
let fresh = reopen(&fixture.uri, "b4f_gen_replace_map").await;
assert_eq!(
read_raw_generated_metadata(&fresh, GEN_OUT)
.await
.as_deref(),
Some(planted.as_str())
);
}
#[tokio::test]
async fn native_mixed_batch_rejects_atomically_no_partial_commit() {
let fixture = create_ordinary_table("b4f_mixed_batch").await;
let planted = plant_valid_generated(&fixture.table, GEN_OUT).await;
let version_before = fixture.table.version().await.unwrap();
let err = fixture
.table
.update_field_metadata(&[
FieldMetadataUpdate::new(CATEGORY).set("unit", "label"),
FieldMetadataUpdate::new(GEN_OUT).remove(GENERATED_COLUMN_METADATA_KEY),
])
.await
.expect_err("mixed batch with forbidden update must reject all-or-none");
assert_not_supported_redacted(&err, "mixed forbidden second", &[&planted]);
assert_eq!(fixture.table.version().await.unwrap(), version_before);
let schema = fixture.table.schema().await.unwrap();
assert!(
!schema
.field_with_name(CATEGORY)
.unwrap()
.metadata()
.contains_key("unit"),
"ordinary metadata must not partially commit"
);
let err = fixture
.table
.update_field_metadata(&[
FieldMetadataUpdate::new(GEN_OUT).set(GENERATED_COLUMN_METADATA_KEY, planted.clone()),
FieldMetadataUpdate::new(CATEGORY).set("unit", "label"),
])
.await
.expect_err("mixed batch with forbidden update first must reject all-or-none");
assert_not_supported_redacted(&err, "mixed forbidden first", &[&planted]);
let fresh = reopen(&fixture.uri, "b4f_mixed_batch").await;
assert_eq!(fresh.version().await.unwrap(), version_before);
assert_eq!(
read_raw_generated_metadata(&fresh, GEN_OUT)
.await
.as_deref(),
Some(planted.as_str())
);
let fresh_schema = fresh.schema().await.unwrap();
assert!(
!fresh_schema
.field_with_name(CATEGORY)
.unwrap()
.metadata()
.contains_key("unit")
);
}
#[tokio::test]
async fn native_malformed_generated_raw_replace_rejects_redacted_and_preserves_raw() {
let fixture = create_ordinary_table("b4f_malformed_replace").await;
let field_id = fixture
.table
.generated_column_binding_snapshot()
.await
.unwrap()
.field(GEN_OUT)
.unwrap()
.field_id();
let planted_raw = format!(
r#"{{"format_version":1,"output_field_id":{field_id},"function_call":{MALFORMED_MARKER},"dependency_epoch":1,"materialized_epoch":1}}"#
);
assert!(planted_raw.contains(MALFORMED_MARKER));
plant_generated_raw(&fixture.table, GEN_OUT, planted_raw.clone()).await;
let version_before = fixture.table.version().await.unwrap();
let err = fixture
.table
.update_field_metadata(&[FieldMetadataUpdate::new(GEN_OUT)
.replace()
.set("unit", "label")])
.await
.expect_err("malformed generated raw must still block replace()");
assert_not_supported_redacted(&err, "malformed replace", &[&planted_raw]);
assert_eq!(fixture.table.version().await.unwrap(), version_before);
let fresh = reopen(&fixture.uri, "b4f_malformed_replace").await;
assert_eq!(
read_raw_generated_metadata(&fresh, GEN_OUT)
.await
.as_deref(),
Some(planted_raw.as_str())
);
}
#[tokio::test]
async fn native_ordinary_metadata_merge_set_remove_replace_still_work() {
let fixture = create_ordinary_table("b4f_ordinary_controls").await;
fixture
.table
.update_field_metadata(&[FieldMetadataUpdate::new(CATEGORY)
.set("unit", "label")
.set("pii", "false")])
.await
.unwrap();
let md = fixture
.table
.schema()
.await
.unwrap()
.field_with_name(CATEGORY)
.unwrap()
.metadata()
.clone();
assert_eq!(md.get("unit").map(String::as_str), Some("label"));
assert_eq!(md.get("pii").map(String::as_str), Some("false"));
fixture
.table
.update_field_metadata(&[FieldMetadataUpdate::new(CATEGORY)
.set("source", "import")
.remove("pii")])
.await
.unwrap();
let md = fixture
.table
.schema()
.await
.unwrap()
.field_with_name(CATEGORY)
.unwrap()
.metadata()
.clone();
assert_eq!(md.get("unit").map(String::as_str), Some("label"));
assert_eq!(md.get("source").map(String::as_str), Some("import"));
assert!(!md.contains_key("pii"));
fixture
.table
.update_field_metadata(&[FieldMetadataUpdate::new(CATEGORY)
.replace()
.set("only", "kept")])
.await
.unwrap();
let md = fixture
.table
.schema()
.await
.unwrap()
.field_with_name(CATEGORY)
.unwrap()
.metadata()
.clone();
assert_eq!(md.len(), 1);
assert_eq!(md.get("only").map(String::as_str), Some("kept"));
assert!(
read_raw_generated_metadata(&fixture.table, CATEGORY)
.await
.is_none()
);
}
#[cfg(feature = "remote")]
mod remote_explicit_key_guard {
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use async_trait::async_trait;
use super::*;
use crate::Error;
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-Test-Header".to_string(),
"must-not-be-requested".to_string(),
)]))
}
}
fn panic_handler(
calls: Arc<AtomicUsize>,
) -> impl Fn(reqwest::Request) -> http::Response<String> + Clone + Send + Sync + 'static {
move |_request| {
calls.fetch_add(1, Ordering::SeqCst);
panic!("remote reserved-key update must not invoke the HTTP handler");
}
}
#[tokio::test]
async fn remote_explicit_set_rejects_before_handler_and_header_provider() {
let handler_calls = Arc::new(AtomicUsize::new(0));
let header_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 table = Table::new_with_handler_and_config(
"my_table",
panic_handler(handler_calls.clone()),
config,
);
let err = table
.update_field_metadata(&[FieldMetadataUpdate::new(CATEGORY)
.set(GENERATED_COLUMN_METADATA_KEY, r#"{"format_version":1}"#)])
.await
.expect_err("remote explicit reserved-key set must reject");
assert!(
matches!(err, Error::NotSupported { .. }),
"expected NotSupported, got {err:?}"
);
assert_eq!(handler_calls.load(Ordering::SeqCst), 0);
assert_eq!(header_calls.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn remote_explicit_remove_rejects_before_handler_and_header_provider() {
let handler_calls = Arc::new(AtomicUsize::new(0));
let header_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 table = Table::new_with_handler_and_config(
"my_table",
panic_handler(handler_calls.clone()),
config,
);
let err = table
.update_field_metadata(&[
FieldMetadataUpdate::new(CATEGORY).remove(GENERATED_COLUMN_METADATA_KEY)
])
.await
.expect_err("remote explicit reserved-key remove must reject");
assert!(
matches!(err, Error::NotSupported { .. }),
"expected NotSupported, got {err:?}"
);
assert_eq!(handler_calls.load(Ordering::SeqCst), 0);
assert_eq!(header_calls.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn remote_ordinary_metadata_update_sends_exact_body_and_succeeds() {
let table = Table::new_with_handler("my_table", |request| {
assert_eq!(request.method(), "POST");
assert_eq!(
request.url().path(),
"/v1/table/my_table/update_field_metadata/"
);
let body = request
.body()
.expect("ordinary update must send a body")
.as_bytes()
.expect("body is in-memory");
let parsed: serde_json::Value = serde_json::from_slice(body).unwrap();
assert_eq!(
parsed,
serde_json::json!({
"updates": [{
"path": "category",
"metadata": { "unit": "label" },
"replace": false
}]
})
);
http::Response::builder()
.status(200)
.body(r#"{"version": 7}"#.to_string())
.unwrap()
});
let result = table
.update_field_metadata(&[FieldMetadataUpdate::new(CATEGORY).set("unit", "label")])
.await
.unwrap();
assert_eq!(result.version, 7);
}
}
File diff suppressed because it is too large Load Diff
+36 -13
View File
@@ -10,7 +10,7 @@ use arrow_array::{
use arrow_schema::{DataType, Field, Fields, Schema};
use futures::TryStreamExt;
use lance::Dataset;
use lance_file::version::LanceFileVersion;
use lance_file::version::{ConcreteFileVersion, LanceFileVersion};
use lancedb::{
Connection, Error, Result, Table,
blob::{BlobRangeRequest, blob},
@@ -61,7 +61,7 @@ async fn create_inline_blob_table(
Ok(table)
}
async fn storage_format_version(table: &Table) -> LanceFileVersion {
async fn storage_format_version(table: &Table) -> ConcreteFileVersion {
table
.as_native()
.unwrap()
@@ -69,9 +69,18 @@ async fn storage_format_version(table: &Table) -> LanceFileVersion {
.await
.unwrap()
.data_storage_format
.lance_file_version()
.unwrap()
.resolve()
.lance_file_format()
}
/// Blob v2 storage capability for the current concrete formats.
///
/// Exact formats deliberately have no Ord: capability is not implied by release
/// order. Enumerate every current concrete variant explicitly.
fn supports_blob_v2_storage(version: ConcreteFileVersion) -> bool {
match version {
ConcreteFileVersion::V2_2 | ConcreteFileVersion::V2_3 => true,
ConcreteFileVersion::V1 | ConcreteFileVersion::V2_0 | ConcreteFileVersion::V2_1 => false,
}
}
async fn uses_stable_row_ids(table: &Table) -> bool {
@@ -112,7 +121,9 @@ async fn declaring_blob_column_bumps_format_and_enables_stable_row_ids() -> Resu
.execute()
.await?;
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
assert!(supports_blob_v2_storage(
storage_format_version(&table).await
));
assert!(uses_stable_row_ids(&table).await);
Ok(())
}
@@ -127,7 +138,9 @@ async fn explicit_stable_row_id_setting_wins_over_blob_default() -> Result<()> {
.execute()
.await?;
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
assert!(supports_blob_v2_storage(
storage_format_version(&table).await
));
assert!(!uses_stable_row_ids(&table).await);
Ok(())
}
@@ -139,7 +152,9 @@ async fn non_blob_table_keeps_default_format_and_row_id_setting() -> Result<()>
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int64, false)]));
let table = db.create_empty_table("t", schema).execute().await?;
assert!(storage_format_version(&table).await < LanceFileVersion::V2_2);
assert!(!supports_blob_v2_storage(
storage_format_version(&table).await
));
assert!(!uses_stable_row_ids(&table).await);
Ok(())
}
@@ -171,7 +186,9 @@ async fn creating_with_blob_data_bumps_format() -> Result<()> {
.unwrap();
let table = db.create_table("t", batch).execute().await?;
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
assert!(supports_blob_v2_storage(
storage_format_version(&table).await
));
assert!(uses_stable_row_ids(&table).await);
assert_eq!(table.count_rows(None).await?, 1);
Ok(())
@@ -281,7 +298,9 @@ async fn connection_level_stable_row_id_setting_wins_over_blob_default() -> Resu
.execute()
.await?;
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
assert!(supports_blob_v2_storage(
storage_format_version(&table).await
));
assert!(!uses_stable_row_ids(&table).await);
Ok(())
}
@@ -297,7 +316,9 @@ async fn namespace_create_applies_blob_defaults() -> Result<()> {
.execute()
.await?;
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
assert!(supports_blob_v2_storage(
storage_format_version(&table).await
));
assert!(uses_stable_row_ids(&table).await);
Ok(())
}
@@ -474,7 +495,9 @@ async fn fetch_blobs_round_trips_nested_blob_column() -> Result<()> {
let batch = RecordBatch::try_new(schema, vec![Arc::new(info_array) as ArrayRef]).unwrap();
let table = db.create_table("t", batch).execute().await?;
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
assert!(supports_blob_v2_storage(
storage_format_version(&table).await
));
assert!(uses_stable_row_ids(&table).await);
let ids = collect_row_ids(&table).await?;
@@ -1305,7 +1328,7 @@ async fn optimize_preserves_blob_v2_null_and_empty_distinction() -> Result<()> {
.await?;
table.add(null_empty_input_batch()).execute().await?;
assert!(
storage_format_version(&table).await >= LanceFileVersion::V2_2,
supports_blob_v2_storage(storage_format_version(&table).await),
"blob v2 columns require storage >= 2.2"
);
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,837 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! Contract tests for FunctionDefinition registration input (FF-007 / B1c).
//!
//! These tests pin the intended public surface under [`lancedb::function`] for
//! Python definition transport only. They intentionally fail to compile until
//! that API exists.
//!
//! Rejection cases are judged by `Result` structure (`is_err` / `is_ok`), never
//! by diagnostic message substrings.
use std::collections::BTreeSet;
use arrow_schema::DataType;
use lancedb::Result;
use lancedb::function::{
Function, FunctionCapability, FunctionDefinition, FunctionId, FunctionOutput,
FunctionParameter, FunctionSignature, PythonFunctionDefinition,
};
use serde_json::Value;
fn sample_signature() -> Result<FunctionSignature> {
FunctionSignature::try_new(
vec![
FunctionParameter::new("text", DataType::Utf8),
FunctionParameter::new("limit", DataType::Int32),
],
FunctionOutput::new(DataType::Utf8, true),
)
}
fn sample_source() -> &'static str {
"def normalize(text, limit):\n return text[:limit]\n"
}
fn sample_python_definition() -> Result<PythonFunctionDefinition> {
PythonFunctionDefinition::try_new(
"normalize_mod",
"normalize",
sample_source(),
"3.12",
vec!["Unidecode==1.3.8".to_string()],
)
}
fn sample_capabilities() -> Result<Vec<FunctionCapability>> {
Ok(vec![
FunctionCapability::try_network("https://api.example.com")?,
FunctionCapability::try_secret("secret://team/api-token", "API_TOKEN")?,
])
}
fn sample_definition() -> Result<FunctionDefinition> {
FunctionDefinition::try_new(
sample_signature()?,
sample_python_definition()?,
sample_capabilities()?,
)
}
fn assert_json_object_keys_exact(value: &Value, expected: &[&str]) {
let object = value
.as_object()
.unwrap_or_else(|| panic!("expected JSON object, got {value}"));
let keys: BTreeSet<&str> = object.keys().map(|k| k.as_str()).collect();
let expected: BTreeSet<&str> = expected.iter().copied().collect();
assert_eq!(
keys, expected,
"JSON object key set must match exactly (iteration order is not a contract); got {keys:?}, expected {expected:?} in {value}"
);
}
fn assert_json_object_keys_subset(value: &Value, allowed: &[&str]) {
let object = value
.as_object()
.unwrap_or_else(|| panic!("expected JSON object, got {value}"));
for key in object.keys() {
assert!(
allowed.contains(&key.as_str()),
"unexpected JSON key `{key}` in {value}"
);
}
}
/// Identity / lineage / artifact / runtime fields that must not appear as public
/// object keys on FunctionDefinition wire. Parameter object key `name` is not
/// listed here: FunctionSignature legitimately uses it under `signature.parameters`.
const FORBIDDEN_DEFINITION_KEYS: &[&str] = &[
"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",
];
fn assert_object_keys_not_forbidden(value: &Value, context: &str) {
let object = value.as_object().unwrap_or_else(|| {
panic!("expected JSON object at {context}, got {value}");
});
for key in object.keys() {
assert!(
!FORBIDDEN_DEFINITION_KEYS.contains(&key.as_str()),
"FunctionDefinition wire must not contain forbidden key `{key}` at {context}: {value}"
);
}
}
fn assert_forbidden_definition_keys_absent(value: &Value) {
// Top-level definition key set (order-independent).
assert_json_object_keys_exact(
value,
&[
"format_version",
"signature",
"implementation",
"capabilities",
],
);
assert_object_keys_not_forbidden(value, "definition");
// Catalog / function identity name is absent at definition root; parameter
// `name` is allowed only under signature.parameters.
assert!(
value.get("name").is_none(),
"top-level FunctionDefinition wire must not contain catalog/function identity key `name`: {value}"
);
assert!(
value.get("catalog_name").is_none(),
"top-level FunctionDefinition wire must not contain `catalog_name`: {value}"
);
let implementation = value.get("implementation").expect("implementation object");
assert_json_object_keys_exact(
implementation,
&["kind", "module", "callable", "source", "python", "packages"],
);
assert_object_keys_not_forbidden(implementation, "implementation");
assert!(
implementation.get("name").is_none(),
"implementation must not contain catalog/function identity key `name`: {implementation}"
);
assert!(
implementation.get("catalog_name").is_none(),
"implementation must not contain `catalog_name`: {implementation}"
);
let capabilities = value
.get("capabilities")
.and_then(Value::as_array)
.expect("capabilities array");
for (idx, capability) in capabilities.iter().enumerate() {
let kind = capability
.get("kind")
.and_then(Value::as_str)
.unwrap_or_else(|| panic!("capabilities[{idx}] missing kind"));
match kind {
"network" => assert_json_object_keys_exact(capability, &["kind", "origin"]),
"secret" => assert_json_object_keys_exact(
capability,
&["kind", "reference", "environment_variable"],
),
other => panic!("unexpected capability kind `{other}` in contract fixture"),
}
let context = format!("capabilities[{idx}]");
assert_object_keys_not_forbidden(capability, &context);
assert!(
capability.get("name").is_none(),
"{context} must not contain catalog/function identity key `name`: {capability}"
);
assert!(
capability.get("catalog_name").is_none(),
"{context} must not contain `catalog_name`: {capability}"
);
}
// Signature may carry parameter objects with key `name`. Still reject
// identity/lineage/runtime keys and function-identity `name`/`catalog_name`
// on the signature and output objects themselves.
let signature = value.get("signature").expect("signature object");
assert_object_keys_not_forbidden(signature, "signature");
assert!(
signature.get("name").is_none(),
"signature object must not contain catalog/function identity key `name`: {signature}"
);
assert!(
signature.get("catalog_name").is_none(),
"signature object must not contain `catalog_name`: {signature}"
);
if let Some(parameters) = signature.get("parameters").and_then(Value::as_array) {
for (idx, parameter) in parameters.iter().enumerate() {
let context = format!("signature.parameters[{idx}]");
assert_object_keys_not_forbidden(parameter, &context);
assert!(
parameter.get("catalog_name").is_none(),
"{context} must not contain `catalog_name`: {parameter}"
);
// `name` is intentionally allowed on parameter objects.
assert!(
parameter.get("name").is_some(),
"{context} must include parameter `name`"
);
}
}
if let Some(output) = signature.get("output") {
assert_object_keys_not_forbidden(output, "signature.output");
assert!(
output.get("name").is_none(),
"signature.output must not contain catalog/function identity key `name`: {output}"
);
assert!(
output.get("catalog_name").is_none(),
"signature.output must not contain `catalog_name`: {output}"
);
}
}
#[test]
fn definition_json_round_trip_pins_exact_wire_shape_and_order() -> Result<()> {
let definition = sample_definition()?;
assert_eq!(definition.signature().parameters().len(), 2);
assert_eq!(definition.signature().parameters()[0].name(), "text");
assert_eq!(
definition.signature().parameters()[0].data_type(),
&DataType::Utf8
);
assert_eq!(definition.signature().parameters()[1].name(), "limit");
assert_eq!(
definition.signature().parameters()[1].data_type(),
&DataType::Int32
);
assert_eq!(definition.signature().output().data_type(), &DataType::Utf8);
assert!(definition.signature().output().nullable());
let python = definition.python_definition();
assert_eq!(python.module(), "normalize_mod");
assert_eq!(python.callable(), "normalize");
assert_eq!(python.source(), sample_source());
assert_eq!(python.python(), "3.12");
assert_eq!(python.packages(), &["Unidecode==1.3.8".to_string()]);
let capabilities = definition.capabilities();
assert_eq!(capabilities.len(), 2);
assert_eq!(capabilities[0].origin(), Some("https://api.example.com"));
assert_eq!(capabilities[0].reference(), None);
assert_eq!(capabilities[0].environment_variable(), None);
assert_eq!(capabilities[1].reference(), Some("secret://team/api-token"));
assert_eq!(capabilities[1].environment_variable(), Some("API_TOKEN"));
assert_eq!(capabilities[1].origin(), None);
let json = serde_json::to_value(&definition).expect("serialize FunctionDefinition");
assert_json_object_keys_exact(
&json,
&[
"format_version",
"signature",
"implementation",
"capabilities",
],
);
assert_eq!(json["format_version"], 1);
let signature = json
.get("signature")
.and_then(Value::as_object)
.expect("signature object");
assert_json_object_keys_subset(&Value::Object(signature.clone()), &["parameters", "output"]);
let parameters = signature
.get("parameters")
.and_then(Value::as_array)
.expect("parameters array");
assert_eq!(parameters.len(), 2);
assert_eq!(parameters[0]["name"], Value::String("text".into()));
assert_eq!(parameters[1]["name"], Value::String("limit".into()));
for parameter in parameters {
assert_json_object_keys_subset(parameter, &["name", "data_type_ipc"]);
assert!(
parameter
.get("data_type_ipc")
.and_then(Value::as_str)
.is_some_and(|s| !s.is_empty()),
"parameter data_type_ipc must be non-empty base64"
);
}
let output = signature
.get("output")
.and_then(Value::as_object)
.expect("output object");
assert_json_object_keys_subset(
&Value::Object(output.clone()),
&["data_type_ipc", "nullable"],
);
assert_eq!(output.get("nullable"), Some(&Value::Bool(true)));
let implementation = json.get("implementation").expect("implementation object");
assert_json_object_keys_exact(
implementation,
&["kind", "module", "callable", "source", "python", "packages"],
);
assert_eq!(implementation["kind"], Value::String("python".into()));
assert_eq!(
implementation["module"],
Value::String("normalize_mod".into())
);
assert_eq!(
implementation["callable"],
Value::String("normalize".into())
);
assert_eq!(
implementation["source"],
Value::String(sample_source().into())
);
assert_eq!(implementation["python"], Value::String("3.12".into()));
assert_eq!(
implementation["packages"],
Value::Array(vec![Value::String("Unidecode==1.3.8".into())])
);
let capabilities_json = json
.get("capabilities")
.and_then(Value::as_array)
.expect("capabilities array");
assert_eq!(capabilities_json.len(), 2);
assert_json_object_keys_exact(&capabilities_json[0], &["kind", "origin"]);
assert_eq!(
capabilities_json[0]["kind"],
Value::String("network".into())
);
assert_eq!(
capabilities_json[0]["origin"],
Value::String("https://api.example.com".into())
);
assert_json_object_keys_exact(
&capabilities_json[1],
&["kind", "reference", "environment_variable"],
);
assert_eq!(capabilities_json[1]["kind"], Value::String("secret".into()));
assert_eq!(
capabilities_json[1]["reference"],
Value::String("secret://team/api-token".into())
);
assert_eq!(
capabilities_json[1]["environment_variable"],
Value::String("API_TOKEN".into())
);
// Same ordered signature IPC representation as Function handle transport.
let function = Function::new(FunctionId::try_new("fn.wire.compare")?, sample_signature()?);
let function_json = serde_json::to_value(&function).expect("serialize Function");
assert_eq!(json["signature"], function_json["signature"]);
let restored: FunctionDefinition =
serde_json::from_value(json.clone()).expect("deserialize FunctionDefinition");
assert_eq!(
restored.signature().parameters()[0].name(),
definition.signature().parameters()[0].name()
);
assert_eq!(
restored.signature().parameters()[0].data_type(),
definition.signature().parameters()[0].data_type()
);
assert_eq!(
restored.signature().parameters()[1].name(),
definition.signature().parameters()[1].name()
);
assert_eq!(
restored.signature().parameters()[1].data_type(),
definition.signature().parameters()[1].data_type()
);
assert_eq!(
restored.signature().output().data_type(),
definition.signature().output().data_type()
);
assert_eq!(
restored.signature().output().nullable(),
definition.signature().output().nullable()
);
assert_eq!(
restored.python_definition().module(),
definition.python_definition().module()
);
assert_eq!(
restored.python_definition().callable(),
definition.python_definition().callable()
);
assert_eq!(
restored.python_definition().source(),
definition.python_definition().source()
);
assert_eq!(
restored.python_definition().python(),
definition.python_definition().python()
);
assert_eq!(
restored.python_definition().packages(),
definition.python_definition().packages()
);
assert_eq!(restored.capabilities().len(), 2);
assert_eq!(
restored.capabilities()[0].origin(),
definition.capabilities()[0].origin()
);
assert_eq!(
restored.capabilities()[1].reference(),
definition.capabilities()[1].reference()
);
assert_eq!(
restored.capabilities()[1].environment_variable(),
definition.capabilities()[1].environment_variable()
);
// Package and capability order are part of the structural wire.
let multi_pkg = FunctionDefinition::try_new(
sample_signature()?,
PythonFunctionDefinition::try_new(
"normalize_mod",
"normalize",
sample_source(),
"3.12",
vec![
"Unidecode==1.3.8".to_string(),
"requests==2.32.3".to_string(),
],
)?,
vec![
FunctionCapability::try_secret("secret://team/api-token", "API_TOKEN")?,
FunctionCapability::try_network("https://api.example.com")?,
FunctionCapability::try_network("https://other.example.com")?,
],
)?;
let multi_json = serde_json::to_value(&multi_pkg).expect("serialize multi-order definition");
assert_eq!(
multi_json["implementation"]["packages"],
Value::Array(vec![
Value::String("Unidecode==1.3.8".into()),
Value::String("requests==2.32.3".into()),
])
);
assert_eq!(
multi_json["capabilities"][0]["kind"],
Value::String("secret".into())
);
assert_eq!(
multi_json["capabilities"][1]["origin"],
Value::String("https://api.example.com".into())
);
assert_eq!(
multi_json["capabilities"][2]["origin"],
Value::String("https://other.example.com".into())
);
let multi_restored: FunctionDefinition =
serde_json::from_value(multi_json.clone()).expect("deserialize multi-order definition");
assert_eq!(
multi_restored.python_definition().packages(),
&[
"Unidecode==1.3.8".to_string(),
"requests==2.32.3".to_string()
]
);
assert_eq!(
multi_restored.capabilities()[0].reference(),
Some("secret://team/api-token")
);
assert_eq!(
multi_restored.capabilities()[1].origin(),
Some("https://api.example.com")
);
assert_eq!(
multi_restored.capabilities()[2].origin(),
Some("https://other.example.com")
);
// Byte-for-byte repeated serde_json encoding for the same value.
let encoded_a = serde_json::to_string(&definition).expect("encode a");
let encoded_b = serde_json::to_string(&definition).expect("encode b");
assert_eq!(encoded_a, encoded_b);
assert_eq!(
serde_json::to_value(&restored).expect("re-serialize restored"),
json
);
assert_eq!(
serde_json::to_string(&multi_restored).expect("re-encode multi"),
serde_json::to_string(&multi_pkg).expect("encode multi")
);
Ok(())
}
#[test]
fn definition_wire_excludes_identity_lineage_artifact_and_runtime_fields() -> Result<()> {
let definition = sample_definition()?;
let json = serde_json::to_value(&definition).expect("serialize FunctionDefinition");
// Forbidden public fields are object keys at structural levels only.
// Do not substring-scan encoded JSON: user source/reference/package text
// may legitimately contain those tokens.
assert_forbidden_definition_keys_absent(&json);
Ok(())
}
#[test]
fn definition_decode_fails_closed_for_unknown_version_fields_and_kinds() -> Result<()> {
let definition = sample_definition()?;
let json = serde_json::to_value(&definition).expect("serialize FunctionDefinition");
let mut unknown_version = json.clone();
unknown_version["format_version"] = Value::from(2);
assert!(
serde_json::from_value::<FunctionDefinition>(unknown_version).is_err(),
"format_version other than 1 must fail closed"
);
let mut unknown_outer = json.clone();
unknown_outer
.as_object_mut()
.unwrap()
.insert("unexpected_field".into(), Value::Bool(true));
assert!(
serde_json::from_value::<FunctionDefinition>(unknown_outer).is_err(),
"unknown outer field must fail closed"
);
let mut unknown_implementation_field = json.clone();
unknown_implementation_field["implementation"]
.as_object_mut()
.unwrap()
.insert("entrypoint".into(), Value::String("main".into()));
assert!(
serde_json::from_value::<FunctionDefinition>(unknown_implementation_field).is_err(),
"unknown nested implementation field must fail closed"
);
let mut unknown_capability_field = json.clone();
unknown_capability_field["capabilities"][0]
.as_object_mut()
.unwrap()
.insert("headers".into(), Value::Object(Default::default()));
assert!(
serde_json::from_value::<FunctionDefinition>(unknown_capability_field).is_err(),
"unknown nested capability field must fail closed"
);
let mut unknown_implementation_kind = json.clone();
unknown_implementation_kind["implementation"]["kind"] = Value::String("builtin".into());
assert!(
serde_json::from_value::<FunctionDefinition>(unknown_implementation_kind).is_err(),
"unknown implementation kind must fail closed"
);
let mut unknown_capability_kind = json.clone();
unknown_capability_kind["capabilities"][0]["kind"] = Value::String("filesystem".into());
assert!(
serde_json::from_value::<FunctionDefinition>(unknown_capability_kind).is_err(),
"unknown capability kind must fail closed"
);
Ok(())
}
#[test]
fn constructors_and_decode_reject_empty_fields_and_duplicate_packages() -> Result<()> {
let signature = sample_signature()?;
let packages = vec!["Unidecode==1.3.8".to_string()];
let capabilities = sample_capabilities()?;
assert!(
PythonFunctionDefinition::try_new(
"",
"normalize",
sample_source(),
"3.12",
packages.clone()
)
.is_err(),
"empty module must be rejected"
);
assert!(
PythonFunctionDefinition::try_new(
"normalize_mod",
"",
sample_source(),
"3.12",
packages.clone()
)
.is_err(),
"empty callable must be rejected"
);
assert!(
PythonFunctionDefinition::try_new(
"normalize_mod",
"normalize",
"",
"3.12",
packages.clone()
)
.is_err(),
"empty source must be rejected"
);
assert!(
PythonFunctionDefinition::try_new(
"normalize_mod",
"normalize",
sample_source(),
"",
packages.clone()
)
.is_err(),
"empty python runtime request must be rejected"
);
assert!(
PythonFunctionDefinition::try_new(
"normalize_mod",
"normalize",
sample_source(),
"3.12",
vec!["".to_string()],
)
.is_err(),
"empty package requirement must be rejected"
);
assert!(
PythonFunctionDefinition::try_new(
"normalize_mod",
"normalize",
sample_source(),
"3.12",
vec![
"Unidecode==1.3.8".to_string(),
"Unidecode==1.3.8".to_string(),
],
)
.is_err(),
"duplicate package requirements must be rejected"
);
assert!(
FunctionCapability::try_network("").is_err(),
"empty network origin must be rejected"
);
assert!(
FunctionCapability::try_secret("", "API_TOKEN").is_err(),
"empty secret reference must be rejected"
);
assert!(
FunctionCapability::try_secret("secret://team/api-token", "").is_err(),
"empty secret environment variable must be rejected"
);
// Decode path must enforce the same emptiness / uniqueness rules.
let definition = FunctionDefinition::try_new(
signature.clone(),
sample_python_definition()?,
capabilities.clone(),
)?;
let json = serde_json::to_value(&definition).expect("serialize FunctionDefinition");
for (pointer, empty) in [
("/implementation/module", ""),
("/implementation/callable", ""),
("/implementation/source", ""),
("/implementation/python", ""),
("/capabilities/0/origin", ""),
("/capabilities/1/reference", ""),
("/capabilities/1/environment_variable", ""),
] {
let mut invalid = json.clone();
let target = invalid
.pointer_mut(pointer)
.unwrap_or_else(|| panic!("missing pointer {pointer}"));
*target = Value::String(empty.into());
assert!(
serde_json::from_value::<FunctionDefinition>(invalid).is_err(),
"decode must reject empty value at {pointer}"
);
}
let mut empty_package = json.clone();
empty_package["implementation"]["packages"] = Value::Array(vec![Value::String("".into())]);
assert!(
serde_json::from_value::<FunctionDefinition>(empty_package).is_err(),
"decode must reject empty package requirement"
);
let mut duplicate_packages = json.clone();
duplicate_packages["implementation"]["packages"] = Value::Array(vec![
Value::String("Unidecode==1.3.8".into()),
Value::String("Unidecode==1.3.8".into()),
]);
assert!(
serde_json::from_value::<FunctionDefinition>(duplicate_packages).is_err(),
"decode must reject duplicate package requirements"
);
// Keep the constructor path for FunctionDefinition itself structurally valid
// when children are valid; emptiness is owned by child constructors above.
assert!(
FunctionDefinition::try_new(signature, sample_python_definition()?, capabilities).is_ok()
);
Ok(())
}
#[test]
fn secret_capability_wire_rejects_plaintext_value_fields() -> Result<()> {
let definition = sample_definition()?;
let json = serde_json::to_value(&definition).expect("serialize FunctionDefinition");
let secret = &json["capabilities"][1];
assert_json_object_keys_exact(secret, &["kind", "reference", "environment_variable"]);
assert!(secret.get("value").is_none());
assert!(secret.get("plaintext_secret").is_none());
let mut with_value = json.clone();
with_value["capabilities"][1]
.as_object_mut()
.unwrap()
.insert("value".into(), Value::String("super-secret".into()));
assert!(
serde_json::from_value::<FunctionDefinition>(with_value).is_err(),
"secret capability must reject `value`"
);
let mut with_plaintext = json.clone();
with_plaintext["capabilities"][1]
.as_object_mut()
.unwrap()
.insert(
"plaintext_secret".into(),
Value::String("super-secret".into()),
);
assert!(
serde_json::from_value::<FunctionDefinition>(with_plaintext).is_err(),
"secret capability must reject `plaintext_secret`"
);
Ok(())
}
#[test]
fn debug_redacts_source_and_secret_reference_while_getters_remain_exact() -> Result<()> {
let python = sample_python_definition()?;
let source = sample_source();
assert_eq!(python.source(), source);
let python_debug = format!("{python:?}");
assert!(
!python_debug.contains(source),
"PythonFunctionDefinition Debug must not contain source body: {python_debug}"
);
assert!(
!python_debug.contains("return text[:limit]"),
"PythonFunctionDefinition Debug must not leak source fragments: {python_debug}"
);
let secret = FunctionCapability::try_secret("secret://team/api-token", "API_TOKEN")?;
assert_eq!(secret.reference(), Some("secret://team/api-token"));
assert_eq!(secret.environment_variable(), Some("API_TOKEN"));
let secret_debug = format!("{secret:?}");
assert!(
!secret_debug.contains("secret://team/api-token"),
"FunctionCapability secret Debug must not contain reference: {secret_debug}"
);
let definition = sample_definition()?;
assert_eq!(definition.python_definition().source(), source);
assert_eq!(
definition.capabilities()[1].reference(),
Some("secret://team/api-token")
);
let definition_debug = format!("{definition:?}");
assert!(
!definition_debug.contains(source),
"FunctionDefinition Debug must not contain source body: {definition_debug}"
);
assert!(
!definition_debug.contains("secret://team/api-token"),
"FunctionDefinition Debug must not contain secret reference: {definition_debug}"
);
Ok(())
}
#[test]
fn definition_has_no_identity_before_registration_and_is_not_a_function_handle() -> Result<()> {
let definition = sample_definition()?;
let json = serde_json::to_value(&definition).expect("serialize FunctionDefinition");
assert!(json.get("id").is_none());
assert!(json.get("function_id").is_none());
assert_json_object_keys_exact(
&json,
&[
"format_version",
"signature",
"implementation",
"capabilities",
],
);
// Identity exists only on the immutable Function handle after registration.
// Definition remains a separate authoring value and does not borrow or mint an ID.
let registered = Function::new(
FunctionId::try_new("fn.published.after.registration")?,
definition.signature().clone(),
);
assert_eq!(registered.id().as_str(), "fn.published.after.registration");
assert_eq!(
registered.signature().parameters().len(),
definition.signature().parameters().len()
);
let definition_again = serde_json::to_value(&definition).expect("re-serialize definition");
assert!(definition_again.get("id").is_none());
assert!(definition_again.get("function_id").is_none());
assert_ne!(
serde_json::to_value(&registered).expect("serialize Function"),
definition_again,
"Function handle wire must remain distinct from FunctionDefinition wire"
);
Ok(())
}
@@ -0,0 +1,186 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! Public contract tests for first-class Function error codes (FF-006).
//!
//! These tests pin the stable `FunctionErrorCode` wire strings, direct
//! `Error::Function` projection, and optional `JobFailure.error_code`.
//! They intentionally fail to compile until that public API exists.
//!
//! Categories are judged only by structural enum matching / equality, never
//! by parsing diagnostic message text.
use lancedb::error::FunctionErrorCode;
use lancedb::{Error, JobFailure};
use serde_json::{Value, json};
/// Exact stable wire strings for the eight known Function error categories.
const KNOWN_WIRE_CODES: &[(&str, FunctionErrorCode)] = &[
(
"definition_validation_failure",
FunctionErrorCode::DefinitionValidationFailure,
),
(
"name_or_function_not_found",
FunctionErrorCode::NameOrFunctionNotFound,
),
("name_conflict", FunctionErrorCode::NameConflict),
(
"unsupported_runtime_or_capability",
FunctionErrorCode::UnsupportedRuntimeOrCapability,
),
("revoked_function", FunctionErrorCode::RevokedFunction),
(
"udf_execution_failure",
FunctionErrorCode::UdfExecutionFailure,
),
(
"generated_column_incomplete",
FunctionErrorCode::GeneratedColumnIncomplete,
),
(
"stale_or_conflicting_input",
FunctionErrorCode::StaleOrConflictingInput,
),
];
fn assert_known_variant(code: &FunctionErrorCode, expected: &FunctionErrorCode) {
assert_eq!(
code, expected,
"FunctionErrorCode must match structurally; got {code:?}, expected {expected:?}"
);
assert!(
!matches!(code, FunctionErrorCode::Unrecognized(_)),
"known wire string must not deserialize as Unrecognized: {code:?}"
);
}
#[test]
fn function_error_code_known_variants_use_exact_stable_json_strings() {
for (wire, expected) in KNOWN_WIRE_CODES {
let encoded = serde_json::to_value(expected).expect("serialize FunctionErrorCode");
assert_eq!(
encoded,
Value::String((*wire).to_string()),
"stable JSON string for {expected:?}"
);
let decoded: FunctionErrorCode = serde_json::from_value(Value::String((*wire).to_string()))
.unwrap_or_else(|e| panic!("deserialize `{wire}`: {e}"));
assert_known_variant(&decoded, expected);
let round_trip = serde_json::to_value(&decoded).expect("re-serialize");
assert_eq!(round_trip, Value::String((*wire).to_string()));
}
}
#[test]
fn unrecognized_error_code_preserves_exact_string_and_does_not_become_known() {
let raw = "enterprise_future_category_xyz";
let decoded: FunctionErrorCode = serde_json::from_value(json!(raw))
.unwrap_or_else(|e| panic!("unknown code must deserialize, not fail: {e}"));
match &decoded {
FunctionErrorCode::Unrecognized(preserved) => {
assert_eq!(preserved, raw, "unknown code must be preserved verbatim");
}
other => panic!("expected FunctionErrorCode::Unrecognized, got {other:?}"),
}
for (_, known) in KNOWN_WIRE_CODES {
assert_ne!(
&decoded, known,
"unrecognized code must not equal known variant {known:?}"
);
}
let encoded = serde_json::to_value(&decoded).expect("serialize Unrecognized");
assert_eq!(encoded, json!(raw));
let again: FunctionErrorCode =
serde_json::from_value(encoded).expect("Unrecognized must round-trip");
match again {
FunctionErrorCode::Unrecognized(preserved) => assert_eq!(preserved, raw),
other => panic!("round-trip must stay Unrecognized, got {other:?}"),
}
}
#[test]
fn error_function_carries_code_plus_diagnostic_message() {
let err = Error::Function {
code: FunctionErrorCode::NameConflict,
message: "sanitized diagnostic only".to_string(),
};
match err {
Error::Function { code, message } => {
assert_known_variant(&code, &FunctionErrorCode::NameConflict);
assert_eq!(message, "sanitized diagnostic only");
}
other => panic!("expected Error::Function, got {other:?}"),
}
}
#[test]
fn error_function_category_is_the_code_field_not_the_message() {
// Message text deliberately names a different category; structural code wins.
let err = Error::Function {
code: FunctionErrorCode::GeneratedColumnIncomplete,
message: "looks like udf_execution_failure to a string parser".to_string(),
};
match err {
Error::Function { code, .. } => {
assert_known_variant(&code, &FunctionErrorCode::GeneratedColumnIncomplete);
assert_ne!(code, FunctionErrorCode::UdfExecutionFailure);
}
other => panic!("expected Error::Function, got {other:?}"),
}
}
#[test]
fn job_failure_has_optional_error_code() {
let with_code = JobFailure {
error_code: Some(FunctionErrorCode::RevokedFunction),
phase: Some("execute".to_string()),
message: Some("revoked".to_string()),
retryable: Some(false),
source: None,
};
match &with_code.error_code {
Some(code) => assert_known_variant(code, &FunctionErrorCode::RevokedFunction),
None => panic!("error_code must be present when set"),
}
let without_code = JobFailure {
phase: Some("execute".to_string()),
message: Some("older backend failure without a category".to_string()),
retryable: Some(true),
..Default::default()
};
assert!(
without_code.error_code.is_none(),
"missing error_code must stay None; diagnostics must not invent a category"
);
}
#[test]
fn job_failure_diagnostics_do_not_overwrite_error_code() {
let failure = JobFailure {
error_code: Some(FunctionErrorCode::StaleOrConflictingInput),
phase: Some("commit".to_string()),
message: Some("definition_validation_failure in worker logs".to_string()),
retryable: Some(true),
source: None,
};
match &failure.error_code {
Some(code) => {
assert_known_variant(code, &FunctionErrorCode::StaleOrConflictingInput);
assert_ne!(code, &FunctionErrorCode::DefinitionValidationFailure);
}
None => panic!("explicit error_code must remain set"),
}
assert_eq!(failure.phase.as_deref(), Some("commit"));
assert_eq!(failure.retryable, Some(true));
}
@@ -0,0 +1,366 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! Contract tests for JobResult value/wire (FF-012).
//!
//! These tests pin the intended public non-resource [`lancedb::JobResult`]
//! surface under [`lancedb::job`]. They intentionally fail to compile until
//! that API exists.
//!
//! Scope is JobResult value and JSON wire only. Job::wait behavior, remote
//! describe shape, missing-result handling, local outcome, Python, Node, and
//! Sophon are out of scope.
//!
//! Rejection cases are judged by `Result` structure (`is_err` / `is_ok`) or
//! serde decode failure, never by diagnostic message substrings.
//! JSON map iteration order is not a contract; exact key sets are compared
//! independently. Byte reproducibility means repeated encoding of the same
//! in-memory value.
use std::collections::BTreeSet;
use arrow_schema::DataType;
use lancedb::JobResult;
use lancedb::Result;
use lancedb::function::{
Function, FunctionId, FunctionOutput, FunctionParameter, FunctionSignature,
};
use serde_json::Value;
fn sample_function() -> Result<Function> {
let id = FunctionId::try_new("fn.exact.job-result")?;
let signature = FunctionSignature::try_new(
vec![
FunctionParameter::new("x", DataType::Int32),
FunctionParameter::new("label", DataType::Utf8),
],
FunctionOutput::new(DataType::Int32, true),
)?;
Ok(Function::new(id, signature))
}
fn assert_json_object_keys_exact(value: &Value, expected: &[&str]) {
let object = value
.as_object()
.unwrap_or_else(|| panic!("expected JSON object, got {value}"));
let keys: BTreeSet<&str> = object.keys().map(|k| k.as_str()).collect();
let expected: BTreeSet<&str> = expected.iter().copied().collect();
assert_eq!(
keys, expected,
"JSON object key set must match exactly (iteration order is not a contract); got {keys:?}, expected {expected:?} in {value}"
);
}
/// Outer JobResult object keys that must not appear. Nested Function signature
/// keys (including parameter `name` and Function `id`) are legitimate and are
/// not scanned here. Opaque string contents are not recursively searched.
const FORBIDDEN_OUTER_RESULT_KEYS: &[&str] = &[
"name",
"definition",
"FunctionDefinition",
"source",
"runtime",
"packages",
"capability",
"capabilities",
"artifact",
"artifact_digest",
"digest",
"storage",
"storage_location",
"location",
"table",
"table_name",
"table_ref",
"version",
"function_version",
"FunctionVersion",
"user_version",
"lineage",
"id",
"job_id",
"jobId",
"type",
"job_type",
"state",
"status",
"lifecycle",
"failure",
"attempt",
"attempt_id",
"timestamp",
"created_at",
"updated_at",
"retry",
"retry_key",
"idempotency",
"idempotency_key",
"commit_token",
"secret",
"compatibility",
"deterministic",
"null_policy",
"nullPolicy",
];
fn assert_outer_forbidden_keys_absent(value: &Value) {
let object = value
.as_object()
.unwrap_or_else(|| panic!("expected outer JobResult JSON object, got {value}"));
for key in object.keys() {
assert!(
!FORBIDDEN_OUTER_RESULT_KEYS.contains(&key.as_str()),
"outer JobResult wire must not contain forbidden key `{key}`: {value}"
);
}
}
fn assert_function_handle_exact(actual: &Function, expected: &Function) {
assert_eq!(actual.id().as_str(), expected.id().as_str());
assert_eq!(
actual.signature().parameters().len(),
expected.signature().parameters().len()
);
for (actual_param, expected_param) in actual
.signature()
.parameters()
.iter()
.zip(expected.signature().parameters().iter())
{
assert_eq!(actual_param.name(), expected_param.name());
assert_eq!(actual_param.data_type(), expected_param.data_type());
}
assert_eq!(
actual.signature().output().data_type(),
expected.signature().output().data_type()
);
assert_eq!(
actual.signature().output().nullable(),
expected.signature().output().nullable()
);
}
#[test]
fn none_and_function_exact_key_sets_helpers_round_trip_and_bytes() -> Result<()> {
let none = JobResult::None;
assert_eq!(none.format_version(), 1);
assert!(none.function().is_none());
assert!(matches!(none, JobResult::None));
let none_json = serde_json::to_value(&none).expect("serialize JobResult::None");
assert_json_object_keys_exact(&none_json, &["format_version", "kind"]);
assert_eq!(none_json["format_version"], 1);
assert_eq!(none_json["kind"], Value::String("none".into()));
assert!(none_json.get("function").is_none());
let none_restored: JobResult =
serde_json::from_value(none_json.clone()).expect("deserialize JobResult::None");
assert_eq!(none_restored.format_version(), 1);
assert!(none_restored.function().is_none());
assert!(matches!(none_restored, JobResult::None));
assert_eq!(none_restored, none);
assert_eq!(
serde_json::to_value(&none_restored).expect("re-serialize None"),
none_json
);
let none_a = serde_json::to_string(&none).expect("encode None a");
let none_b = serde_json::to_string(&none).expect("encode None b");
assert_eq!(
none_a, none_b,
"repeated None encoding must be byte-identical"
);
let function = sample_function()?;
let expected_function_wire =
serde_json::to_value(&function).expect("serialize nested Function");
let function_result = JobResult::Function(function.clone());
assert_eq!(function_result.format_version(), 1);
assert!(matches!(function_result, JobResult::Function(_)));
assert_function_handle_exact(
function_result.function().expect("Function variant"),
&function,
);
let function_json =
serde_json::to_value(&function_result).expect("serialize JobResult::Function");
assert_json_object_keys_exact(&function_json, &["format_version", "kind", "function"]);
assert_eq!(function_json["format_version"], 1);
assert_eq!(function_json["kind"], Value::String("function".into()));
assert_eq!(
function_json["function"], expected_function_wire,
"nested function must be the exact existing Function wire"
);
assert_json_object_keys_exact(
&function_json["function"],
&["format_version", "id", "signature"],
);
assert_eq!(
function_json["function"]["id"],
Value::String("fn.exact.job-result".into())
);
assert_eq!(function_json["function"]["format_version"], 1);
let function_restored: JobResult =
serde_json::from_value(function_json.clone()).expect("deserialize JobResult::Function");
assert_eq!(function_restored.format_version(), 1);
assert_function_handle_exact(
function_restored.function().expect("Function variant"),
&function,
);
assert_eq!(
function_restored
.function()
.expect("Function variant")
.id()
.as_str(),
"fn.exact.job-result"
);
assert_eq!(function_restored, function_result);
assert_eq!(
serde_json::to_value(&function_restored).expect("re-serialize Function"),
function_json
);
let function_a = serde_json::to_string(&function_result).expect("encode Function a");
let function_b = serde_json::to_string(&function_result).expect("encode Function b");
assert_eq!(
function_a, function_b,
"repeated Function encoding must be byte-identical"
);
Ok(())
}
#[test]
fn function_and_into_function_accessors_for_both_variants() -> Result<()> {
let none = JobResult::None;
assert!(none.function().is_none());
assert!(none.into_function().is_none());
let function = sample_function()?;
let function_result = JobResult::Function(function.clone());
assert_function_handle_exact(
function_result.function().expect("borrowed Function"),
&function,
);
let owned = function_result
.into_function()
.expect("owned Function from Function variant");
assert_function_handle_exact(&owned, &function);
assert_eq!(owned.id().as_str(), "fn.exact.job-result");
Ok(())
}
#[test]
fn unknown_kind_field_version_and_malformed_function_fail_closed() -> Result<()> {
let none = JobResult::None;
let none_json = serde_json::to_value(&none).expect("serialize None");
let mut unknown_kind = none_json.clone();
unknown_kind["kind"] = Value::String("artifact".into());
assert!(
serde_json::from_value::<JobResult>(unknown_kind).is_err(),
"unknown kind must fail closed and must not become None"
);
let mut unknown_field = none_json.clone();
unknown_field
.as_object_mut()
.unwrap()
.insert("unexpected_field".into(), Value::Bool(true));
assert!(
serde_json::from_value::<JobResult>(unknown_field).is_err(),
"unknown outer field must fail closed"
);
let mut unknown_version = none_json.clone();
unknown_version["format_version"] = Value::from(2);
assert!(
serde_json::from_value::<JobResult>(unknown_version).is_err(),
"unsupported format_version must fail closed"
);
let mut unexpected_function_on_none = none_json.clone();
unexpected_function_on_none.as_object_mut().unwrap().insert(
"function".into(),
serde_json::to_value(&sample_function()?).expect("nested Function"),
);
assert!(
serde_json::from_value::<JobResult>(unexpected_function_on_none).is_err(),
"kind=none with unexpected function field must fail closed"
);
let function = sample_function()?;
let function_json =
serde_json::to_value(JobResult::Function(function)).expect("serialize Function");
let mut missing_function = function_json.clone();
missing_function.as_object_mut().unwrap().remove("function");
assert!(
serde_json::from_value::<JobResult>(missing_function).is_err(),
"kind=function without function field must fail closed"
);
let mut empty_function_id = function_json.clone();
empty_function_id["function"]["id"] = Value::String("".into());
assert!(
serde_json::from_value::<JobResult>(empty_function_id).is_err(),
"empty nested Function ID must fail closed"
);
let mut unknown_nested_function_field = function_json.clone();
unknown_nested_function_field["function"]
.as_object_mut()
.unwrap()
.insert("unexpected_field".into(), Value::Bool(true));
assert!(
serde_json::from_value::<JobResult>(unknown_nested_function_field).is_err(),
"unknown nested Function field must fail closed"
);
let mut malformed_nested_version = function_json.clone();
malformed_nested_version["function"]["format_version"] = Value::from(2);
assert!(
serde_json::from_value::<JobResult>(malformed_nested_version).is_err(),
"malformed nested Function must fail closed"
);
// Unknown wire must not be preserved as a public variant or downgraded to None.
let unknown_raw = serde_json::json!({
"format_version": 1,
"kind": "future_result_kind",
"raw": {"keep": true}
});
assert!(
serde_json::from_value::<JobResult>(unknown_raw).is_err(),
"unknown kind must fail closed without a public raw/unknown variant"
);
Ok(())
}
#[test]
fn outer_result_excludes_forbidden_fields() -> Result<()> {
let none_json = serde_json::to_value(&JobResult::None).expect("serialize None");
assert_json_object_keys_exact(&none_json, &["format_version", "kind"]);
assert_outer_forbidden_keys_absent(&none_json);
let function = sample_function()?;
let function_json =
serde_json::to_value(JobResult::Function(function)).expect("serialize Function");
assert_json_object_keys_exact(&function_json, &["format_version", "kind", "function"]);
assert_outer_forbidden_keys_absent(&function_json);
// Nested Function may carry `id` and signature parameter `name`; those are
// not outer result keys and must remain present on the nested object.
assert_eq!(
function_json["function"]["id"],
Value::String("fn.exact.job-result".into())
);
assert!(
function_json["function"]["signature"]["parameters"][0]
.get("name")
.is_some(),
"nested signature parameter `name` remains legitimate"
);
Ok(())
}
@@ -0,0 +1,537 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! Contract tests for RegisterFunctionJobSpec (FF-008 / B1d).
//!
//! These tests pin the intended public surface under [`lancedb::function`] for
//! registration Job operation input only. They intentionally fail to compile
//! until that API exists.
//!
//! Rejection cases are judged by `Result` structure (`is_err` / `is_ok`), never
//! by diagnostic message substrings. Catalog execution, typed Job wait, and
//! result Function publication are out of scope.
use std::collections::BTreeSet;
use arrow_schema::DataType;
use lancedb::Result;
use lancedb::function::{
FunctionCapability, FunctionDefinition, FunctionId, FunctionOutput, FunctionParameter,
FunctionSignature, PythonFunctionDefinition, RegisterFunctionJobSpec,
};
use serde_json::Value;
fn sample_signature() -> Result<FunctionSignature> {
FunctionSignature::try_new(
vec![
FunctionParameter::new("text", DataType::Utf8),
FunctionParameter::new("limit", DataType::Int32),
],
FunctionOutput::new(DataType::Utf8, true),
)
}
fn sample_source() -> &'static str {
"def normalize(text, limit):\n return text[:limit]\n"
}
fn sample_python_definition() -> Result<PythonFunctionDefinition> {
PythonFunctionDefinition::try_new(
"normalize_mod",
"normalize",
sample_source(),
"3.12",
vec!["Unidecode==1.3.8".to_string()],
)
}
fn sample_capabilities() -> Result<Vec<FunctionCapability>> {
Ok(vec![
FunctionCapability::try_network("https://api.example.com")?,
FunctionCapability::try_secret("secret://team/api-token", "API_TOKEN")?,
])
}
fn sample_definition() -> Result<FunctionDefinition> {
FunctionDefinition::try_new(
sample_signature()?,
sample_python_definition()?,
sample_capabilities()?,
)
}
fn sample_create_spec() -> Result<RegisterFunctionJobSpec> {
RegisterFunctionJobSpec::try_new("text.normalize", sample_definition()?, None)
}
fn sample_replace_spec() -> Result<RegisterFunctionJobSpec> {
RegisterFunctionJobSpec::try_new(
"text.normalize",
sample_definition()?,
Some(FunctionId::try_new("fn.existing.exact")?),
)
}
fn assert_json_object_keys_exact(value: &Value, expected: &[&str]) {
let object = value
.as_object()
.unwrap_or_else(|| panic!("expected JSON object, got {value}"));
let keys: BTreeSet<&str> = object.keys().map(|k| k.as_str()).collect();
let expected: BTreeSet<&str> = expected.iter().copied().collect();
assert_eq!(
keys, expected,
"JSON object key set must match exactly (iteration order is not a contract); got {keys:?}, expected {expected:?} in {value}"
);
}
/// Lifecycle / identity / artifact / runtime keys that must not appear as public
/// object keys on RegisterFunctionJobSpec wire. Exact key matching only: do not
/// substring-scan encoded JSON (source text, parameter `name`, and
/// `expected_current_function_id` must not false-match).
const FORBIDDEN_SPEC_KEYS: &[&str] = &[
"id",
"function_id",
"FunctionId",
"new_function_id",
"generated_function_id",
"result_function_id",
"version",
"function_version",
"FunctionVersion",
"user_version",
"lineage",
"idempotency_key",
"retry_key",
"idempotency",
"job_id",
"jobId",
"state",
"status",
"attempt",
"attempt_id",
"timestamp",
"created_at",
"updated_at",
"table",
"table_name",
"table_ref",
"executor",
"environment",
"artifact",
"artifact_digest",
"digest",
"storage",
"storage_location",
"location",
"worker",
"scheduler",
"replica",
"placement",
];
fn assert_object_keys_not_forbidden(value: &Value, context: &str) {
let object = value.as_object().unwrap_or_else(|| {
panic!("expected JSON object at {context}, got {value}");
});
for key in object.keys() {
assert!(
!FORBIDDEN_SPEC_KEYS.contains(&key.as_str()),
"RegisterFunctionJobSpec wire must not contain forbidden key `{key}` at {context}: {value}"
);
}
}
fn assert_forbidden_spec_keys_absent(value: &Value) {
assert_json_object_keys_exact(
value,
&[
"format_version",
"name",
"definition",
"expected_current_function_id",
],
);
assert_object_keys_not_forbidden(value, "RegisterFunctionJobSpec");
// Precondition field is allowed; a generated/new Function ID key is not.
assert!(
value.get("function_id").is_none(),
"spec must not carry generated/new `function_id`; use expected_current_function_id only: {value}"
);
assert!(
value.get("id").is_none(),
"spec must not carry generated/new Function `id`: {value}"
);
let definition = value.get("definition").expect("definition object");
assert_json_object_keys_exact(
definition,
&[
"format_version",
"signature",
"implementation",
"capabilities",
],
);
assert_object_keys_not_forbidden(definition, "definition");
// Catalog/function identity name belongs on the spec, not the nested definition.
assert!(
definition.get("name").is_none(),
"nested definition must not contain catalog name: {definition}"
);
let implementation = definition
.get("implementation")
.expect("implementation object");
assert_json_object_keys_exact(
implementation,
&["kind", "module", "callable", "source", "python", "packages"],
);
assert_object_keys_not_forbidden(implementation, "definition.implementation");
let capabilities = definition
.get("capabilities")
.and_then(Value::as_array)
.expect("capabilities array");
for (idx, capability) in capabilities.iter().enumerate() {
let context = format!("definition.capabilities[{idx}]");
assert_object_keys_not_forbidden(capability, &context);
}
let signature = definition.get("signature").expect("signature object");
assert_object_keys_not_forbidden(signature, "definition.signature");
if let Some(parameters) = signature.get("parameters").and_then(Value::as_array) {
for (idx, parameter) in parameters.iter().enumerate() {
let context = format!("definition.signature.parameters[{idx}]");
assert_object_keys_not_forbidden(parameter, &context);
// Parameter object key `name` is legitimate and must not be treated
// as a forbidden catalog/function identity field.
assert!(
parameter.get("name").is_some(),
"{context} must include parameter `name`"
);
}
}
}
#[test]
fn create_and_replace_round_trip_pins_name_definition_and_precondition() -> Result<()> {
let create = sample_create_spec()?;
assert_eq!(create.format_version(), 1);
assert_eq!(create.name(), "text.normalize");
assert!(create.expected_current_function_id().is_none());
assert_eq!(
create.definition().python_definition().source(),
sample_source()
);
assert_eq!(create.definition().capabilities().len(), 2);
assert_eq!(
create.definition().capabilities()[0].origin(),
Some("https://api.example.com")
);
assert_eq!(
create.definition().capabilities()[1].reference(),
Some("secret://team/api-token")
);
let create_json = serde_json::to_value(&create).expect("serialize create spec");
assert_eq!(create_json["format_version"], 1);
assert_eq!(create_json["name"], Value::String("text.normalize".into()));
assert_eq!(create_json["expected_current_function_id"], Value::Null);
let create_restored: RegisterFunctionJobSpec =
serde_json::from_value(create_json.clone()).expect("deserialize create spec");
assert_eq!(create_restored.format_version(), 1);
assert_eq!(create_restored.name(), "text.normalize");
assert!(create_restored.expected_current_function_id().is_none());
assert_eq!(
create_restored.definition().python_definition().module(),
create.definition().python_definition().module()
);
assert_eq!(
create_restored.definition().python_definition().callable(),
create.definition().python_definition().callable()
);
assert_eq!(
create_restored.definition().python_definition().source(),
create.definition().python_definition().source()
);
assert_eq!(
create_restored.definition().python_definition().python(),
create.definition().python_definition().python()
);
assert_eq!(
create_restored.definition().python_definition().packages(),
create.definition().python_definition().packages()
);
assert_eq!(
create_restored.definition().capabilities()[1].reference(),
create.definition().capabilities()[1].reference()
);
assert_eq!(
create_restored.definition().capabilities()[1].environment_variable(),
create.definition().capabilities()[1].environment_variable()
);
let replace = sample_replace_spec()?;
assert_eq!(replace.name(), "text.normalize");
assert_eq!(
replace
.expected_current_function_id()
.map(FunctionId::as_str),
Some("fn.existing.exact")
);
let replace_json = serde_json::to_value(&replace).expect("serialize replace spec");
assert_eq!(
replace_json["expected_current_function_id"],
Value::String("fn.existing.exact".into())
);
let replace_restored: RegisterFunctionJobSpec =
serde_json::from_value(replace_json.clone()).expect("deserialize replace spec");
assert_eq!(
replace_restored
.expected_current_function_id()
.map(FunctionId::as_str),
Some("fn.existing.exact")
);
assert_eq!(replace_restored.name(), replace.name());
assert_eq!(
replace_restored.definition().python_definition().source(),
replace.definition().python_definition().source()
);
// Byte-for-byte repeated serde_json encoding for the same value.
let create_a = serde_json::to_string(&create).expect("encode create a");
let create_b = serde_json::to_string(&create).expect("encode create b");
assert_eq!(create_a, create_b);
let replace_a = serde_json::to_string(&replace).expect("encode replace a");
let replace_b = serde_json::to_string(&replace).expect("encode replace b");
assert_eq!(replace_a, replace_b);
assert_eq!(
serde_json::to_value(&create_restored).expect("re-serialize create"),
create_json
);
assert_eq!(
serde_json::to_value(&replace_restored).expect("re-serialize replace"),
replace_json
);
Ok(())
}
#[test]
fn create_wire_pins_exact_key_set_and_null_precondition() -> Result<()> {
let create = sample_create_spec()?;
let json = serde_json::to_value(&create).expect("serialize create spec");
// Exact object key set; Map iteration order is not part of the contract.
assert_json_object_keys_exact(
&json,
&[
"format_version",
"name",
"definition",
"expected_current_function_id",
],
);
assert_eq!(json["format_version"], 1);
assert_eq!(json["name"], Value::String("text.normalize".into()));
// Create-if-absent always serializes the precondition key as JSON null.
assert_eq!(json["expected_current_function_id"], Value::Null);
assert!(json["expected_current_function_id"].is_null());
let replace = sample_replace_spec()?;
let replace_json = serde_json::to_value(&replace).expect("serialize replace spec");
assert_json_object_keys_exact(
&replace_json,
&[
"format_version",
"name",
"definition",
"expected_current_function_id",
],
);
assert_eq!(
replace_json["expected_current_function_id"],
Value::String("fn.existing.exact".into())
);
Ok(())
}
#[test]
fn empty_name_and_expected_id_unknown_field_and_version_fail_closed() -> Result<()> {
let definition = sample_definition()?;
assert!(
RegisterFunctionJobSpec::try_new("", definition.clone(), None).is_err(),
"empty name must be rejected by constructor"
);
assert!(
RegisterFunctionJobSpec::try_new(
"text.normalize",
definition,
Some(FunctionId::try_new("fn.existing.exact")?),
)
.is_ok(),
"non-empty name with exact expected ID must construct"
);
let create = sample_create_spec()?;
let json = serde_json::to_value(&create).expect("serialize create spec");
let mut empty_name = json.clone();
empty_name["name"] = Value::String("".into());
assert!(
serde_json::from_value::<RegisterFunctionJobSpec>(empty_name).is_err(),
"decode must reject empty name"
);
let mut empty_expected_id = json.clone();
empty_expected_id["expected_current_function_id"] = Value::String("".into());
assert!(
serde_json::from_value::<RegisterFunctionJobSpec>(empty_expected_id).is_err(),
"decode must reject empty expected_current_function_id"
);
let mut unknown_field = json.clone();
unknown_field
.as_object_mut()
.unwrap()
.insert("unexpected_field".into(), Value::Bool(true));
assert!(
serde_json::from_value::<RegisterFunctionJobSpec>(unknown_field).is_err(),
"unknown outer field must fail closed"
);
let mut unknown_version = json.clone();
unknown_version["format_version"] = Value::from(2);
assert!(
serde_json::from_value::<RegisterFunctionJobSpec>(unknown_version).is_err(),
"format_version other than 1 must fail closed"
);
Ok(())
}
#[test]
fn spec_wire_excludes_generated_identity_job_lifecycle_and_artifact_fields() -> Result<()> {
let create = sample_create_spec()?;
let create_json = serde_json::to_value(&create).expect("serialize create spec");
assert_forbidden_spec_keys_absent(&create_json);
let replace = sample_replace_spec()?;
let replace_json = serde_json::to_value(&replace).expect("serialize replace spec");
assert_forbidden_spec_keys_absent(&replace_json);
// Replace may carry the exact opaque precondition string only.
assert_eq!(
replace_json["expected_current_function_id"],
Value::String("fn.existing.exact".into())
);
Ok(())
}
#[test]
fn debug_redacts_source_and_secret_reference_while_nested_getters_remain_exact() -> Result<()> {
let source = sample_source();
let secret_reference = "secret://team/api-token";
let create = sample_create_spec()?;
assert_eq!(create.definition().python_definition().source(), source);
assert_eq!(
create.definition().capabilities()[1].reference(),
Some(secret_reference)
);
let create_debug = format!("{create:?}");
assert!(
!create_debug.contains(source),
"create RegisterFunctionJobSpec Debug must not contain source body: {create_debug}"
);
assert!(
!create_debug.contains("return text[:limit]"),
"create RegisterFunctionJobSpec Debug must not leak source fragments: {create_debug}"
);
assert!(
!create_debug.contains(secret_reference),
"create RegisterFunctionJobSpec Debug must not contain secret reference: {create_debug}"
);
let replace = sample_replace_spec()?;
assert_eq!(replace.definition().python_definition().source(), source);
assert_eq!(
replace.definition().capabilities()[1].reference(),
Some(secret_reference)
);
assert_eq!(
replace
.expected_current_function_id()
.map(FunctionId::as_str),
Some("fn.existing.exact")
);
let replace_debug = format!("{replace:?}");
assert!(
!replace_debug.contains(source),
"replace RegisterFunctionJobSpec Debug must not contain source body: {replace_debug}"
);
assert!(
!replace_debug.contains(secret_reference),
"replace RegisterFunctionJobSpec Debug must not contain secret reference: {replace_debug}"
);
Ok(())
}
#[test]
fn nested_definition_wire_is_exact_ff007_function_definition() -> Result<()> {
let definition = sample_definition()?;
let expected_definition_wire =
serde_json::to_value(&definition).expect("serialize FunctionDefinition");
let create = RegisterFunctionJobSpec::try_new("text.normalize", definition.clone(), None)?;
let create_json = serde_json::to_value(&create).expect("serialize create spec");
assert_eq!(
create_json["definition"], expected_definition_wire,
"nested definition must be the exact FF-007 FunctionDefinition wire"
);
assert_json_object_keys_exact(
&create_json["definition"],
&[
"format_version",
"signature",
"implementation",
"capabilities",
],
);
assert_eq!(
create_json["definition"]["implementation"]["source"],
Value::String(sample_source().into())
);
assert_eq!(
create_json["definition"]["capabilities"][1]["reference"],
Value::String("secret://team/api-token".into())
);
// Not a summarized / digested / artifact reference indirection.
assert!(create_json["definition"].get("digest").is_none());
assert!(create_json["definition"].get("artifact").is_none());
assert!(create_json["definition"].get("artifact_digest").is_none());
assert!(create_json["definition"].get("storage").is_none());
assert!(
create_json["definition"]["implementation"]
.get("digest")
.is_none()
);
assert!(
create_json["definition"]["implementation"]
.get("artifact")
.is_none()
);
let replace = RegisterFunctionJobSpec::try_new(
"text.normalize",
definition,
Some(FunctionId::try_new("fn.existing.exact")?),
)?;
let replace_json = serde_json::to_value(&replace).expect("serialize replace spec");
assert_eq!(
replace_json["definition"], expected_definition_wire,
"replace nested definition must remain the exact FF-007 wire"
);
Ok(())
}
@@ -0,0 +1,495 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
//! Contract tests for GeneratedColumnBindingSnapshot (FF-029 / FF-030).
//!
//! Pins the hidden value projection, Table seam, and bound-call field
//! validation used by generated-column call binding. These tests intentionally
//! fail to compile until that API exists. They do not submit Jobs, mutate
//! generated-column state, or resolve authored Function calls.
use std::sync::Arc;
use arrow_array::{ArrayRef, Int32Array, RecordBatch, StringArray};
use arrow_schema::{DataType, Field, Schema};
use lance::dataset::NewColumnTransform;
use lancedb::connect;
use lancedb::function::{
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
FunctionSignature, GeneratedColumnBindingEntry, GeneratedColumnBindingSnapshot,
};
use lancedb::table::ColumnAlteration;
use lancedb::{Error, Result};
use tempfile::tempdir;
fn sample_fields() -> Vec<arrow_schema::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)),
]
}
fn sample_output() -> FunctionOutput {
FunctionOutput::new(DataType::Int32, true)
}
/// Parameter names intentionally differ from table column names so any
/// name-based validation would fail these fixtures.
fn two_field_function() -> Result<Function> {
let id = FunctionId::try_new("fn.exact.binding.validate")?;
let signature = FunctionSignature::try_new(
vec![
FunctionParameter::new("input_payload", DataType::Utf8),
FunctionParameter::new("metric_value", DataType::Int32),
],
sample_output(),
)?;
Ok(Function::new(id, signature))
}
fn one_field_function() -> Result<Function> {
let id = FunctionId::try_new("fn.exact.binding.one-field")?;
let signature = FunctionSignature::try_new(
vec![FunctionParameter::new("payload_arg", DataType::Utf8)],
sample_output(),
)?;
Ok(Function::new(id, signature))
}
fn literal_only_function() -> Result<Function> {
let id = FunctionId::try_new("fn.exact.binding.literal-only")?;
let signature = FunctionSignature::try_new(
vec![FunctionParameter::new("constant_arg", DataType::Int32)],
sample_output(),
)?;
Ok(Function::new(id, signature))
}
fn mixed_function() -> Result<Function> {
let id = FunctionId::try_new("fn.exact.binding.mixed")?;
let signature = FunctionSignature::try_new(
vec![
FunctionParameter::new("payload_arg", DataType::Utf8),
FunctionParameter::new("constant_arg", DataType::Int32),
],
sample_output(),
)?;
Ok(Function::new(id, signature))
}
fn int_literal(value: Option<i32>) -> Result<FunctionArgument> {
FunctionArgument::try_literal(Arc::new(Int32Array::from(vec![value])) as ArrayRef)
}
#[test]
fn try_new_preserves_version_order_and_exact_lookup() -> Result<()> {
let fields = sample_fields();
let snapshot = GeneratedColumnBindingSnapshot::try_new(7, fields.clone(), vec![3, 5, 9])?;
assert_eq!(snapshot.version(), 7);
let entries = snapshot.entries();
assert_eq!(entries.len(), 3);
assert_eq!(entries[0].field_id(), 3);
assert_eq!(entries[0].field().name(), "text");
assert_eq!(entries[0].field().data_type(), &DataType::Utf8);
assert_eq!(entries[1].field_id(), 5);
assert_eq!(entries[1].field().name(), "score");
assert_eq!(entries[2].field_id(), 9);
assert_eq!(entries[2].field().name(), "a.b");
let by_name = snapshot.field("score").expect("exact name");
assert_eq!(by_name.field_id(), 5);
assert!(snapshot.field("Score").is_none());
assert!(snapshot.field("a").is_none());
let dotted = snapshot
.field("a.b")
.expect("literal dotted top-level name");
assert_eq!(dotted.field_id(), 9);
assert_eq!(dotted.field().as_ref(), fields[2].as_ref());
// Type existence pin for the entry surface used by the next binding slice.
let _: &GeneratedColumnBindingEntry = by_name;
Ok(())
}
#[test]
fn try_new_rejects_invalid_projections() {
let fields = sample_fields();
assert!(matches!(
GeneratedColumnBindingSnapshot::try_new(1, fields.clone(), vec![1, 2]),
Err(Error::InvalidInput { .. })
));
assert!(matches!(
GeneratedColumnBindingSnapshot::try_new(1, fields.clone(), vec![1, 2, -1]),
Err(Error::InvalidInput { .. })
));
assert!(matches!(
GeneratedColumnBindingSnapshot::try_new(1, fields.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 { .. })
));
}
#[test]
fn validate_field_arguments_accepts_valid_and_mixed_bindings() -> Result<()> {
let snapshot = GeneratedColumnBindingSnapshot::try_new(3, sample_fields(), vec![3, 5, 9])?;
let one = one_field_function()?;
let valid = FunctionCall::try_new(
&one,
vec![(
"payload_arg".to_string(),
FunctionArgument::try_field(3, DataType::Utf8)?,
)],
)?;
snapshot.validate_field_arguments(&valid)?;
let two = two_field_function()?;
let multi = FunctionCall::try_new(
&two,
vec![
(
"input_payload".to_string(),
FunctionArgument::try_field(3, DataType::Utf8)?,
),
(
"metric_value".to_string(),
FunctionArgument::try_field(5, DataType::Int32)?,
),
],
)?;
snapshot.validate_field_arguments(&multi)?;
let mixed_fn = mixed_function()?;
let mixed = FunctionCall::try_new(
&mixed_fn,
vec![
(
"payload_arg".to_string(),
FunctionArgument::try_field(3, DataType::Utf8)?,
),
("constant_arg".to_string(), int_literal(Some(42))?),
],
)?;
snapshot.validate_field_arguments(&mixed)?;
Ok(())
}
#[test]
fn validate_field_arguments_rejects_missing_id_and_type_mismatch() -> Result<()> {
let snapshot = GeneratedColumnBindingSnapshot::try_new(3, sample_fields(), vec![3, 5, 9])?;
let one = one_field_function()?;
let missing = FunctionCall::try_new(
&one,
vec![(
"payload_arg".to_string(),
FunctionArgument::try_field(99, DataType::Utf8)?,
)],
)?;
let err = snapshot
.validate_field_arguments(&missing)
.expect_err("missing stable field id");
assert!(matches!(err, Error::InvalidInput { .. }));
let message = err.to_string();
assert!(
message.contains("99"),
"diagnostics may name field id: {message}"
);
assert!(
!message.contains("text") && !message.contains("score") && !message.contains("a.b"),
"diagnostics must not invent or use a column name: {message}"
);
// Same stable ID, different Arrow type: exact-type equality must reject.
// This covers Remote/other producer projections that keep the ID.
let type_mismatch = FunctionCall::try_new(
&one,
vec![(
"payload_arg".to_string(),
FunctionArgument::try_field(5, DataType::Utf8)?,
)],
)?;
let err = snapshot
.validate_field_arguments(&type_mismatch)
.expect_err("same-id exact type mismatch");
assert!(matches!(err, Error::InvalidInput { .. }));
let message = err.to_string();
assert!(
message.contains("5"),
"diagnostics may name field id: {message}"
);
assert!(
message.contains("Utf8") && message.contains("Int32"),
"diagnostics may identify expected/current types: {message}"
);
assert!(
!message.contains("score") && !message.contains("text"),
"diagnostics must not invent or use a column name: {message}"
);
Ok(())
}
#[test]
fn validate_field_arguments_literal_only_ignores_table_fields() -> Result<()> {
// Snapshot has no field that a name-based binder could match to "constant_arg".
let snapshot = GeneratedColumnBindingSnapshot::try_new(
1,
vec![Arc::new(Field::new("unrelated", DataType::Utf8, true))],
vec![11],
)?;
let function = literal_only_function()?;
let call = FunctionCall::try_new(
&function,
vec![("constant_arg".to_string(), int_literal(Some(7))?)],
)?;
snapshot.validate_field_arguments(&call)?;
Ok(())
}
#[tokio::test]
async fn table_seam_returns_atomic_native_snapshot() -> Result<()> {
let tmp = tempdir().unwrap();
let db = connect(tmp.path().to_str().unwrap()).execute().await?;
let schema = Arc::new(Schema::new(vec![
Field::new("text", DataType::Utf8, true),
Field::new("score", DataType::Int32, false),
]));
let batch = RecordBatch::try_new(
schema,
vec![
Arc::new(StringArray::from(vec![Some("a")])),
Arc::new(Int32Array::from(vec![1])),
],
)?;
let table = db.create_table("binding", batch).execute().await?;
let snapshot = table.generated_column_binding_snapshot().await?;
let public_schema = table.schema().await?;
let version = table.version().await?;
assert_eq!(snapshot.version(), version);
assert_eq!(snapshot.entries().len(), public_schema.fields().len());
for (entry, field) in snapshot.entries().iter().zip(public_schema.fields()) {
assert_eq!(entry.field().name(), field.name());
assert_eq!(entry.field().data_type(), field.data_type());
assert!(entry.field_id() >= 0);
assert!(!field.metadata().contains_key("lance:field_id"));
assert!(!entry.field().metadata().contains_key("lance:field_id"));
}
Ok(())
}
#[tokio::test]
async fn validate_field_arguments_survives_rename_on_real_table() -> Result<()> {
let tmp = tempdir().unwrap();
let db = connect(tmp.path().to_str().unwrap()).execute().await?;
// Column names deliberately differ from Function parameter names.
let schema = Arc::new(Schema::new(vec![
Field::new("source_text", DataType::Utf8, true),
Field::new("source_score", DataType::Int32, false),
]));
let batch = RecordBatch::try_new(
schema,
vec![
Arc::new(StringArray::from(vec![Some("hello")])),
Arc::new(Int32Array::from(vec![7])),
],
)?;
let table = db.create_table("binding_rename", batch).execute().await?;
let before = table.generated_column_binding_snapshot().await?;
let text_entry = before.field("source_text").expect("source_text");
let score_entry = before.field("source_score").expect("source_score");
let text_id = text_entry.field_id();
let score_id = score_entry.field_id();
assert_eq!(text_entry.field().data_type(), &DataType::Utf8);
assert_eq!(score_entry.field().data_type(), &DataType::Int32);
let function = two_field_function()?;
let call = FunctionCall::try_new(
&function,
vec![
(
"input_payload".to_string(),
FunctionArgument::try_field(text_id, DataType::Utf8)?,
),
(
"metric_value".to_string(),
FunctionArgument::try_field(score_id, DataType::Int32)?,
),
],
)?;
before.validate_field_arguments(&call)?;
table
.alter_columns(&[ColumnAlteration::new("source_text".into()).rename("renamed_text".into())])
.await?;
let after = table.generated_column_binding_snapshot().await?;
assert!(after.field("source_text").is_none());
let renamed = after.field("renamed_text").expect("renamed_text");
assert_eq!(renamed.field_id(), text_id);
assert_eq!(renamed.field().data_type(), &DataType::Utf8);
assert_eq!(
after
.field("source_score")
.expect("source_score")
.field_id(),
score_id
);
after.validate_field_arguments(&call)?;
Ok(())
}
#[tokio::test]
async fn validate_field_arguments_rejects_drop_recreate_same_name_type() -> Result<()> {
let tmp = tempdir().unwrap();
let db = connect(tmp.path().to_str().unwrap()).execute().await?;
let schema = Arc::new(Schema::new(vec![
Field::new("keep_col", DataType::Int32, false),
Field::new("bound_col", DataType::Utf8, true),
]));
let batch = RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from(vec![1])),
Arc::new(StringArray::from(vec![Some("v")])),
],
)?;
let table = db
.create_table("binding_drop_recreate", batch)
.execute()
.await?;
let before = table.generated_column_binding_snapshot().await?;
let bound = before.field("bound_col").expect("bound_col");
let old_id = bound.field_id();
assert_eq!(bound.field().data_type(), &DataType::Utf8);
let function = one_field_function()?;
let call = FunctionCall::try_new(
&function,
vec![(
"payload_arg".to_string(),
FunctionArgument::try_field(old_id, DataType::Utf8)?,
)],
)?;
before.validate_field_arguments(&call)?;
table.drop_columns(&["bound_col"]).await?;
table
.add_columns()
.transform(NewColumnTransform::SqlExpressions(vec![(
"bound_col".into(),
"cast(NULL as string)".into(),
)]))
.execute()
.await?;
let after = table.generated_column_binding_snapshot().await?;
let recreated = after.field("bound_col").expect("recreated bound_col");
assert_eq!(recreated.field().data_type(), &DataType::Utf8);
assert_ne!(
recreated.field_id(),
old_id,
"drop/recreate must allocate a new stable field id"
);
let err = after
.validate_field_arguments(&call)
.expect_err("old call must not bind by name");
assert!(matches!(err, Error::InvalidInput { .. }));
let message = err.to_string();
assert!(
message.contains(&old_id.to_string()),
"diagnostics may name missing field id: {message}"
);
assert!(
!message.contains("bound_col") && !message.contains("keep_col"),
"diagnostics must not invent or use a column name: {message}"
);
Ok(())
}
#[tokio::test]
async fn validate_field_arguments_rejects_cast_that_allocates_new_field_id() -> Result<()> {
// Native Lance cast_to allocates a new stable field ID. The old bound call
// must fail because that ID is absent. Same-ID exact-type mismatch is proved
// separately via manually constructed snapshots (Remote/other producers).
let tmp = tempdir().unwrap();
let db = connect(tmp.path().to_str().unwrap()).execute().await?;
let schema = Arc::new(Schema::new(vec![
Field::new("label_col", DataType::Utf8, true),
Field::new("metric_col", DataType::Int32, false),
]));
let batch = RecordBatch::try_new(
schema,
vec![
Arc::new(StringArray::from(vec![Some("x")])),
Arc::new(Int32Array::from(vec![3])),
],
)?;
let table = db
.create_table("binding_type_change", batch)
.execute()
.await?;
let before = table.generated_column_binding_snapshot().await?;
let metric = before.field("metric_col").expect("metric_col");
let old_metric_id = metric.field_id();
assert_eq!(metric.field().data_type(), &DataType::Int32);
let id = FunctionId::try_new("fn.exact.binding.type-change")?;
let function = Function::new(
id,
FunctionSignature::try_new(
vec![FunctionParameter::new("metric_value", DataType::Int32)],
sample_output(),
)?,
);
let call = FunctionCall::try_new(
&function,
vec![(
"metric_value".to_string(),
FunctionArgument::try_field(old_metric_id, DataType::Int32)?,
)],
)?;
before.validate_field_arguments(&call)?;
table
.alter_columns(&[ColumnAlteration::new("metric_col".into()).cast_to(DataType::Int64)])
.await?;
let after = table.generated_column_binding_snapshot().await?;
let casted = after.field("metric_col").expect("metric_col");
assert_eq!(casted.field().data_type(), &DataType::Int64);
assert_ne!(
casted.field_id(),
old_metric_id,
"Native Lance cast_to must allocate a new stable field id"
);
let err = after
.validate_field_arguments(&call)
.expect_err("old call must fail because the prior stable field id is absent");
assert!(matches!(err, Error::InvalidInput { .. }));
let message = err.to_string();
assert!(
message.contains(&old_metric_id.to_string()),
"diagnostics may name missing field id: {message}"
);
assert!(
!message.contains("metric_col") && !message.contains("label_col"),
"diagnostics must not invent or use a column name: {message}"
);
Ok(())
}

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