mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-04 12:38:38 +00:00
Compare commits
113 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 3a3ddfda01 | |||
| c4371eb500 | |||
| 7a08580400 | |||
| 2fccab172f | |||
| e478b80985 | |||
| 72767b17fa | |||
| d7d25cd5ef | |||
| 4843445a7e | |||
| 713510375b | |||
| a9617bf830 | |||
| 208787ae7b | |||
| 9b46e7a448 | |||
| 7a46da2e67 | |||
| 705f7e7760 | |||
| d38a566282 | |||
| 5a27c71ab8 | |||
| 76f6487d92 | |||
| 6f6c3c33e0 | |||
| 1aa3665d67 | |||
| 0194f2317a | |||
| c2a647189d | |||
| df67ee4028 | |||
| d9e41228c8 | |||
| 68597070d2 | |||
| c825737780 | |||
| 0a43795996 | |||
| 0fa2fa05ad | |||
| 93ba442ac2 | |||
| 7a94ab7d6c | |||
| 6ed1a25439 | |||
| ca1d04db25 | |||
| efe3300404 | |||
| ecf87f6371 | |||
| 47213e31f8 | |||
| f65bf89c98 | |||
| d902144605 | |||
| a49dc5c71d | |||
| 98fed41efa | |||
| 1524ee0669 | |||
| 29be3e5509 | |||
| 8cedd50495 | |||
| b71ada0fae | |||
| 206efd98ff | |||
| 65c0968c0f | |||
| 2b10f2a7ce | |||
| f8bb90405f | |||
| 76aac96749 | |||
| 0093bc8179 | |||
| ac35a687f1 | |||
| 203f6536a6 | |||
| 9d3d0d0640 | |||
| a9ed8dba27 | |||
| 04acf1d3b5 | |||
| 3746118374 | |||
| d0b5cbe510 | |||
| 7b195adc3a | |||
| 818d6d1f59 | |||
| 9d589bea44 | |||
| 1798ece362 | |||
| 82b82711ba | |||
| a615306f39 | |||
| 920fc0e455 | |||
| 5acce6782e | |||
| 12405a4077 | |||
| 36054be576 | |||
| 77a93fee76 | |||
| 7bb501839a | |||
| 5b347afd99 | |||
| 706a9c327f | |||
| be290447d9 | |||
| 79ba076429 | |||
| ec21e37040 | |||
| 6ba80a960c | |||
| 11f24b1df4 | |||
| 2ba7407dc3 | |||
| 607e556927 | |||
| 564e5d0d56 | |||
| dd5cb4d805 | |||
| dbc3687c7b | |||
| ec80acb668 | |||
| fc44535cee | |||
| 4048150fdd | |||
| 2922c171f7 | |||
| c5f9efefe9 | |||
| f4c668e244 | |||
| b1cfe6edb1 | |||
| 001237c7a4 | |||
| 369b10a377 | |||
| 1c3cd1d918 | |||
| 9707966943 | |||
| 62fe413a52 | |||
| 1493ece3de | |||
| e6444ecc05 | |||
| cc0139c136 | |||
| b20696ef9c | |||
| 772bdeced8 | |||
| c1a3fa7f51 | |||
| 0ba82873c5 | |||
| 3af51541a0 | |||
| 2c06a48bd8 | |||
| ac8b28c010 | |||
| 173f889d2a | |||
| 03b52e5877 | |||
| 798e5364fb | |||
| f1f34dfdd3 | |||
| 123c921c4f | |||
| 99a68db78c | |||
| 9e73d440a3 | |||
| 3956d9dbfa | |||
| 16e1967efc | |||
| 27dd92c67e | |||
| 9e2e711c7a | |||
| c3176a47ce |
+1
-1
@@ -1,5 +1,5 @@
|
||||
[tool.bumpversion]
|
||||
current_version = "0.37.1-beta.0"
|
||||
current_version = "0.37.1-beta.1"
|
||||
parse = """(?x)
|
||||
(?P<major>0|[1-9]\\d*)\\.
|
||||
(?P<minor>0|[1-9]\\d*)\\.
|
||||
|
||||
@@ -0,0 +1,243 @@
|
||||
name: Check doc links
|
||||
|
||||
# Checking external links is inherently noisy: third-party sites rate-limit
|
||||
# automated clients, reject non-browser user agents, and go down temporarily.
|
||||
# Blocking pull requests on that trades a lot of false failures for very little
|
||||
# signal, so this runs on a schedule and reports findings in a single tracking
|
||||
# issue instead of failing anyone's build.
|
||||
on:
|
||||
schedule:
|
||||
- cron: "0 7 * * *"
|
||||
workflow_dispatch:
|
||||
|
||||
# The report lives in one repository-global issue, so runs must not overlap: a
|
||||
# lookup racing a create produces duplicate issues, and a healthy run closing
|
||||
# the issue while a failing run only rewrites its body would leave a broken
|
||||
# report closed. The group is deliberately ref-independent so that a manual
|
||||
# dispatch serializes against the scheduled run.
|
||||
concurrency:
|
||||
group: docs-link-check
|
||||
cancel-in-progress: false
|
||||
|
||||
permissions: {}
|
||||
|
||||
env:
|
||||
REPORT_TITLE: "Docs link checker report"
|
||||
|
||||
jobs:
|
||||
scan:
|
||||
name: Scan links
|
||||
runs-on: ubuntu-24.04
|
||||
# lychee-action is pinned by SHA, but its wrapper downloads the lychee
|
||||
# release tarball at run time without verifying a digest, and hands the
|
||||
# resulting binary a GitHub token. Release assets remain replaceable, so
|
||||
# that binary is confined to a job whose token can only read public
|
||||
# content; everything that writes runs in the report job below.
|
||||
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
|
||||
with:
|
||||
# workflow_dispatch can run from any ref, but the report is
|
||||
# repository-global. Always measure the default branch so a manual
|
||||
# run from a topic branch cannot close a report that main warrants,
|
||||
# or overwrite it with branch-only findings.
|
||||
ref: ${{ github.event.repository.default_branch }}
|
||||
persist-credentials: false
|
||||
|
||||
- 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
|
||||
# API reference (the js/ tree comes from `npm run docs` in nodejs)
|
||||
# and the hand-written pages use mkdocstrings cross-references and
|
||||
# nav-relative paths that only resolve in the site mkdocs builds,
|
||||
# not in this checkout, so relative links would be reported as
|
||||
# broken on every run.
|
||||
args: >-
|
||||
--scheme https
|
||||
--scheme http
|
||||
--no-progress
|
||||
--max-retries 3
|
||||
--timeout 20
|
||||
'docs/src/**/*.md'
|
||||
format: json
|
||||
output: ./lychee/out.json
|
||||
jobSummary: false
|
||||
# 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 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: |
|
||||
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.validate.outputs.status == 'findings'
|
||||
uses: actions/upload-artifact@v7
|
||||
with:
|
||||
name: link-report
|
||||
path: ./lychee/out.json
|
||||
retention-days: 7
|
||||
|
||||
report:
|
||||
name: Update report issue
|
||||
needs: scan
|
||||
runs-on: ubuntu-24.04
|
||||
# Deliberately no checkout: this job needs the report artifact and the
|
||||
# issues API, not the repository contents.
|
||||
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: Find existing report issue
|
||||
id: report
|
||||
# Matched on title alone, and through search rather than a listing:
|
||||
# the issue action applies labels in a separate call after creating the
|
||||
# issue, so a label filter misses a half-created report, and this
|
||||
# repository has far more open issues than one listing page holds.
|
||||
# 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 a problem recurs.
|
||||
run: |
|
||||
match=$(gh issue list --repo "$GITHUB_REPOSITORY" --state all \
|
||||
--search "in:title \"$REPORT_TITLE\" author:app/github-actions" \
|
||||
--limit 50 --json number,title,state \
|
||||
--jq "[.[] | select(.title == \"$REPORT_TITLE\")] | sort_by(.number) | first // empty")
|
||||
echo "number=$(jq -r '.number // empty' <<<"$match")" >> "$GITHUB_OUTPUT"
|
||||
echo "state=$(jq -r '.state // empty' <<<"$match")" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Download report
|
||||
if: env.STATUS == 'findings'
|
||||
uses: actions/download-artifact@v8
|
||||
with:
|
||||
name: link-report
|
||||
path: ./lychee
|
||||
|
||||
- name: Compose report
|
||||
if: env.STATUS == 'findings'
|
||||
run: |
|
||||
run_url="$GITHUB_SERVER_URL/$GITHUB_REPOSITORY/actions/runs/$GITHUB_RUN_ID"
|
||||
{
|
||||
echo "Broken documentation links found by [\`$GITHUB_WORKFLOW\`]($run_url)."
|
||||
echo
|
||||
echo "This issue is rewritten by every scheduled run and closed automatically once all links resolve."
|
||||
echo
|
||||
echo "Entries can be false positives: some sites rate-limit or block automated clients while working fine in a browser. Confirm before editing the docs, and add persistent offenders to \`--exclude\` in \`.github/workflows/docs-link-check.yml\`."
|
||||
echo
|
||||
# Timeouts are reported alongside errors: entries land in
|
||||
# timeout_map with a status text instead of an HTTP code.
|
||||
jq -r '
|
||||
"\(.errors) of \(.total) links failed, \(.timeouts) timed out.",
|
||||
"",
|
||||
([(.error_map | to_entries[]), (.timeout_map | to_entries[])]
|
||||
| group_by(.key)[] |
|
||||
"### Errors in \(.[0].key)",
|
||||
"",
|
||||
(map(.value[])[] | "* [\(.status.code // .status.text // "ERR")] <\(.url)> — \(.status.details // .status.text // "unknown error")"),
|
||||
"")
|
||||
' ./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, 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 "The documentation link checker reported a problem again in [the latest run]($run_url)."
|
||||
|
||||
- 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
|
||||
# the same issue is updated in place.
|
||||
issue-number: ${{ steps.report.outputs.number }}
|
||||
title: ${{ env.REPORT_TITLE }}
|
||||
content-filepath: ./lychee/issue.md
|
||||
labels: documentation
|
||||
|
||||
- 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.STATUS == 'healthy' &&
|
||||
steps.report.outputs.state == 'OPEN'
|
||||
env:
|
||||
ISSUE_NUMBER: ${{ steps.report.outputs.number }}
|
||||
run: |
|
||||
run_url="$GITHUB_SERVER_URL/$GITHUB_REPOSITORY/actions/runs/$GITHUB_RUN_ID"
|
||||
gh issue close "$ISSUE_NUMBER" --repo "$GITHUB_REPOSITORY" \
|
||||
--comment "All documentation links resolved in [the latest run]($run_url)."
|
||||
@@ -296,16 +296,18 @@ jobs:
|
||||
cargo update -p aws-types --precise 1.3.9
|
||||
cargo update -p aws-sigv4 --precise 1.3.5
|
||||
cargo update -p aws-credential-types --precise 1.2.8
|
||||
cargo update -p aws-smithy-checksums --precise 0.63.9
|
||||
# aws-smithy-checksums must stay at or above 0.63.13: OpenDAL's S3
|
||||
# service needs crc-fast ~1.9, and older releases pin it to ~1.3.
|
||||
cargo update -p aws-smithy-checksums --precise 0.63.13
|
||||
cargo update -p aws-smithy-runtime --precise 1.9.3
|
||||
cargo update -p aws-smithy-http --precise 0.62.4
|
||||
cargo update -p aws-smithy-eventstream --precise 0.60.12
|
||||
cargo update -p aws-smithy-http --precise 0.62.6
|
||||
cargo update -p aws-smithy-eventstream --precise 0.60.14
|
||||
cargo update -p aws-smithy-http-client --precise 1.1.3
|
||||
cargo update -p aws-smithy-observability --precise 0.1.4
|
||||
cargo update -p aws-smithy-query --precise 0.60.8
|
||||
cargo update -p aws-smithy-runtime-api --precise 1.9.1
|
||||
cargo update -p aws-smithy-async --precise 1.2.6
|
||||
cargo update -p aws-smithy-types --precise 1.3.5
|
||||
cargo update -p aws-smithy-runtime-api --precise 1.9.3
|
||||
cargo update -p aws-smithy-async --precise 1.2.7
|
||||
cargo update -p aws-smithy-types --precise 1.3.6
|
||||
cargo update -p aws-smithy-xml --precise 0.60.11
|
||||
cargo update -p home --precise 0.5.9
|
||||
- name: cargo +${{ matrix.msrv }} check
|
||||
|
||||
Generated
+264
-239
File diff suppressed because it is too large
Load Diff
+15
-15
@@ -13,20 +13,20 @@ categories = ["database-implementations"]
|
||||
rust-version = "1.91.0"
|
||||
|
||||
[workspace.dependencies]
|
||||
lance = { "version" = "=10.1.0-beta.1", default-features = false, "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-core = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datagen = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-file = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-io = { "version" = "=10.1.0-beta.1", default-features = false, "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-index = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-linalg = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace-impls = { "version" = "=10.1.0-beta.1", default-features = false, "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-table = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-testing = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datafusion = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-encoding = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-arrow = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "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 }
|
||||
@@ -52,7 +52,7 @@ env_logger = "0.11"
|
||||
half = { "version" = "2.7.1", default-features = false, features = [
|
||||
"num-traits",
|
||||
] }
|
||||
futures = "0"
|
||||
futures = "0.3"
|
||||
log = "0.4"
|
||||
metrics = "0.24"
|
||||
metrics-util = "0.19"
|
||||
|
||||
@@ -14,7 +14,7 @@ Add the following dependency to your `pom.xml`:
|
||||
<dependency>
|
||||
<groupId>com.lancedb</groupId>
|
||||
<artifactId>lancedb-core</artifactId>
|
||||
<version>0.37.1-beta.0</version>
|
||||
<version>0.37.1-beta.1</version>
|
||||
</dependency>
|
||||
```
|
||||
|
||||
|
||||
@@ -431,9 +431,10 @@ Read the [LsmWriteSpec](../interfaces/LsmWriteSpec.md) currently installed on th
|
||||
|
||||
Resolves to `undefined` when the MemWAL LSM write path is not enabled (no
|
||||
spec has been set, or it was removed with [Table#unsetLsmWriteSpec](Table.md#unsetlsmwritespec)).
|
||||
The returned spec — including its `maintainedIndexes` and
|
||||
`writerConfigDefaults` — mirrors what was passed to
|
||||
[Table#setLsmWriteSpec](Table.md#setlsmwritespec).
|
||||
The returned spec mirrors what was passed to
|
||||
[Table#setLsmWriteSpec](Table.md#setlsmwritespec), except that `maintainedIndexes` always
|
||||
reports the concrete list resolved when the spec was set — `undefined`
|
||||
never round-trips.
|
||||
|
||||
#### Returns
|
||||
|
||||
@@ -806,6 +807,11 @@ All variants require the table to have an unenforced primary key
|
||||
([Table#setUnenforcedPrimaryKey](Table.md#setunenforcedprimarykey)); bucket sharding additionally
|
||||
requires it to be the single column being bucketed.
|
||||
|
||||
Omitting `maintainedIndexes` maintains every index on the table, resolved
|
||||
here, failing if one cannot be maintained — name them to install anyway.
|
||||
Naming them pins an exact set, and a still-building index is rejected
|
||||
rather than quietly omitted.
|
||||
|
||||
#### Parameters
|
||||
|
||||
* **spec**: [`LsmWriteSpec`](../interfaces/LsmWriteSpec.md)
|
||||
|
||||
@@ -34,7 +34,9 @@ Bucket and identity variants: the sharding column.
|
||||
optional maintainedIndexes: string[];
|
||||
```
|
||||
|
||||
Names of indexes the MemWAL should keep up to date during writes.
|
||||
Indexes the MemWAL keeps up to date. Omit to maintain every supported
|
||||
index, resolved on install — a snapshot, so indexes created later are not
|
||||
maintained. Pass `[]` for none.
|
||||
|
||||
***
|
||||
|
||||
|
||||
@@ -44,4 +44,7 @@ The number of rows in the table
|
||||
totalBytes: number;
|
||||
```
|
||||
|
||||
The total number of bytes in the table
|
||||
The total size, in bytes, of the table's data files, index files, and
|
||||
overlay files
|
||||
|
||||
Read from the manifest, so this excludes deletion files and manifests.
|
||||
|
||||
@@ -31,7 +31,7 @@ is also an [asynchronous API client](#connections-asynchronous).
|
||||
## Namespaces (Synchronous)
|
||||
|
||||
A namespace-backed connection resolves tables through a
|
||||
[Lance namespace](https://lancedb.github.io/lance-namespace/) service instead of
|
||||
[Lance namespace](https://lance-format.github.io/lance-namespace/) service instead of
|
||||
listing a storage directory.
|
||||
|
||||
::: lancedb.connect_namespace
|
||||
|
||||
@@ -8,7 +8,7 @@
|
||||
<parent>
|
||||
<groupId>com.lancedb</groupId>
|
||||
<artifactId>lancedb-parent</artifactId>
|
||||
<version>0.37.1-beta.0</version>
|
||||
<version>0.37.1-beta.1</version>
|
||||
<relativePath>../pom.xml</relativePath>
|
||||
</parent>
|
||||
|
||||
|
||||
+2
-2
@@ -6,7 +6,7 @@
|
||||
|
||||
<groupId>com.lancedb</groupId>
|
||||
<artifactId>lancedb-parent</artifactId>
|
||||
<version>0.37.1-beta.0</version>
|
||||
<version>0.37.1-beta.1</version>
|
||||
<packaging>pom</packaging>
|
||||
<name>${project.artifactId}</name>
|
||||
<description>LanceDB Java SDK Parent POM</description>
|
||||
@@ -28,7 +28,7 @@
|
||||
<properties>
|
||||
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
|
||||
<arrow.version>15.0.0</arrow.version>
|
||||
<lance-core.version>10.1.0-beta.1</lance-core.version>
|
||||
<lance-core.version>11.0.0-beta.3</lance-core.version>
|
||||
<spotless.skip>false</spotless.skip>
|
||||
<spotless.version>2.30.0</spotless.version>
|
||||
<spotless.java.googlejavaformat.version>1.7</spotless.java.googlejavaformat.version>
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[package]
|
||||
name = "lancedb-nodejs"
|
||||
edition.workspace = true
|
||||
version = "0.37.1-beta.0"
|
||||
version = "0.37.1-beta.1"
|
||||
publish = false
|
||||
license.workspace = true
|
||||
description.workspace = true
|
||||
|
||||
@@ -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"];
|
||||
@@ -197,6 +204,35 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
|
||||
expect(table.getChild("d")?.toJSON()).toEqual([9n, 10n, null]);
|
||||
});
|
||||
|
||||
it("will use a provided FixedSizeList schema with typed array values", function () {
|
||||
const schema = new Schema([
|
||||
new Field("text", new Utf8(), false),
|
||||
new Field(
|
||||
"vector",
|
||||
new FixedSizeList(3, new Field("item", new Float32(), false)),
|
||||
false,
|
||||
),
|
||||
]);
|
||||
|
||||
const table = makeArrowTable(
|
||||
[
|
||||
{
|
||||
text: "foo",
|
||||
vector: new Float32Array([1, 2, 3]),
|
||||
},
|
||||
],
|
||||
{ schema },
|
||||
);
|
||||
|
||||
expect(table.getChild("text")?.toJSON()).toEqual(["foo"]);
|
||||
expect(
|
||||
table
|
||||
.getChild("vector")
|
||||
?.toJSON()
|
||||
.map((value) => value.toJSON()),
|
||||
).toEqual([[1, 2, 3]]);
|
||||
});
|
||||
|
||||
it("will assume the column `vector` is FixedSizeList<Float32> by default", async function () {
|
||||
const schema = new Schema([
|
||||
new Field("a", new Float(Precision.DOUBLE), true),
|
||||
@@ -1025,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()),
|
||||
|
||||
@@ -11,8 +11,11 @@ import {
|
||||
Float16,
|
||||
Float32,
|
||||
Float64,
|
||||
Int32,
|
||||
Schema,
|
||||
Utf8,
|
||||
fromDataToBuffer,
|
||||
tableFromIPC,
|
||||
} from "../lancedb/arrow";
|
||||
import { EmbeddingFunction, LanceSchema } from "../lancedb/embedding";
|
||||
import { getRegistry, register } from "../lancedb/embedding/registry";
|
||||
@@ -184,6 +187,63 @@ describe("embedding functions", () => {
|
||||
const vector0 = JSON.parse(JSON.stringify(arr[0].vector));
|
||||
expect(vector0).toEqual([1, 2, 3]);
|
||||
});
|
||||
|
||||
it("should append generated vectors to a non-nullable schema", async () => {
|
||||
@register("non_nullable_schema_test")
|
||||
class MockEmbeddingFunction extends EmbeddingFunction<string> {
|
||||
ndims() {
|
||||
return 3;
|
||||
}
|
||||
embeddingDataType(): Float {
|
||||
return new Float64();
|
||||
}
|
||||
async computeSourceEmbeddings(data: string[]) {
|
||||
return data.map(() => [1, 2, 3]);
|
||||
}
|
||||
}
|
||||
|
||||
const schema = new Schema([
|
||||
new Field("id", new Int32()),
|
||||
new Field("text", new Utf8()),
|
||||
new Field("type", new Utf8()),
|
||||
new Field(
|
||||
"vector",
|
||||
new FixedSizeList(3, new Field("item", new Float64())),
|
||||
),
|
||||
]);
|
||||
const func = new MockEmbeddingFunction();
|
||||
const db = await connect(tmpDir.name);
|
||||
const table = await db.createEmptyTable("test_non_nullable", schema, {
|
||||
embeddingFunction: {
|
||||
function: func,
|
||||
sourceColumn: "text",
|
||||
},
|
||||
});
|
||||
|
||||
const data = [
|
||||
{ id: 1, text: "Carrot", type: "vegetable" },
|
||||
{ id: 2, text: "Apple", type: "fruit" },
|
||||
];
|
||||
const buffer = await fromDataToBuffer(
|
||||
data,
|
||||
undefined,
|
||||
await table.schema(),
|
||||
);
|
||||
const generatedTable = tableFromIPC(buffer);
|
||||
const vectorField = generatedTable.schema.fields.find(
|
||||
(field) => field.name === "vector",
|
||||
);
|
||||
expect(vectorField?.nullable).toBe(false);
|
||||
|
||||
await table.add(data);
|
||||
|
||||
const rows = await table.query().toArray();
|
||||
expect(rows).toHaveLength(2);
|
||||
for (const row of rows) {
|
||||
expect([...row.vector]).toEqual([1, 2, 3]);
|
||||
}
|
||||
});
|
||||
|
||||
it("should error when appending to a table with an unregistered embedding function", async () => {
|
||||
@register("mock")
|
||||
class MockEmbeddingFunction extends EmbeddingFunction<string> {
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
import packageJson = require("../package.json");
|
||||
|
||||
describe("package metadata", () => {
|
||||
it("requires Node.js type declarations compatible with the runtime", () => {
|
||||
expect(packageJson.engines.node).toBe(">= 18");
|
||||
expect(packageJson.peerDependencies["@types/node"]).toBe(">=18");
|
||||
expect(packageJson.peerDependenciesMeta["@types/node"]).toEqual({
|
||||
optional: true,
|
||||
});
|
||||
});
|
||||
});
|
||||
@@ -110,6 +110,81 @@ describe("Query outputSchema", () => {
|
||||
});
|
||||
});
|
||||
|
||||
describe("Search pagination", () => {
|
||||
let tmpDir: tmp.DirResult;
|
||||
let table: Table;
|
||||
|
||||
beforeEach(async () => {
|
||||
tmpDir = tmp.dirSync({ unsafeCleanup: true });
|
||||
const db = await connect(tmpDir.name);
|
||||
const schema = new Schema([
|
||||
new Field("id", new Int64(), false),
|
||||
new Field("text", new Utf8(), false),
|
||||
new Field(
|
||||
"vector",
|
||||
new FixedSizeList(2, new Field("item", new Float32())),
|
||||
false,
|
||||
),
|
||||
]);
|
||||
const data = makeArrowTable(
|
||||
[
|
||||
{ id: 1n, text: "common", vector: [0, 0] },
|
||||
{ id: 2n, text: "common common", vector: [1, 1] },
|
||||
{ id: 3n, text: "common common common", vector: [2, 2] },
|
||||
{ id: 4n, text: "common common common common", vector: [3, 3] },
|
||||
],
|
||||
{ schema },
|
||||
);
|
||||
table = await db.createTable("test", data);
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
tmpDir.removeCallback();
|
||||
});
|
||||
|
||||
it("applies offset after the vector search limit", async () => {
|
||||
const allResults = await table
|
||||
.vectorSearch([0, 0])
|
||||
.select(["id"])
|
||||
.limit(4)
|
||||
.toArray();
|
||||
const secondPage = await table
|
||||
.vectorSearch([0, 0])
|
||||
.select(["id"])
|
||||
.limit(2)
|
||||
.offset(2)
|
||||
.toArray();
|
||||
|
||||
expect(allResults).toHaveLength(4);
|
||||
expect(secondPage).toHaveLength(2);
|
||||
expect(secondPage.map((row) => row.id)).toEqual(
|
||||
allResults.slice(2, 4).map((row) => row.id),
|
||||
);
|
||||
});
|
||||
|
||||
it("applies offset after the full-text search limit", async () => {
|
||||
await table.createIndex("text", { config: Index.fts() });
|
||||
|
||||
const allResults = await table
|
||||
.search("common", "fts")
|
||||
.select(["id"])
|
||||
.limit(4)
|
||||
.toArray();
|
||||
const secondPage = await table
|
||||
.search("common", "fts")
|
||||
.select(["id"])
|
||||
.limit(2)
|
||||
.offset(2)
|
||||
.toArray();
|
||||
|
||||
expect(allResults).toHaveLength(4);
|
||||
expect(secondPage).toHaveLength(2);
|
||||
expect(secondPage.map((row) => row.id)).toEqual(
|
||||
allResults.slice(2, 4).map((row) => row.id),
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
describe("Query orderBy", () => {
|
||||
let tmpDir: tmp.DirResult;
|
||||
let table: Table;
|
||||
|
||||
@@ -170,6 +170,38 @@ describe("remote connection", () => {
|
||||
);
|
||||
});
|
||||
|
||||
it("surfaces JSON server errors from remote table operations", async () => {
|
||||
await withMockDatabase(
|
||||
(req, res) => {
|
||||
const path = req.url ?? "";
|
||||
if (path.endsWith("/describe/")) {
|
||||
res.writeHead(200, { "Content-Type": "application/json" }).end(
|
||||
JSON.stringify({
|
||||
name: "broken_table",
|
||||
version: 1,
|
||||
schema: { fields: [] },
|
||||
}),
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
if (path.endsWith("/count_rows/")) {
|
||||
res
|
||||
.writeHead(400, { "Content-Type": "application/json" })
|
||||
.end(JSON.stringify({ error: "count rows failed" }));
|
||||
return;
|
||||
}
|
||||
|
||||
res.writeHead(404).end();
|
||||
},
|
||||
async (db) => {
|
||||
const table = await db.openTable("broken_table");
|
||||
|
||||
await expect(table.countRows()).rejects.toThrow("count rows failed");
|
||||
},
|
||||
);
|
||||
});
|
||||
|
||||
it("should pass on requested extra headers", async () => {
|
||||
await withMockDatabase(
|
||||
(req, res) => {
|
||||
|
||||
@@ -86,6 +86,44 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
|
||||
await expect(table.countRows()).resolves.toBe(3);
|
||||
});
|
||||
|
||||
it("should support a foreign Float64 vector schema end to end", async () => {
|
||||
const conn = await connect(tmpDir.name);
|
||||
const schema = new arrow.Schema([
|
||||
new arrow.Field("resource_id", new arrow.Int32(), false),
|
||||
new arrow.Field(
|
||||
"vector",
|
||||
new arrow.FixedSizeList(
|
||||
3,
|
||||
new arrow.Field("value", new arrow.Float64(), true),
|
||||
),
|
||||
false,
|
||||
),
|
||||
]);
|
||||
const data = [
|
||||
{
|
||||
// biome-ignore lint/style/useNamingConvention: matches the reported schema
|
||||
resource_id: 0,
|
||||
vector: [0.1, 0.1, 0.1],
|
||||
},
|
||||
];
|
||||
|
||||
const resources = await conn.createTable("resources", data, { schema });
|
||||
|
||||
const existing = await resources
|
||||
.query()
|
||||
.where("resource_id = 0")
|
||||
.limit(1)
|
||||
.toArray();
|
||||
expect(existing).toHaveLength(1);
|
||||
|
||||
const matched = await resources
|
||||
.search(Float64Array.from(data[0].vector))
|
||||
.limit(1)
|
||||
.toArray();
|
||||
expect(matched).toHaveLength(1);
|
||||
expect(matched[0]["resource_id"]).toBe(0);
|
||||
});
|
||||
|
||||
it("should support branches", async () => {
|
||||
await table.add([{ id: 1 }]);
|
||||
expect(await table.countRows()).toBe(1);
|
||||
@@ -239,8 +277,16 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
|
||||
},
|
||||
numIndices: 0,
|
||||
numRows: 3,
|
||||
totalBytes: 44,
|
||||
// Full on-disk size of the two data files, footers and metadata included.
|
||||
totalBytes: 684,
|
||||
});
|
||||
|
||||
// Index files count toward totalBytes too (only deletion files and
|
||||
// manifests are excluded).
|
||||
await table.createIndex("id", { config: Index.btree() });
|
||||
const statsWithIndex = await table.stats();
|
||||
expect(statsWithIndex.numIndices).toBe(1);
|
||||
expect(statsWithIndex.totalBytes).toBeGreaterThan(684);
|
||||
});
|
||||
|
||||
it("should overwrite data if asked", async () => {
|
||||
|
||||
+174
-29
@@ -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 = {
|
||||
|
||||
+14
-4
@@ -197,7 +197,11 @@ export interface LsmWriteSpec {
|
||||
column?: string;
|
||||
/** Bucket variant: the number of buckets, in `[1, 1024]`. */
|
||||
numBuckets?: number;
|
||||
/** Names of indexes the MemWAL should keep up to date during writes. */
|
||||
/**
|
||||
* Indexes the MemWAL keeps up to date. Omit to maintain every supported
|
||||
* index, resolved on install — a snapshot, so indexes created later are not
|
||||
* maintained. Pass `[]` for none.
|
||||
*/
|
||||
maintainedIndexes?: string[];
|
||||
/** Default `ShardWriter` configuration recorded in the MemWAL index. */
|
||||
writerConfigDefaults?: Record<string, string>;
|
||||
@@ -595,6 +599,11 @@ export abstract class Table {
|
||||
* All variants require the table to have an unenforced primary key
|
||||
* ({@link Table#setUnenforcedPrimaryKey}); bucket sharding additionally
|
||||
* requires it to be the single column being bucketed.
|
||||
*
|
||||
* Omitting `maintainedIndexes` maintains every index on the table, resolved
|
||||
* here, failing if one cannot be maintained — name them to install anyway.
|
||||
* Naming them pins an exact set, and a still-building index is rejected
|
||||
* rather than quietly omitted.
|
||||
* @param {LsmWriteSpec} spec The sharding spec to install.
|
||||
* @returns {Promise<void>}
|
||||
* @example
|
||||
@@ -622,9 +631,10 @@ export abstract class Table {
|
||||
*
|
||||
* Resolves to `undefined` when the MemWAL LSM write path is not enabled (no
|
||||
* spec has been set, or it was removed with {@link Table#unsetLsmWriteSpec}).
|
||||
* The returned spec — including its `maintainedIndexes` and
|
||||
* `writerConfigDefaults` — mirrors what was passed to
|
||||
* {@link Table#setLsmWriteSpec}.
|
||||
* The returned spec mirrors what was passed to
|
||||
* {@link Table#setLsmWriteSpec}, except that `maintainedIndexes` always
|
||||
* reports the concrete list resolved when the spec was set — `undefined`
|
||||
* never round-trips.
|
||||
* @returns {Promise<LsmWriteSpec | undefined>}
|
||||
*/
|
||||
abstract getLsmWriteSpec(): Promise<LsmWriteSpec | undefined>;
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-darwin-arm64",
|
||||
"version": "0.37.1-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": ["darwin"],
|
||||
"cpu": ["arm64"],
|
||||
"main": "lancedb.darwin-arm64.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-linux-arm64-gnu",
|
||||
"version": "0.37.1-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": ["linux"],
|
||||
"cpu": ["arm64"],
|
||||
"main": "lancedb.linux-arm64-gnu.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-linux-arm64-musl",
|
||||
"version": "0.37.1-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": ["linux"],
|
||||
"cpu": ["arm64"],
|
||||
"main": "lancedb.linux-arm64-musl.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-linux-x64-gnu",
|
||||
"version": "0.37.1-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": ["linux"],
|
||||
"cpu": ["x64"],
|
||||
"main": "lancedb.linux-x64-gnu.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-linux-x64-musl",
|
||||
"version": "0.37.1-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": ["linux"],
|
||||
"cpu": ["x64"],
|
||||
"main": "lancedb.linux-x64-musl.node",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-win32-arm64-msvc",
|
||||
"version": "0.37.1-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": [
|
||||
"win32"
|
||||
],
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb-win32-x64-msvc",
|
||||
"version": "0.37.1-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"os": ["win32"],
|
||||
"cpu": ["x64"],
|
||||
"main": "lancedb.win32-x64-msvc.node",
|
||||
|
||||
Generated
+8
-2
@@ -1,12 +1,12 @@
|
||||
{
|
||||
"name": "@lancedb/lancedb",
|
||||
"version": "0.37.1-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"lockfileVersion": 3,
|
||||
"requires": true,
|
||||
"packages": {
|
||||
"": {
|
||||
"name": "@lancedb/lancedb",
|
||||
"version": "0.37.1-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"cpu": [
|
||||
"x64",
|
||||
"arm64"
|
||||
@@ -55,7 +55,13 @@
|
||||
"openai": "4.29.2"
|
||||
},
|
||||
"peerDependencies": {
|
||||
"@types/node": ">=18",
|
||||
"apache-arrow": ">=15.0.0 <=18.1.0"
|
||||
},
|
||||
"peerDependenciesMeta": {
|
||||
"@types/node": {
|
||||
"optional": true
|
||||
}
|
||||
}
|
||||
},
|
||||
"node_modules/@aws-crypto/crc32": {
|
||||
|
||||
+7
-1
@@ -11,7 +11,7 @@
|
||||
"ann"
|
||||
],
|
||||
"private": false,
|
||||
"version": "0.37.1-beta.0",
|
||||
"version": "0.37.1-beta.1",
|
||||
"main": "dist/index.js",
|
||||
"exports": {
|
||||
".": "./dist/index.js",
|
||||
@@ -101,6 +101,12 @@
|
||||
"openai": "4.29.2"
|
||||
},
|
||||
"peerDependencies": {
|
||||
"@types/node": ">=18",
|
||||
"apache-arrow": ">=15.0.0 <=18.1.0"
|
||||
},
|
||||
"peerDependenciesMeta": {
|
||||
"@types/node": {
|
||||
"optional": true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+11
-1
@@ -42,9 +42,19 @@ impl Job {
|
||||
}
|
||||
|
||||
/// Wait until the operation reaches a terminal state.
|
||||
///
|
||||
/// Jobs that complete without a resource result resolve successfully.
|
||||
/// Resource results are not exposed on this binding yet; unsupported
|
||||
/// success results reject with a generic error.
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn wait(&self) -> napi::Result<()> {
|
||||
self.inner.wait().await.default_error()
|
||||
match self.inner.wait().await.default_error()? {
|
||||
lancedb::JobResult::None => Ok(()),
|
||||
// JobResult is non_exhaustive; Function and future variants fail closed.
|
||||
_ => Err(napi::Error::from_reason(
|
||||
"unsupported job result".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Request cancellation. Cancelling a finished operation is a no-op.
|
||||
|
||||
+10
-7
@@ -772,7 +772,8 @@ pub struct LsmWriteSpec {
|
||||
pub column: Option<String>,
|
||||
/// Bucket variant: the number of buckets, in `[1, 1024]`.
|
||||
pub num_buckets: Option<u32>,
|
||||
/// Names of indexes the MemWAL should keep up to date during writes.
|
||||
/// Indexes the MemWAL keeps up to date. Omitted resolves every
|
||||
/// maintainable index on install; an empty array means none.
|
||||
pub maintained_indexes: Option<Vec<String>>,
|
||||
/// Default `ShardWriter` configuration recorded in the MemWAL index.
|
||||
pub writer_config_defaults: Option<HashMap<String, String>>,
|
||||
@@ -782,7 +783,6 @@ impl TryFrom<LsmWriteSpec> for lancedb::table::LsmWriteSpec {
|
||||
type Error = napi::Error;
|
||||
|
||||
fn try_from(value: LsmWriteSpec) -> napi::Result<Self> {
|
||||
let maintained = value.maintained_indexes.unwrap_or_default();
|
||||
let writer_config_defaults = value.writer_config_defaults.unwrap_or_default();
|
||||
let spec = match value.spec_type.as_str() {
|
||||
"bucket" => {
|
||||
@@ -809,7 +809,7 @@ impl TryFrom<LsmWriteSpec> for lancedb::table::LsmWriteSpec {
|
||||
}
|
||||
};
|
||||
Ok(spec
|
||||
.with_maintained_indexes(maintained)
|
||||
.with_maintained_indexes(value.maintained_indexes)
|
||||
.with_writer_config_defaults(writer_config_defaults))
|
||||
}
|
||||
}
|
||||
@@ -827,7 +827,7 @@ impl From<lancedb::table::LsmWriteSpec> for LsmWriteSpec {
|
||||
spec_type: "bucket".to_string(),
|
||||
column: Some(column),
|
||||
num_buckets: Some(num_buckets),
|
||||
maintained_indexes: Some(maintained_indexes),
|
||||
maintained_indexes,
|
||||
writer_config_defaults: Some(writer_config_defaults),
|
||||
},
|
||||
Native::Identity {
|
||||
@@ -838,7 +838,7 @@ impl From<lancedb::table::LsmWriteSpec> for LsmWriteSpec {
|
||||
spec_type: "identity".to_string(),
|
||||
column: Some(column),
|
||||
num_buckets: None,
|
||||
maintained_indexes: Some(maintained_indexes),
|
||||
maintained_indexes,
|
||||
writer_config_defaults: Some(writer_config_defaults),
|
||||
},
|
||||
Native::Unsharded {
|
||||
@@ -848,7 +848,7 @@ impl From<lancedb::table::LsmWriteSpec> for LsmWriteSpec {
|
||||
spec_type: "unsharded".to_string(),
|
||||
column: None,
|
||||
num_buckets: None,
|
||||
maintained_indexes: Some(maintained_indexes),
|
||||
maintained_indexes,
|
||||
writer_config_defaults: Some(writer_config_defaults),
|
||||
},
|
||||
}
|
||||
@@ -1043,7 +1043,10 @@ impl From<lancedb::index::IndexStatistics> for IndexStatistics {
|
||||
|
||||
#[napi(object)]
|
||||
pub struct TableStatistics {
|
||||
/// The total number of bytes in the table
|
||||
/// The total size, in bytes, of the table's data files, index files, and
|
||||
/// overlay files
|
||||
///
|
||||
/// Read from the manifest, so this excludes deletion files and manifests.
|
||||
pub total_bytes: i64,
|
||||
|
||||
/// The number of rows in the table
|
||||
|
||||
+3
-3
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "lancedb-python"
|
||||
version = "0.37.1-beta.0"
|
||||
version = "0.37.1-beta.1"
|
||||
publish = false
|
||||
edition.workspace = true
|
||||
description = "Python bindings for LanceDB"
|
||||
@@ -26,7 +26,7 @@ lance-namespace-impls.workspace = true
|
||||
lance-io.workspace = true
|
||||
env_logger.workspace = true
|
||||
log.workspace = true
|
||||
pyo3 = { version = "0.28", features = ["extension-module", "abi3-py39", "chrono"] }
|
||||
pyo3 = { version = "0.28", features = ["extension-module", "abi3-py310", "chrono"] }
|
||||
chrono = { version = "0.4", default-features = false, features = ["clock"] }
|
||||
pyo3-async-runtimes = { version = "0.28", features = [
|
||||
"attributes",
|
||||
@@ -43,7 +43,7 @@ libc = "0.2"
|
||||
[build-dependencies]
|
||||
pyo3-build-config = { version = "0.28", features = [
|
||||
"extension-module",
|
||||
"abi3-py39",
|
||||
"abi3-py310",
|
||||
] }
|
||||
|
||||
[features]
|
||||
|
||||
@@ -60,7 +60,7 @@ tests = [
|
||||
"pytest-asyncio>=0.21",
|
||||
"duckdb>=0.9.0",
|
||||
"pytz>=2023.3",
|
||||
"polars>=0.19, <=1.3.0",
|
||||
"polars>=0.19, <=1.32.3",
|
||||
"pyarrow<25",
|
||||
"pyarrow-stubs>=16.0",
|
||||
"pylance==9.0.0rc1",
|
||||
@@ -140,6 +140,7 @@ include = [
|
||||
"python/lancedb/remote/errors.py",
|
||||
"python/lancedb/embeddings/__init__.py",
|
||||
"python/lancedb/_lancedb.pyi",
|
||||
"python/type_tests/connect.py",
|
||||
]
|
||||
exclude = ["python/tests/"]
|
||||
pythonVersion = "3.13"
|
||||
|
||||
@@ -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__",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Private first-class Function namespace facades for database connections.
|
||||
|
||||
These helpers are internal submission and lookup surfaces. They are not durable
|
||||
resources and are not part of the public top-level export surface.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from . import _udf
|
||||
from ._lancedb import Function
|
||||
from .job import AsyncJob, Job
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .db import AsyncConnection, DBConnection
|
||||
|
||||
|
||||
class _SyncFunctions:
|
||||
"""Synchronous `db.functions` facade."""
|
||||
|
||||
__slots__ = ("_connection",)
|
||||
|
||||
def __init__(self, connection: DBConnection) -> None:
|
||||
self._connection = connection
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return "_SyncFunctions()"
|
||||
|
||||
def register(self, name: str, decorated_udf: Callable[..., object]) -> Job:
|
||||
"""Register a decorated UDF and return a synchronous [Job][lancedb.job.Job]."""
|
||||
definition = _udf._build_function_definition(decorated_udf)
|
||||
native_job = self._connection._submit_register_function(name, definition)
|
||||
return Job(AsyncJob(native_job))
|
||||
|
||||
def replace(
|
||||
self, name: str, current: Function, decorated_udf: Callable[..., object]
|
||||
) -> Job:
|
||||
"""Conditionally replace a Function; return sync [Job][lancedb.job.Job]."""
|
||||
definition = _udf._build_function_definition(decorated_udf)
|
||||
native_job = self._connection._submit_replace_function(
|
||||
name, current, definition
|
||||
)
|
||||
return Job(AsyncJob(native_job))
|
||||
|
||||
def get(self, name: str) -> Function:
|
||||
"""Return the Function currently bound to a database-scoped name."""
|
||||
return self._connection._lookup_function_by_name(name)
|
||||
|
||||
def get_by_id(self, function_id: str) -> Function:
|
||||
"""Return the immutable Function for an exact Function ID."""
|
||||
return self._connection._lookup_function_by_id(function_id)
|
||||
|
||||
def remove(self, name: str, current: Function) -> None:
|
||||
"""Conditionally remove a Function catalog name binding."""
|
||||
return self._connection._remove_function_name(name, current)
|
||||
|
||||
def revoke(self, function: Function) -> None:
|
||||
"""Revoke an exact immutable Function by administrator set-bit."""
|
||||
return self._connection._revoke_function(function)
|
||||
|
||||
|
||||
class _AsyncFunctions:
|
||||
"""Asynchronous `async_db.functions` facade."""
|
||||
|
||||
__slots__ = ("_connection",)
|
||||
|
||||
def __init__(self, connection: AsyncConnection) -> None:
|
||||
self._connection = connection
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return "_AsyncFunctions()"
|
||||
|
||||
async def register(
|
||||
self, name: str, decorated_udf: Callable[..., object]
|
||||
) -> AsyncJob:
|
||||
"""Register a decorated UDF and return an [AsyncJob][lancedb.job.AsyncJob]."""
|
||||
definition = _udf._build_function_definition(decorated_udf)
|
||||
native_job = await self._connection._register_function(name, definition)
|
||||
return AsyncJob(native_job)
|
||||
|
||||
async def replace(
|
||||
self, name: str, current: Function, decorated_udf: Callable[..., object]
|
||||
) -> AsyncJob:
|
||||
"""Conditionally replace a Function; return [AsyncJob][lancedb.job.AsyncJob]."""
|
||||
definition = _udf._build_function_definition(decorated_udf)
|
||||
native_job = await self._connection._replace_function(name, current, definition)
|
||||
return AsyncJob(native_job)
|
||||
|
||||
async def get(self, name: str) -> Function:
|
||||
"""Return the Function currently bound to a database-scoped name."""
|
||||
return await self._connection._lookup_function_by_name(name)
|
||||
|
||||
async def get_by_id(self, function_id: str) -> Function:
|
||||
"""Return the immutable Function for an exact Function ID."""
|
||||
return await self._connection._lookup_function_by_id(function_id)
|
||||
|
||||
async def remove(self, name: str, current: Function) -> None:
|
||||
"""Conditionally remove a Function catalog name binding."""
|
||||
return await self._connection._remove_function_name(name, current)
|
||||
|
||||
async def revoke(self, function: Function) -> None:
|
||||
"""Revoke an exact immutable Function by administrator set-bit."""
|
||||
return await self._connection._revoke_function(function)
|
||||
@@ -153,6 +153,16 @@ class Connection(object):
|
||||
async def job_history(
|
||||
self, job_id: Optional[str] = None
|
||||
) -> List[pa.RecordBatch]: ...
|
||||
async def _register_function(
|
||||
self, name: str, definition: "_FunctionDefinition"
|
||||
) -> Job: ...
|
||||
async def _replace_function(
|
||||
self, name: str, current: Function, definition: "_FunctionDefinition"
|
||||
) -> Job: ...
|
||||
async def _lookup_function_by_name(self, name: str) -> Function: ...
|
||||
async def _lookup_function_by_id(self, function_id: str) -> Function: ...
|
||||
async def _remove_function_name(self, name: str, current: Function) -> None: ...
|
||||
async def _revoke_function(self, function: Function) -> None: ...
|
||||
async def create_table(
|
||||
self,
|
||||
name: str,
|
||||
@@ -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]): ...
|
||||
@@ -355,6 +413,10 @@ class Table:
|
||||
async def set_lsm_write_spec(self, spec: LsmWriteSpec) -> None: ...
|
||||
async def unset_lsm_write_spec(self) -> None: ...
|
||||
async def get_lsm_write_spec(self) -> Optional[LsmWriteSpec]: ...
|
||||
async def checkpoint_lsm(self) -> None: ...
|
||||
async def flush_lsm(self) -> None: ...
|
||||
async def compact_lsm(self) -> None: ...
|
||||
async def get_lsm_stats(self, include_generation_rows: bool) -> Optional[dict]: ...
|
||||
async def close_lsm_writers(self) -> None: ...
|
||||
@property
|
||||
def tags(self) -> Tags: ...
|
||||
@@ -649,9 +711,10 @@ class LsmWriteSpec:
|
||||
def identity(column: str) -> "LsmWriteSpec": ...
|
||||
@staticmethod
|
||||
def unsharded() -> "LsmWriteSpec": ...
|
||||
def with_maintained_indexes(self, indexes: List[str]) -> "LsmWriteSpec":
|
||||
"""Return a copy of this spec asking the MemWAL to keep the named
|
||||
indexes up to date as rows are appended."""
|
||||
def with_maintained_indexes(self, indexes: Optional[List[str]]) -> "LsmWriteSpec":
|
||||
"""Set which indexes the MemWAL keeps up to date. None resolves every
|
||||
index on the table at install, failing if one cannot be maintained;
|
||||
a list is verbatim, empty means none."""
|
||||
...
|
||||
def with_writer_config_defaults(self, defaults: Dict[str, str]) -> "LsmWriteSpec":
|
||||
"""Return a copy of this spec recording the given default
|
||||
@@ -666,7 +729,9 @@ class LsmWriteSpec:
|
||||
@property
|
||||
def num_buckets(self) -> Optional[int]: ...
|
||||
@property
|
||||
def maintained_indexes(self) -> List[str]: ...
|
||||
def maintained_indexes(self) -> Optional[List[str]]:
|
||||
"""Indexes the MemWAL keeps up to date, or None for every supported one."""
|
||||
...
|
||||
@property
|
||||
def writer_config_defaults(self) -> Dict[str, str]: ...
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -101,8 +101,7 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
|
||||
|
||||
@weak_lru(maxsize=1)
|
||||
def ndims(self):
|
||||
model = self.get_model()
|
||||
return model.encode("foo").shape[0]
|
||||
return len(self.generate_embeddings([[self.source_instruction, "foo"]])[0])
|
||||
|
||||
def compute_query_embeddings(self, query: str, *args, **kwargs) -> List[np.array]:
|
||||
return self.generate_embeddings([[self.query_instruction, query]])
|
||||
|
||||
@@ -87,12 +87,13 @@ class JinaEmbeddings(EmbeddingFunction):
|
||||
if isinstance(image, bytes):
|
||||
image_dict = {"image": base64.b64encode(image).decode("utf-8")}
|
||||
elif isinstance(image, (str, Path)):
|
||||
parsed = urlparse.urlparse(image)
|
||||
# TODO handle drive letter on windows.
|
||||
parsed = urlparse(str(image))
|
||||
PIL_Image = attempt_import_or_raise("PIL.Image", "pillow")
|
||||
if parsed.scheme == "file":
|
||||
pil_image = PIL_Image.open(parsed.path)
|
||||
elif parsed.scheme == "":
|
||||
elif parsed.scheme == "" or (os.name == "nt" and len(parsed.scheme) == 1):
|
||||
# A Windows drive letter parses as a one-character scheme
|
||||
# ("C:\\img.png" -> scheme="c"), so treat it as a local path.
|
||||
pil_image = PIL_Image.open(image if os.name == "nt" else parsed.path)
|
||||
elif parsed.scheme.startswith("http"):
|
||||
pil_image = PIL_Image.open(io.BytesIO(url_retrieve(image)))
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
|
||||
"""Custom exception handling"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class MissingValueError(ValueError):
|
||||
"""Exception raised when a required value is missing."""
|
||||
@@ -26,12 +28,47 @@ class MissingColumnError(KeyError):
|
||||
|
||||
|
||||
class JobFailedError(RuntimeError):
|
||||
"""Exception raised when an asynchronous job reaches the failed state."""
|
||||
"""Exception raised when an asynchronous job reaches the failed state.
|
||||
|
||||
pass
|
||||
``error_code`` is the optional exact category string projected from the
|
||||
native job failure when the backend supplied one. The RuntimeError
|
||||
message remains the existing diagnostic text and must not be used to
|
||||
recover or override the code.
|
||||
"""
|
||||
|
||||
__slots__ = ("_error_code",)
|
||||
|
||||
def __init__(self, message: str, error_code: Optional[str] = None) -> None:
|
||||
super().__init__(message)
|
||||
self._error_code = error_code
|
||||
|
||||
@property
|
||||
def error_code(self) -> Optional[str]:
|
||||
"""Exact job failure error category string, when supplied."""
|
||||
return self._error_code
|
||||
|
||||
|
||||
class JobCancelledError(RuntimeError):
|
||||
"""Exception raised when an asynchronous job was cancelled."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class FunctionError(RuntimeError):
|
||||
"""Exception raised when a first-class Function operation fails.
|
||||
|
||||
``code`` is the stable semantic category from the native error. The
|
||||
message is a sanitized client diagnostic and must not be used to recover
|
||||
or override the code.
|
||||
"""
|
||||
|
||||
__slots__ = ("_code",)
|
||||
|
||||
def __init__(self, message: str, code: str) -> None:
|
||||
super().__init__(message)
|
||||
self._code = code
|
||||
|
||||
@property
|
||||
def code(self) -> str:
|
||||
"""Stable Function error category string."""
|
||||
return self._code
|
||||
|
||||
@@ -10,6 +10,7 @@ from typing import Optional
|
||||
from lancedb.background_loop import LOOP
|
||||
|
||||
from . import _lancedb
|
||||
from ._lancedb import Function
|
||||
|
||||
|
||||
class AsyncJob:
|
||||
@@ -44,18 +45,22 @@ class AsyncJob:
|
||||
return "finished"
|
||||
return await self._inner.status()
|
||||
|
||||
async def wait(self, timeout: Optional[timedelta] = None):
|
||||
async def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
|
||||
"""Wait until the operation reaches a terminal state.
|
||||
|
||||
Returns the success result when present (currently a
|
||||
:class:`~lancedb.Function`), or `None` when the job finished without
|
||||
one.
|
||||
|
||||
Raises `JobFailedError` if the operation failed, `JobCancelledError`
|
||||
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
|
||||
"""
|
||||
if self._inner is None:
|
||||
return
|
||||
return None
|
||||
if timeout is None:
|
||||
await self._inner.wait()
|
||||
return await self._inner.wait()
|
||||
else:
|
||||
await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
|
||||
return await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
|
||||
|
||||
async def cancel(self):
|
||||
"""Request cancellation. Cancelling a finished operation is a no-op."""
|
||||
@@ -88,15 +93,19 @@ class Job:
|
||||
return "finished"
|
||||
return LOOP.run(self._inner.status())
|
||||
|
||||
def wait(self, timeout: Optional[timedelta] = None):
|
||||
def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
|
||||
"""Block until the operation reaches a terminal state.
|
||||
|
||||
Returns the success result when present (currently a
|
||||
:class:`~lancedb.Function`), or `None` when the job finished without
|
||||
one.
|
||||
|
||||
Raises `JobFailedError` if the operation failed, `JobCancelledError`
|
||||
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
|
||||
"""
|
||||
if self._inner is None:
|
||||
return
|
||||
LOOP.run(self._inner.wait(timeout))
|
||||
return None
|
||||
return LOOP.run(self._inner.wait(timeout))
|
||||
|
||||
def cancel(self):
|
||||
"""Request cancellation. Cancelling a finished operation is a no-op."""
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -153,6 +153,16 @@ def Vector(
|
||||
return FixedSizeList
|
||||
|
||||
|
||||
def _raise_bare_vector_error(*_args):
|
||||
raise TypeError("Vector must be parameterized with a dimension, e.g. Vector(128).")
|
||||
|
||||
|
||||
# Pydantic v1 and v2 otherwise treat the bare Vector factory as a field validator
|
||||
# and inspect its signature, which produces misleading errors about internal types.
|
||||
setattr(Vector, "__get_validators__", _raise_bare_vector_error)
|
||||
setattr(Vector, "__get_pydantic_core_schema__", _raise_bare_vector_error)
|
||||
|
||||
|
||||
def MultiVector(
|
||||
dim: int, value_type: pa.DataType = pa.float32(), nullable: bool = True
|
||||
) -> Type:
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -108,6 +108,11 @@ def _should_push_down_query_table(
|
||||
return namespace_client is not None and "QueryTable" in pushdown_operations
|
||||
|
||||
|
||||
def _polars_predicate_pushdown_barrier(frame: Any) -> Any:
|
||||
"""Return a Polars frame unchanged while blocking predicate pushdown."""
|
||||
return frame
|
||||
|
||||
|
||||
_MODEL_BACKED_TOKENIZER_PREFIXES = ("jieba", "lindera")
|
||||
_MODEL_BACKED_TOKENIZER_ERRORS = (
|
||||
"unknown base tokenizer",
|
||||
@@ -180,6 +185,7 @@ if TYPE_CHECKING:
|
||||
LsmWriteSpec,
|
||||
MergeResult,
|
||||
UpdateResult,
|
||||
_FunctionCall,
|
||||
)
|
||||
from .index import IndexConfig
|
||||
import pandas
|
||||
@@ -864,12 +870,18 @@ class Table(ABC):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def to_polars(self, **kwargs) -> "pl.DataFrame":
|
||||
"""Return the table as a polars.DataFrame.
|
||||
def to_polars(self, **kwargs) -> "pl.LazyFrame":
|
||||
"""Return the table as a Polars LazyFrame.
|
||||
|
||||
Note
|
||||
----
|
||||
The Polars streaming engine is not supported because it does not currently
|
||||
implement Python PyArrow dataset scans. Use the default engine when collecting
|
||||
this LazyFrame.
|
||||
|
||||
Returns
|
||||
-------
|
||||
polars.DataFrame
|
||||
polars.LazyFrame
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -996,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.
|
||||
@@ -2569,6 +2618,9 @@ class LanceTable(Table):
|
||||
2. Currently we've disabled push-down of the filters from polars
|
||||
because polars pushdown into pyarrow uses pyarrow compute
|
||||
expressions rather than SQl strings (which LanceDB supports)
|
||||
3. The Polars streaming engine is not supported because it does not
|
||||
currently implement Python PyArrow dataset scans. Use the default
|
||||
engine when collecting this LazyFrame.
|
||||
|
||||
Returns
|
||||
-------
|
||||
@@ -2577,8 +2629,12 @@ class LanceTable(Table):
|
||||
from lancedb.integrations.pyarrow import PyarrowDatasetAdapter
|
||||
|
||||
dataset = PyarrowDatasetAdapter(self)
|
||||
return pl.scan_pyarrow_dataset(
|
||||
dataset, allow_pyarrow_filter=False, batch_size=batch_size
|
||||
# Polars 1.32's non-PyArrow callback path passes batch_size twice. Keep
|
||||
# the compatible PyArrow path, but block predicates because this adapter
|
||||
# cannot translate PyArrow expressions into LanceDB filters.
|
||||
return pl.scan_pyarrow_dataset(dataset, batch_size=batch_size).map_batches(
|
||||
_polars_predicate_pushdown_barrier,
|
||||
predicate_pushdown=False,
|
||||
)
|
||||
|
||||
# New unified API overload
|
||||
@@ -2831,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,
|
||||
@@ -3958,6 +4051,28 @@ class LanceTable(Table):
|
||||
[`AsyncTable.get_lsm_write_spec`][lancedb.AsyncTable.get_lsm_write_spec]."""
|
||||
return LOOP.run(self._table.get_lsm_write_spec())
|
||||
|
||||
def checkpoint_lsm(self) -> None:
|
||||
"""Synchronous version of
|
||||
[`AsyncTable.checkpoint_lsm`][lancedb.AsyncTable.checkpoint_lsm]."""
|
||||
return LOOP.run(self._table.checkpoint_lsm())
|
||||
|
||||
def flush_lsm(self) -> None:
|
||||
"""Synchronous version of
|
||||
[`AsyncTable.flush_lsm`][lancedb.AsyncTable.flush_lsm]."""
|
||||
return LOOP.run(self._table.flush_lsm())
|
||||
|
||||
def compact_lsm(self) -> None:
|
||||
"""Synchronous version of
|
||||
[`AsyncTable.compact_lsm`][lancedb.AsyncTable.compact_lsm]."""
|
||||
return LOOP.run(self._table.compact_lsm())
|
||||
|
||||
def get_lsm_stats(self, *, include_generation_rows: bool = False) -> Optional[dict]:
|
||||
"""Synchronous version of
|
||||
[`AsyncTable.get_lsm_stats`][lancedb.AsyncTable.get_lsm_stats]."""
|
||||
return LOOP.run(
|
||||
self._table.get_lsm_stats(include_generation_rows=include_generation_rows)
|
||||
)
|
||||
|
||||
def close_lsm_writers(self) -> None:
|
||||
"""Close cached MemWAL shard writers. See
|
||||
[`AsyncTable.close_lsm_writers`][lancedb.AsyncTable.close_lsm_writers]."""
|
||||
@@ -4636,6 +4751,13 @@ class AsyncTable:
|
||||
via [`set_unenforced_primary_key`]; bucket sharding additionally
|
||||
requires it to be the single column being bucketed.
|
||||
|
||||
By default the MemWAL maintains every index on the table, resolved
|
||||
here — a snapshot, so an index created afterwards needs the spec unset
|
||||
and set again. This fails if one cannot be maintained; name the set
|
||||
with ``with_maintained_indexes`` to install anyway. That pins an exact
|
||||
set (a still-building index is rejected, not omitted); ``[]`` maintains
|
||||
none.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
spec : LsmWriteSpec
|
||||
@@ -4662,12 +4784,73 @@ class AsyncTable:
|
||||
|
||||
Returns ``None`` when the MemWAL LSM write path is not enabled (no
|
||||
spec has been set, or it was removed with `unset_lsm_write_spec`).
|
||||
The returned spec — including its ``maintained_indexes`` and
|
||||
``writer_config_defaults`` — mirrors what was passed to
|
||||
`set_lsm_write_spec`.
|
||||
The returned spec mirrors what was passed to `set_lsm_write_spec`,
|
||||
except that ``maintained_indexes`` always reports the concrete list
|
||||
resolved when the spec was set — ``None`` never round-trips.
|
||||
"""
|
||||
return await self._inner.get_lsm_write_spec()
|
||||
|
||||
async def checkpoint_lsm(self) -> None:
|
||||
"""Converge this table's LSM write path into its base table.
|
||||
|
||||
One flush, sealing every memtable into L0, then compaction triggers
|
||||
until every generation that existed at that moment has reached base.
|
||||
The loop runs client-side, reading progress from ``get_lsm_stats``.
|
||||
|
||||
Best-effort: generations created *while* it runs are deliberately not
|
||||
waited on, which is what lets it terminate on a table taking writes.
|
||||
Idempotent and safe on a cadence.
|
||||
|
||||
There is no deadline, and the caller owns that. It returns when the
|
||||
target generations are gone, raises on a terminal server fault, and
|
||||
otherwise waits however long the server takes. A slow table and a
|
||||
stuck one are the same picture from the client: the compactor pool is
|
||||
shared across every table on the node, so a checkpoint queued behind
|
||||
unrelated work looks exactly like one that is merging. Wrap this in
|
||||
``asyncio.wait_for`` for a wall-clock bound; abandoning it partway
|
||||
costs nothing.
|
||||
"""
|
||||
return await self._inner.checkpoint_lsm()
|
||||
|
||||
async def flush_lsm(self) -> None:
|
||||
"""Seal every bucket's active memtable into L0.
|
||||
|
||||
Does not touch the base table — moving L0 into base is
|
||||
`compact_lsm`. On a node that has not claimed this table, this claims
|
||||
it and replays its WAL log first.
|
||||
"""
|
||||
return await self._inner.flush_lsm()
|
||||
|
||||
async def compact_lsm(self) -> None:
|
||||
"""Trigger a background L0 to base compaction pass per bucket.
|
||||
|
||||
Returns once the passes are dispatched, not once they finish: watch
|
||||
``get_lsm_stats`` for progress, or use ``checkpoint_lsm`` to loop
|
||||
until the current L0 has reached base.
|
||||
"""
|
||||
return await self._inner.compact_lsm()
|
||||
|
||||
async def get_lsm_stats(
|
||||
self, *, include_generation_rows: bool = False
|
||||
) -> Optional[dict]:
|
||||
"""Read live per-bucket LSM state.
|
||||
|
||||
Answers "how far behind is my fresh tier", "which bucket is hot", and
|
||||
"why is my fresh-tier vector search brute-force". Mutates no table
|
||||
state, though on a node that has not claimed this table it claims it,
|
||||
exactly as a read would.
|
||||
|
||||
Returns ``None`` only when the LSM write path is not enabled.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
include_generation_rows
|
||||
Report a row count per L0 generation. Off by default: each count
|
||||
opens an uncached Lance dataset, and ``checkpoint_lsm`` polls this
|
||||
needing only generation numbers.
|
||||
"""
|
||||
return await self._inner.get_lsm_stats(include_generation_rows)
|
||||
|
||||
async def close_lsm_writers(self) -> None:
|
||||
"""Drain and close any cached MemWAL shard writers for this table.
|
||||
|
||||
@@ -4958,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.
|
||||
@@ -6233,7 +6460,9 @@ class TableStatistics:
|
||||
Attributes
|
||||
----------
|
||||
total_bytes: int
|
||||
The total number of bytes in the table.
|
||||
The total size, in bytes, of the table's data files, index files, and
|
||||
overlay files. Read from the manifest, so this excludes deletion files
|
||||
and manifests.
|
||||
num_rows: int
|
||||
The total number of rows in the table.
|
||||
num_indices: int
|
||||
|
||||
@@ -395,6 +395,11 @@ def _(value: dict):
|
||||
)
|
||||
|
||||
|
||||
@value_to_sql.register(pa.Scalar)
|
||||
def _(value: pa.Scalar):
|
||||
return value_to_sql(value.as_py())
|
||||
|
||||
|
||||
@value_to_sql.register(np.ndarray)
|
||||
def _(value: np.ndarray):
|
||||
return value_to_sql(value.tolist())
|
||||
|
||||
@@ -2,9 +2,11 @@
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
|
||||
import inspect
|
||||
import re
|
||||
import sys
|
||||
from datetime import timedelta
|
||||
from importlib import resources
|
||||
import os
|
||||
from types import SimpleNamespace
|
||||
|
||||
@@ -17,6 +19,10 @@ from lance_namespace.errors import NamespaceNotEmptyError, TableNotFoundError
|
||||
from lancedb.pydantic import LanceModel, Vector
|
||||
|
||||
|
||||
def test_package_includes_pep_561_marker():
|
||||
assert resources.files(lancedb).joinpath("py.typed").is_file()
|
||||
|
||||
|
||||
def test_basic(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
|
||||
@@ -62,17 +68,23 @@ def test_basic(tmp_path):
|
||||
assert db.open_table("test").name == db["test"].name
|
||||
|
||||
|
||||
def test_sync_repr_does_not_use_background_loop(tmp_path, monkeypatch):
|
||||
def test_sync_debugger_inspection_does_not_use_background_loop(tmp_path, monkeypatch):
|
||||
from lancedb.background_loop import LOOP
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("test", data=[{"id": 1}])
|
||||
|
||||
def fail_run(*args, **kwargs):
|
||||
raise AssertionError("repr should not use the Python background loop")
|
||||
raise AssertionError("debugger inspection should not use the background loop")
|
||||
|
||||
monkeypatch.setattr(LOOP, "run", fail_run)
|
||||
|
||||
# Debuggers enumerate and evaluate every exposed attribute when expanding a
|
||||
# variable. This must remain safe while their breakpoint suspends LOOP's thread.
|
||||
members = dict(inspect.getmembers(db))
|
||||
|
||||
assert members["uri"] == str(tmp_path)
|
||||
assert members["read_consistency_interval"] is None
|
||||
assert repr(db) == f"LanceDBConnection(uri={str(tmp_path)!r})"
|
||||
assert repr(table) == f"LanceTable(name='test', _conn={db!r})"
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -64,6 +64,23 @@ def test_embedding_function(tmp_path):
|
||||
assert np.allclose(actual, expected)
|
||||
|
||||
|
||||
def test_instructor_ndims_uses_instruction():
|
||||
instructor = get_registry().get("instructor").create()
|
||||
model = MagicMock()
|
||||
model.encode.return_value = np.zeros((1, 384))
|
||||
|
||||
with patch.object(type(instructor), "get_model", return_value=model):
|
||||
assert instructor.ndims() == 384
|
||||
|
||||
model.encode.assert_called_once_with(
|
||||
[[instructor.source_instruction, "foo"]],
|
||||
batch_size=instructor.batch_size,
|
||||
show_progress_bar=instructor.show_progress_bar,
|
||||
normalize_embeddings=instructor.normalize_embeddings,
|
||||
device=instructor.device,
|
||||
)
|
||||
|
||||
|
||||
def test_embedding_function_variables():
|
||||
@register("variable-testing")
|
||||
class VariableTestingFunction(TextEmbeddingFunction):
|
||||
@@ -115,34 +132,16 @@ def test_embedding_function_variables():
|
||||
assert func.safe_model_dump()["secret_key"] == "$var:secret"
|
||||
|
||||
|
||||
def test_parse_functions_with_variables():
|
||||
@register("variable-parsing-test")
|
||||
class VariableParsingFunction(TextEmbeddingFunction):
|
||||
api_key: str
|
||||
base_url: Optional[str] = None
|
||||
|
||||
@staticmethod
|
||||
def sensitive_keys():
|
||||
return ["api_key"]
|
||||
|
||||
def ndims(self):
|
||||
return 10
|
||||
|
||||
def generate_embeddings(self, texts):
|
||||
# Mock implementation that just returns random embeddings
|
||||
# In real usage, this would use the api_key to call an API
|
||||
return [np.random.rand(self.ndims()).tolist() for _ in texts]
|
||||
|
||||
def test_openai_variables_survive_metadata_round_trip():
|
||||
registry = EmbeddingFunctionRegistry.get_instance()
|
||||
|
||||
registry.set_var("test_api_key", "sk-test-key-12345")
|
||||
registry.set_var("test_base_url", "https://api.example.com")
|
||||
|
||||
conf = EmbeddingFunctionConfig(
|
||||
source_column="text",
|
||||
vector_column="vector",
|
||||
function=registry.get("variable-parsing-test").create(
|
||||
api_key="$var:test_api_key", base_url="$var:test_base_url"
|
||||
function=registry.get("openai").create(
|
||||
api_key="$var:test_api_key", base_url="https://api.example.com"
|
||||
),
|
||||
)
|
||||
|
||||
@@ -150,7 +149,10 @@ def test_parse_functions_with_variables():
|
||||
|
||||
# Create a mock arrow table with the metadata
|
||||
schema = pa.schema(
|
||||
[pa.field("text", pa.string()), pa.field("vector", pa.list_(pa.float32(), 10))]
|
||||
[
|
||||
pa.field("text", pa.string()),
|
||||
pa.field("vector", pa.list_(pa.float32(), 1536)),
|
||||
]
|
||||
)
|
||||
table = pa.table({"text": [], "vector": []}, schema=schema)
|
||||
table = table.replace_schema_metadata(metadata)
|
||||
@@ -164,13 +166,15 @@ def test_parse_functions_with_variables():
|
||||
|
||||
assert parsed_func.api_key == "sk-test-key-12345"
|
||||
assert parsed_func.base_url == "https://api.example.com"
|
||||
|
||||
embeddings = parsed_func.generate_embeddings(["test text"])
|
||||
assert len(embeddings) == 1
|
||||
assert len(embeddings[0]) == 10
|
||||
|
||||
assert parsed_func.safe_model_dump()["api_key"] == "$var:test_api_key"
|
||||
|
||||
with patch("lancedb.embeddings.openai.attempt_import_or_raise") as import_openai:
|
||||
parsed_func._openai_client
|
||||
|
||||
import_openai.return_value.OpenAI.assert_called_once_with(
|
||||
api_key="sk-test-key-12345", base_url="https://api.example.com"
|
||||
)
|
||||
|
||||
|
||||
def test_embedding_with_bad_results(tmp_path):
|
||||
@register("null-embedding")
|
||||
@@ -627,3 +631,23 @@ def test_url_retrieve_downloads_image():
|
||||
image_bytes = url_retrieve(image_url)
|
||||
img = Image.open(io.BytesIO(image_bytes))
|
||||
assert img.size[0] > 0 and img.size[1] > 0
|
||||
|
||||
|
||||
def test_jina_generate_image_input_dict_local_path(tmp_path):
|
||||
"""
|
||||
JinaEmbeddings._generate_image_input_dict must accept a local image path
|
||||
(str or Path), not just bytes. Previously it crashed with
|
||||
`AttributeError: 'function' object has no attribute 'urlparse'` on any
|
||||
str/Path input because it called `urlparse.urlparse(image)` instead of
|
||||
`urlparse(image)` (urlparse was imported as a function, not a module).
|
||||
"""
|
||||
Image = pytest.importorskip("PIL.Image")
|
||||
from lancedb.embeddings.jinaai import JinaEmbeddings
|
||||
|
||||
image_path = tmp_path / "test.png"
|
||||
Image.new("RGB", (4, 4), color="red").save(image_path, format="PNG")
|
||||
|
||||
for image in (str(image_path), image_path):
|
||||
image_dict = JinaEmbeddings._generate_image_input_dict(image)
|
||||
assert "image" in image_dict
|
||||
assert isinstance(image_dict["image"], str) and len(image_dict["image"]) > 0
|
||||
|
||||
@@ -0,0 +1,372 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python exact Function handle call authoring (FF-028)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.expr import Expr, col, func, lit
|
||||
|
||||
_CALL_PATH = "/v1/functions/lookup"
|
||||
_CALL_CATALOG_NAME = "text.normalize.call-name"
|
||||
_CALL_FUNCTION_ID = "fn.exact.call-handle"
|
||||
_LITERAL_PAYLOAD_SENTINEL = "LITERAL_PAYLOAD_SENTINEL_call_xyz_42"
|
||||
_INT_PAYLOAD_SENTINEL = 2_147_000_123
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64).
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
_LIST_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////+4AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAABAAAANz///8c"
|
||||
"AAAADAAAAAAAAQxcAAAAAQAAABwAAAAEAAQABAAAABAAFAAQAA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECH"
|
||||
"AAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAQAAABpdGVtAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////8AAAAAFAAAAAAAAAAMABQAEgAMAAgABAAMAAAAnAAAAKAAAAAQAAAAAAAEAAgACAAAAAQACAAAAAQAAAA"
|
||||
"BAAAABAAAANz///8cAAAADAAAAAAAAQxcAAAAAQAAABwAAAAEAAQABAAAABAAFAAQAA4ADwAEAAAACAAQAAAA"
|
||||
"GAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAQAAABpdGVtAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAAAwAAAAEFSUk9XMQ=="
|
||||
)
|
||||
|
||||
_OVERDESIGN_ATTRS = (
|
||||
"id",
|
||||
"function_id",
|
||||
"name",
|
||||
"connection",
|
||||
"table",
|
||||
"snapshot",
|
||||
"field_id",
|
||||
"field_ids",
|
||||
"job",
|
||||
"job_id",
|
||||
"artifact",
|
||||
"digest",
|
||||
"retry_key",
|
||||
"idempotency_key",
|
||||
"user_version",
|
||||
"execute",
|
||||
"status",
|
||||
"wait",
|
||||
"cancel",
|
||||
"to_json",
|
||||
"_to_json",
|
||||
"serialize",
|
||||
"geneva",
|
||||
)
|
||||
|
||||
|
||||
def _sample_function_wire(
|
||||
*,
|
||||
function_id: str = _CALL_FUNCTION_ID,
|
||||
parameters: list[dict[str, str]] | None = None,
|
||||
output_type_ipc: str = _UTF8_TYPE_IPC_B64,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": function_id,
|
||||
"signature": {
|
||||
"parameters": parameters
|
||||
or [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": output_type_ipc,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _lookup_success_body(function: dict[str, Any] | None = None) -> bytes:
|
||||
return json.dumps({"function": function or _sample_function_wire()}).encode("utf-8")
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _lookup_function(function: dict[str, Any] | None = None):
|
||||
body = _lookup_success_body(function)
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _CALL_PATH
|
||||
_read_body(request)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(body)
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
return db.functions.get(_CALL_CATALOG_NAME)
|
||||
|
||||
|
||||
def _authored_call_type():
|
||||
cls = getattr(_native, "_FunctionCall", None)
|
||||
if cls is None:
|
||||
pytest.fail("lancedb._lancedb._FunctionCall is missing")
|
||||
return cls
|
||||
|
||||
|
||||
def _exception_text(exc: BaseException) -> str:
|
||||
parts = [str(exc), repr(exc)]
|
||||
current: BaseException | None = exc
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(f"{type(current).__name__}: {current}")
|
||||
current = current.__cause__ or current.__context__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def test_function_keyword_call_returns_private_frozen_authored_value():
|
||||
function = _lookup_function()
|
||||
assert callable(function)
|
||||
|
||||
authored = function(text=col("text"), limit=8)
|
||||
authored_type = _authored_call_type()
|
||||
assert type(authored) is authored_type
|
||||
assert authored_type.__module__ == "lancedb._lancedb"
|
||||
assert authored_type.__name__ == "_FunctionCall"
|
||||
|
||||
# Keyword order must not matter; bindings store/render in signature order.
|
||||
authored_reversed = function(limit=8, text=col("text"))
|
||||
assert type(authored_reversed) is authored_type
|
||||
rendered = repr(authored_reversed)
|
||||
assert rendered.index("text=") < rendered.index("limit=")
|
||||
assert 'text=field("text")' in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
|
||||
|
||||
def test_function_call_rejects_positional_missing_and_unknown_args():
|
||||
function = _lookup_function()
|
||||
|
||||
with pytest.raises(TypeError, match="keyword"):
|
||||
function(col("text"), 8)
|
||||
|
||||
with pytest.raises((TypeError, ValueError), match="limit"):
|
||||
function(text=col("text"))
|
||||
|
||||
with pytest.raises((TypeError, ValueError), match="text"):
|
||||
function(limit=8)
|
||||
|
||||
with pytest.raises((TypeError, ValueError), match="unknown|extra"):
|
||||
function(text=col("text"), limit=8, extra=1)
|
||||
|
||||
|
||||
def test_function_call_accepts_direct_case_sensitive_column_and_rejects_complex_exprs():
|
||||
function = _lookup_function()
|
||||
|
||||
authored = function(text=col("firstName"), limit=1)
|
||||
assert type(authored) is _authored_call_type()
|
||||
rendered = repr(authored)
|
||||
assert 'text=field("firstName")' in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
|
||||
complex_exprs = (
|
||||
col("text") + lit("x"),
|
||||
col("text").cast(pa.string()),
|
||||
func("lower", col("text")),
|
||||
col("text") == lit("x"),
|
||||
col("text").lower(),
|
||||
)
|
||||
for expr in complex_exprs:
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
function(text=expr, limit=1)
|
||||
|
||||
# Raw native PyExpr is not the public col() wrapper.
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
function(text=col("text")._inner, limit=1)
|
||||
|
||||
# Non-expression / non-literal objects are rejected for field-shaped misuse
|
||||
# when a column binding is required; plain strings are literals for utf8.
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
function(text=object(), limit=1)
|
||||
|
||||
|
||||
def test_function_call_plain_literal_declared_type_null_and_nested():
|
||||
function = _lookup_function()
|
||||
|
||||
authored = function(text="hello", limit=7)
|
||||
assert type(authored) is _authored_call_type()
|
||||
rendered = repr(authored)
|
||||
assert "text=literal(Utf8, null=false)" in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
|
||||
# Plain Python int normalizes to declared Int32 and non-null.
|
||||
authored_int32 = function(text="hello", limit=2_147_483_647)
|
||||
assert type(authored_int32) is _authored_call_type()
|
||||
rendered_int32 = repr(authored_int32)
|
||||
assert "limit=literal(Int32, null=false)" in rendered_int32
|
||||
assert "Int64" not in rendered_int32
|
||||
assert "2147483647" not in rendered_int32
|
||||
|
||||
# Plain None keeps each declared parameter type with null=true.
|
||||
authored_null = function(text=None, limit=None)
|
||||
assert type(authored_null) is _authored_call_type()
|
||||
rendered_null = repr(authored_null)
|
||||
assert "text=literal(Utf8, null=true)" in rendered_null
|
||||
assert "limit=literal(Int32, null=true)" in rendered_null
|
||||
|
||||
list_function = _lookup_function(
|
||||
_sample_function_wire(
|
||||
parameters=[
|
||||
{"name": "values", "data_type_ipc": _LIST_INT32_TYPE_IPC_B64},
|
||||
]
|
||||
)
|
||||
)
|
||||
authored_list = list_function(values=[1, 2, 3])
|
||||
assert type(authored_list) is _authored_call_type()
|
||||
rendered_list = repr(authored_list)
|
||||
assert "values=literal(List(Int32), null=false)" in rendered_list
|
||||
assert "[1, 2, 3]" not in rendered_list
|
||||
|
||||
authored_list_null = list_function(values=None)
|
||||
assert type(authored_list_null) is _authored_call_type()
|
||||
rendered_list_null = repr(authored_list_null)
|
||||
assert "values=literal(List(Int32), null=true)" in rendered_list_null
|
||||
|
||||
|
||||
def test_function_call_direct_literal_expr_exact_type_only():
|
||||
function = _lookup_function()
|
||||
|
||||
# lit(int) is Int64 in the expression builder; int32 parameter must reject it.
|
||||
with pytest.raises((TypeError, ValueError), match="limit|int32|type") as raised:
|
||||
function(text="hello", limit=lit(8))
|
||||
reject_text = _exception_text(raised.value)
|
||||
assert "Int64" in reject_text or "int64" in reject_text.lower()
|
||||
assert "Int32" in reject_text or "int32" in reject_text.lower()
|
||||
|
||||
# Exact utf8 literal expression is accepted and stored as Utf8/non-null.
|
||||
authored = function(text=lit("hello"), limit=8)
|
||||
assert type(authored) is _authored_call_type()
|
||||
rendered = repr(authored)
|
||||
assert "text=literal(Utf8, null=false)" in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
assert "hello" not in rendered
|
||||
|
||||
# Cast / arithmetic around a literal is not a direct Literal node.
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
function(text=lit("hello").cast(pa.string()), limit=8)
|
||||
|
||||
|
||||
def test_function_call_conversion_error_and_repr_are_payload_free():
|
||||
function = _lookup_function()
|
||||
|
||||
with pytest.raises((TypeError, ValueError)) as raised:
|
||||
function(text="ok", limit=_LITERAL_PAYLOAD_SENTINEL)
|
||||
text = _exception_text(raised.value)
|
||||
assert _LITERAL_PAYLOAD_SENTINEL not in text
|
||||
assert "limit" in text
|
||||
assert "int32" in text.lower() or "Int32" in text
|
||||
|
||||
authored = function(text=_LITERAL_PAYLOAD_SENTINEL, limit=_INT_PAYLOAD_SENTINEL)
|
||||
rendered = f"{authored!r}\n{authored!s}"
|
||||
assert _LITERAL_PAYLOAD_SENTINEL not in rendered
|
||||
assert str(_INT_PAYLOAD_SENTINEL) not in rendered
|
||||
assert "text=literal(Utf8, null=false)" in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
assert type(authored).__name__ == "_FunctionCall"
|
||||
assert "_FunctionCall" in rendered
|
||||
|
||||
|
||||
def test_function_call_private_type_nonconstructible_immutable_and_not_exported():
|
||||
function = _lookup_function()
|
||||
authored = function(text=col("text"), limit=1)
|
||||
authored_type = _authored_call_type()
|
||||
|
||||
assert "_FunctionCall" not in getattr(lancedb, "__all__", [])
|
||||
assert not hasattr(lancedb, "_FunctionCall")
|
||||
assert getattr(_native, "_FunctionCall", None) is authored_type
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
authored_type()
|
||||
|
||||
for attr in _OVERDESIGN_ATTRS:
|
||||
assert not hasattr(authored, attr)
|
||||
|
||||
for attr in ("function", "bindings", "arguments", "parameters", "text", "limit"):
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(authored, attr, None)
|
||||
|
||||
# Existing Function handle stays frozen / connection-free / name-free.
|
||||
assert not hasattr(function, "name")
|
||||
assert not hasattr(function, "connection")
|
||||
with pytest.raises(AttributeError):
|
||||
function.id = "mutated"
|
||||
|
||||
|
||||
def test_function_call_does_not_change_col_query_expression_behavior():
|
||||
# Regression guard: authoring must not alter public col()/Expr query behavior.
|
||||
expr = col("firstName") > lit(1)
|
||||
assert isinstance(expr, Expr)
|
||||
assert expr.to_sql() == "(`firstName` > 1)"
|
||||
@@ -0,0 +1,181 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pyarrow as pa
|
||||
|
||||
from lancedb import udf
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"value": pa.int64()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pyarrow==24.0.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
def double_nullable(value):
|
||||
if value is None:
|
||||
return None
|
||||
return value * 2
|
||||
|
||||
|
||||
def test_first_class_function_enterprise_lifecycle():
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
from datetime import timedelta
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import FunctionError
|
||||
from lancedb.expr import col
|
||||
|
||||
host = os.environ.get("LANCEDB_FCF_E2E_HOST")
|
||||
if not host:
|
||||
pytest.skip("LANCEDB_FCF_E2E_HOST is required for the live enterprise test")
|
||||
|
||||
database_uri = os.environ.get("LANCEDB_FCF_E2E_DB_URI", "db://fcf-e2e-local")
|
||||
api_key = os.environ.get("LANCEDB_FCF_E2E_API_KEY", "fake")
|
||||
run_suffix = uuid.uuid4().hex[:12]
|
||||
table_name = f"fcf_e2e_{run_suffix}"
|
||||
function_name = f"fcf_e2e.double_{run_suffix}"
|
||||
job_timeout = timedelta(minutes=5)
|
||||
query_timeout = timedelta(seconds=30)
|
||||
|
||||
def connect():
|
||||
return lancedb.connect(
|
||||
database_uri,
|
||||
api_key=api_key,
|
||||
host_override=host,
|
||||
)
|
||||
|
||||
setup_db = connect()
|
||||
setup_db.create_table(
|
||||
table_name,
|
||||
data=pa.Table.from_pylist(
|
||||
[
|
||||
{"row_id": 1, "value": 2},
|
||||
{"row_id": 2, "value": 5},
|
||||
{"row_id": 3, "value": None},
|
||||
],
|
||||
schema=pa.schema(
|
||||
[
|
||||
pa.field("row_id", pa.int64(), nullable=False),
|
||||
pa.field("value", pa.int64(), nullable=True),
|
||||
]
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
registration_job = setup_db.functions.register(function_name, double_nullable)
|
||||
registration_job_id = registration_job.id
|
||||
assert isinstance(registration_job_id, str) and registration_job_id
|
||||
registered_function = registration_job.wait(timeout=job_timeout)
|
||||
assert type(registered_function) is lancedb.Function
|
||||
assert isinstance(registered_function.id, str) and registered_function.id
|
||||
with pytest.raises(AttributeError):
|
||||
registered_function.id = "mutated"
|
||||
|
||||
catalog_reader = connect()
|
||||
function_by_name = catalog_reader.functions.get(function_name)
|
||||
function_by_id = catalog_reader.functions.get_by_id(registered_function.id)
|
||||
expected_signature = ((("value", pa.int64()),), pa.int64(), True)
|
||||
expected_identity = (
|
||||
registered_function.id,
|
||||
*expected_signature,
|
||||
)
|
||||
for function in (registered_function, function_by_name, function_by_id):
|
||||
assert type(function) is lancedb.Function
|
||||
assert (
|
||||
function.id,
|
||||
function.parameters,
|
||||
function.output_type,
|
||||
function.output_nullable,
|
||||
) == expected_identity
|
||||
|
||||
generated_column_table = catalog_reader.open_table(table_name)
|
||||
generated_column_job = generated_column_table.add_generated_column(
|
||||
"derived",
|
||||
registered_function(value=col("value")),
|
||||
)
|
||||
generated_column_job_id = generated_column_job.id
|
||||
assert isinstance(generated_column_job_id, str) and generated_column_job_id
|
||||
assert generated_column_job.wait(timeout=job_timeout) is None
|
||||
|
||||
complete_reader = connect().open_table(table_name)
|
||||
complete_status = complete_reader.generated_column_status("derived")
|
||||
assert complete_status == "complete"
|
||||
initial_rows = sorted(
|
||||
complete_reader.search()
|
||||
.select(["row_id", "value", "derived"])
|
||||
.limit(3)
|
||||
.to_list(timeout=query_timeout),
|
||||
key=lambda row: row["row_id"],
|
||||
)
|
||||
assert initial_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 5, "derived": 10},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
update_result = complete_reader.update(
|
||||
where="row_id = 2",
|
||||
values={"value": 7},
|
||||
)
|
||||
assert update_result.rows_updated == 1
|
||||
|
||||
incomplete_reader = connect().open_table(table_name)
|
||||
incomplete_status = incomplete_reader.generated_column_status("derived")
|
||||
assert incomplete_status == "incomplete"
|
||||
with pytest.raises(FunctionError) as raised:
|
||||
(
|
||||
incomplete_reader.search()
|
||||
.select(["row_id", "derived"])
|
||||
.limit(3)
|
||||
.to_list(timeout=query_timeout)
|
||||
)
|
||||
assert raised.value.code == "generated_column_incomplete"
|
||||
|
||||
refresh_job = incomplete_reader.refresh_generated_column("derived")
|
||||
refresh_job_id = refresh_job.id
|
||||
assert isinstance(refresh_job_id, str) and refresh_job_id
|
||||
assert refresh_job.wait(timeout=job_timeout) is None
|
||||
|
||||
refreshed_reader = connect().open_table(table_name)
|
||||
refreshed_status = refreshed_reader.generated_column_status("derived")
|
||||
assert refreshed_status == "complete"
|
||||
final_rows = sorted(
|
||||
refreshed_reader.search()
|
||||
.select(["row_id", "value", "derived"])
|
||||
.limit(3)
|
||||
.to_list(timeout=query_timeout),
|
||||
key=lambda row: row["row_id"],
|
||||
)
|
||||
assert final_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 7, "derived": 14},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
evidence = {
|
||||
"run_suffix": run_suffix,
|
||||
"database": database_uri.removeprefix("db://"),
|
||||
"table": table_name,
|
||||
"function": function_name,
|
||||
"function_id": registered_function.id,
|
||||
"job_ids": {
|
||||
"register": registration_job_id,
|
||||
"add_generated_column": generated_column_job_id,
|
||||
"refresh_generated_column": refresh_job_id,
|
||||
},
|
||||
"status_transitions": [
|
||||
complete_status,
|
||||
incomplete_status,
|
||||
refreshed_status,
|
||||
],
|
||||
"final_rows": final_rows,
|
||||
}
|
||||
print(json.dumps(evidence, sort_keys=True, separators=(",", ":")))
|
||||
@@ -0,0 +1,595 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pyarrow as pa
|
||||
|
||||
from lancedb import udf
|
||||
|
||||
|
||||
_RUNNING_DEADLINE_SECONDS = 30
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"value": pa.int64()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pyarrow==24.0.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
def reliable_double(value):
|
||||
if value is None:
|
||||
return None
|
||||
return value * 2
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"value": pa.int64()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pyarrow==24.0.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
def terminate_worker_on_input(value):
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
if len(value) == 0:
|
||||
return value
|
||||
except TypeError:
|
||||
pass
|
||||
|
||||
import os
|
||||
|
||||
os._exit(73)
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"value": pa.int64()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pyarrow==24.0.0"],
|
||||
output_nullable=False,
|
||||
)
|
||||
def slow_triple(value):
|
||||
import time
|
||||
|
||||
time.sleep(0.02)
|
||||
return value * 3
|
||||
|
||||
|
||||
def _require_live() -> str:
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
host = os.environ.get("LANCEDB_FCF_E2E_HOST")
|
||||
if not host:
|
||||
pytest.skip(
|
||||
"LANCEDB_FCF_E2E_HOST is required for live enterprise reliability tests"
|
||||
)
|
||||
return host
|
||||
|
||||
|
||||
def _job_timeout():
|
||||
from datetime import timedelta
|
||||
|
||||
return timedelta(minutes=5)
|
||||
|
||||
|
||||
def _query_timeout():
|
||||
from datetime import timedelta
|
||||
|
||||
return timedelta(seconds=30)
|
||||
|
||||
|
||||
def _connect():
|
||||
import os
|
||||
|
||||
import lancedb
|
||||
|
||||
return lancedb.connect(
|
||||
os.environ.get("LANCEDB_FCF_E2E_DB_URI", "db://fcf-e2e-local"),
|
||||
api_key=os.environ.get("LANCEDB_FCF_E2E_API_KEY", "fake"),
|
||||
host_override=_require_live(),
|
||||
)
|
||||
|
||||
|
||||
def _run_names(case: str) -> tuple[str, str]:
|
||||
import uuid
|
||||
|
||||
suffix = uuid.uuid4().hex[:12]
|
||||
return f"fcf_rel_{case}_{suffix}", f"fcf_rel.{case}_{suffix}"
|
||||
|
||||
|
||||
def _read_rows(table, columns: list[str], row_count: int) -> list[dict]:
|
||||
return sorted(
|
||||
table.search()
|
||||
.select(columns)
|
||||
.limit(row_count)
|
||||
.to_list(timeout=_query_timeout()),
|
||||
key=lambda row: row["row_id"],
|
||||
)
|
||||
|
||||
|
||||
def _emit_evidence(case: str, evidence: dict) -> None:
|
||||
import json
|
||||
|
||||
print(
|
||||
json.dumps(
|
||||
{"case": case, **evidence},
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_enterprise_reliability_core_lifecycle():
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import FunctionError
|
||||
from lancedb.expr import col
|
||||
|
||||
_require_live()
|
||||
table_name, function_name = _run_names("lifecycle")
|
||||
setup_db = _connect()
|
||||
setup_db.create_table(
|
||||
table_name,
|
||||
data=pa.Table.from_pylist(
|
||||
[
|
||||
{"row_id": 1, "value": 2},
|
||||
{"row_id": 2, "value": 5},
|
||||
{"row_id": 3, "value": None},
|
||||
],
|
||||
schema=pa.schema(
|
||||
[
|
||||
pa.field("row_id", pa.int64(), nullable=False),
|
||||
pa.field("value", pa.int64(), nullable=True),
|
||||
]
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
registration_job = setup_db.functions.register(function_name, reliable_double)
|
||||
registration_job_id = registration_job.id
|
||||
assert isinstance(registration_job_id, str) and registration_job_id
|
||||
registered = registration_job.wait(timeout=_job_timeout())
|
||||
assert type(registered) is lancedb.Function
|
||||
assert isinstance(registered.id, str) and registered.id
|
||||
with pytest.raises(AttributeError):
|
||||
registered.id = "mutated"
|
||||
|
||||
catalog_reader = _connect()
|
||||
by_name = catalog_reader.functions.get(function_name)
|
||||
by_id = catalog_reader.functions.get_by_id(registered.id)
|
||||
expected_identity = (
|
||||
registered.id,
|
||||
(("value", pa.int64()),),
|
||||
pa.int64(),
|
||||
True,
|
||||
)
|
||||
for function in (registered, by_name, by_id):
|
||||
assert type(function) is lancedb.Function
|
||||
assert (
|
||||
function.id,
|
||||
function.parameters,
|
||||
function.output_type,
|
||||
function.output_nullable,
|
||||
) == expected_identity
|
||||
|
||||
table = catalog_reader.open_table(table_name)
|
||||
create_job = table.add_generated_column(
|
||||
"derived",
|
||||
registered(value=col("value")),
|
||||
)
|
||||
create_job_id = create_job.id
|
||||
assert isinstance(create_job_id, str) and create_job_id
|
||||
assert create_job.wait(timeout=_job_timeout()) is None
|
||||
|
||||
complete_reader = _connect().open_table(table_name)
|
||||
complete_status = complete_reader.generated_column_status("derived")
|
||||
assert complete_status == "complete"
|
||||
initial_rows = _read_rows(
|
||||
complete_reader,
|
||||
["row_id", "value", "derived"],
|
||||
3,
|
||||
)
|
||||
assert initial_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 5, "derived": 10},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
complete_reader.update(where="row_id = 2", values={"value": 7})
|
||||
|
||||
incomplete_reader = _connect().open_table(table_name)
|
||||
changed_rows = _read_rows(incomplete_reader, ["row_id", "value"], 3)
|
||||
assert changed_rows == [
|
||||
{"row_id": 1, "value": 2},
|
||||
{"row_id": 2, "value": 7},
|
||||
{"row_id": 3, "value": None},
|
||||
]
|
||||
incomplete_status = incomplete_reader.generated_column_status("derived")
|
||||
assert incomplete_status == "incomplete"
|
||||
with pytest.raises(FunctionError) as raised:
|
||||
(
|
||||
incomplete_reader.search()
|
||||
.select(["row_id", "derived"])
|
||||
.limit(3)
|
||||
.to_list(timeout=_query_timeout())
|
||||
)
|
||||
assert raised.value.code == "generated_column_incomplete"
|
||||
|
||||
refresh_job = incomplete_reader.refresh_generated_column("derived")
|
||||
refresh_job_id = refresh_job.id
|
||||
assert isinstance(refresh_job_id, str) and refresh_job_id
|
||||
assert refresh_job.wait(timeout=_job_timeout()) is None
|
||||
|
||||
refreshed_reader = _connect().open_table(table_name)
|
||||
refreshed_status = refreshed_reader.generated_column_status("derived")
|
||||
assert refreshed_status == "complete"
|
||||
final_rows = _read_rows(
|
||||
refreshed_reader,
|
||||
["row_id", "value", "derived"],
|
||||
3,
|
||||
)
|
||||
assert final_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 7, "derived": 14},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
_emit_evidence(
|
||||
"core_lifecycle",
|
||||
{
|
||||
"final_rows": final_rows,
|
||||
"function_id": registered.id,
|
||||
"job_ids": {
|
||||
"create": create_job_id,
|
||||
"refresh": refresh_job_id,
|
||||
"register": registration_job_id,
|
||||
},
|
||||
"status": [
|
||||
complete_status,
|
||||
incomplete_status,
|
||||
refreshed_status,
|
||||
],
|
||||
"table": table_name,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_enterprise_reliability_restart_retention():
|
||||
import json
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
|
||||
_require_live()
|
||||
raw_evidence = os.environ.get("LANCEDB_FCF_E2E_RESTART_EVIDENCE")
|
||||
if not raw_evidence:
|
||||
pytest.skip(
|
||||
"LANCEDB_FCF_E2E_RESTART_EVIDENCE is required for restart retention"
|
||||
)
|
||||
|
||||
try:
|
||||
evidence = json.loads(raw_evidence)
|
||||
except json.JSONDecodeError as error:
|
||||
pytest.fail(f"LANCEDB_FCF_E2E_RESTART_EVIDENCE must be valid JSON: {error.msg}")
|
||||
|
||||
assert isinstance(evidence, dict), (
|
||||
"LANCEDB_FCF_E2E_RESTART_EVIDENCE must be a JSON object"
|
||||
)
|
||||
table_name = evidence.get("table")
|
||||
function_id = evidence.get("function_id")
|
||||
raw_job_ids = evidence.get("job_ids")
|
||||
assert isinstance(table_name, str) and table_name, (
|
||||
"restart evidence must contain a non-empty table"
|
||||
)
|
||||
assert isinstance(function_id, str) and function_id, (
|
||||
"restart evidence must contain a non-empty function_id"
|
||||
)
|
||||
assert isinstance(raw_job_ids, dict), (
|
||||
"restart evidence must contain a job_ids object"
|
||||
)
|
||||
job_ids = {}
|
||||
for job_kind in ("register", "create", "refresh"):
|
||||
job_id = raw_job_ids.get(job_kind)
|
||||
assert isinstance(job_id, str) and job_id, (
|
||||
f"restart evidence must contain a non-empty job_ids.{job_kind}"
|
||||
)
|
||||
job_ids[job_kind] = job_id
|
||||
|
||||
db = _connect()
|
||||
function = db.functions.get_by_id(function_id)
|
||||
expected_identity = (
|
||||
function_id,
|
||||
(("value", pa.int64()),),
|
||||
pa.int64(),
|
||||
True,
|
||||
)
|
||||
assert type(function) is lancedb.Function
|
||||
assert (
|
||||
function.id,
|
||||
function.parameters,
|
||||
function.output_type,
|
||||
function.output_nullable,
|
||||
) == expected_identity
|
||||
|
||||
jobs = {}
|
||||
for job_kind in ("register", "create", "refresh"):
|
||||
job = db.get_job(job_ids[job_kind])
|
||||
assert job is not None
|
||||
assert job.job_id == job_ids[job_kind]
|
||||
assert job.state == "finished"
|
||||
assert job.failure is None
|
||||
jobs[job_kind] = job
|
||||
|
||||
registered_result = jobs["register"].result
|
||||
assert type(registered_result) is lancedb.Function
|
||||
assert (
|
||||
registered_result.id,
|
||||
registered_result.parameters,
|
||||
registered_result.output_type,
|
||||
registered_result.output_nullable,
|
||||
) == expected_identity
|
||||
assert jobs["create"].result is None
|
||||
assert jobs["refresh"].result is None
|
||||
|
||||
table = db.open_table(table_name)
|
||||
status = table.generated_column_status("derived")
|
||||
assert status == "complete"
|
||||
assert table.count_rows() == 3
|
||||
final_rows = _read_rows(table, ["row_id", "value", "derived"], 3)
|
||||
assert final_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 7, "derived": 14},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
_emit_evidence(
|
||||
"restart_retention",
|
||||
{
|
||||
"final_rows": final_rows,
|
||||
"function_id": function_id,
|
||||
"generated_column_status": status,
|
||||
"job_ids": job_ids,
|
||||
"job_states": {
|
||||
job_kind: jobs[job_kind].state
|
||||
for job_kind in ("register", "create", "refresh")
|
||||
},
|
||||
"table": table_name,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_enterprise_reliability_failure_atomicity_and_worker_recovery():
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import JobFailedError
|
||||
from lancedb.expr import col
|
||||
|
||||
_require_live()
|
||||
table_name, failing_function_name = _run_names("worker_failure")
|
||||
_, healthy_function_name = _run_names("worker_recovery")
|
||||
row_count = 4
|
||||
setup_db = _connect()
|
||||
setup_db.create_table(
|
||||
table_name,
|
||||
data=pa.table(
|
||||
{
|
||||
"row_id": list(range(row_count)),
|
||||
"value": [1, 2, 3, 4],
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
registration_job = setup_db.functions.register(
|
||||
failing_function_name,
|
||||
terminate_worker_on_input,
|
||||
)
|
||||
failing_function = registration_job.wait(timeout=_job_timeout())
|
||||
assert type(failing_function) is lancedb.Function
|
||||
|
||||
table = setup_db.open_table(table_name)
|
||||
failed_create_job = table.add_generated_column(
|
||||
"must_not_publish",
|
||||
failing_function(value=col("value")),
|
||||
)
|
||||
failed_job_id = failed_create_job.id
|
||||
assert isinstance(failed_job_id, str) and failed_job_id
|
||||
with pytest.raises(JobFailedError) as raised:
|
||||
failed_create_job.wait(timeout=_job_timeout())
|
||||
assert raised.value.error_code == "udf_execution_failure"
|
||||
|
||||
first_description = _connect().get_job(failed_job_id)
|
||||
second_description = _connect().get_job(failed_job_id)
|
||||
for description in (first_description, second_description):
|
||||
assert description is not None
|
||||
assert description.job_id == failed_job_id
|
||||
assert description.state == "failed"
|
||||
assert description.failure is not None
|
||||
assert description.failure.error_code == "udf_execution_failure"
|
||||
|
||||
atomic_reader = _connect().open_table(table_name)
|
||||
assert "must_not_publish" not in atomic_reader.schema.names
|
||||
assert _read_rows(atomic_reader, ["row_id", "value"], row_count) == [
|
||||
{"row_id": 0, "value": 1},
|
||||
{"row_id": 1, "value": 2},
|
||||
{"row_id": 2, "value": 3},
|
||||
{"row_id": 3, "value": 4},
|
||||
]
|
||||
|
||||
healthy_registration_job = setup_db.functions.register(
|
||||
healthy_function_name,
|
||||
reliable_double,
|
||||
)
|
||||
healthy_function = healthy_registration_job.wait(timeout=_job_timeout())
|
||||
assert type(healthy_function) is lancedb.Function
|
||||
recovery_job = atomic_reader.add_generated_column(
|
||||
"recovered",
|
||||
healthy_function(value=col("value")),
|
||||
)
|
||||
recovery_job_id = recovery_job.id
|
||||
assert isinstance(recovery_job_id, str) and recovery_job_id
|
||||
assert recovery_job.wait(timeout=_job_timeout()) is None
|
||||
|
||||
recovered_reader = _connect().open_table(table_name)
|
||||
assert "must_not_publish" not in recovered_reader.schema.names
|
||||
assert recovered_reader.generated_column_status("recovered") == "complete"
|
||||
recovered_rows = _read_rows(
|
||||
recovered_reader,
|
||||
["row_id", "value", "recovered"],
|
||||
row_count,
|
||||
)
|
||||
assert recovered_rows == [
|
||||
{"row_id": 0, "value": 1, "recovered": 2},
|
||||
{"row_id": 1, "value": 2, "recovered": 4},
|
||||
{"row_id": 2, "value": 3, "recovered": 6},
|
||||
{"row_id": 3, "value": 4, "recovered": 8},
|
||||
]
|
||||
|
||||
_emit_evidence(
|
||||
"failure_atomicity_and_worker_recovery",
|
||||
{
|
||||
"failure_code": first_description.failure.error_code,
|
||||
"failed_job_id": failed_job_id,
|
||||
"recovered_rows": recovered_rows,
|
||||
"recovery_job_id": recovery_job_id,
|
||||
"table": table_name,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_enterprise_reliability_concurrent_refresh_fencing():
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import FunctionError, JobFailedError
|
||||
from lancedb.expr import col
|
||||
|
||||
_require_live()
|
||||
table_name, function_name = _run_names("refresh_fencing")
|
||||
row_count = 1024
|
||||
setup_db = _connect()
|
||||
setup_db.create_table(
|
||||
table_name,
|
||||
data=pa.table(
|
||||
{
|
||||
"row_id": list(range(row_count)),
|
||||
"value": list(range(row_count)),
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
registration_job = setup_db.functions.register(function_name, slow_triple)
|
||||
function = registration_job.wait(timeout=_job_timeout())
|
||||
assert type(function) is lancedb.Function
|
||||
|
||||
table = setup_db.open_table(table_name)
|
||||
create_job = table.add_generated_column(
|
||||
"derived",
|
||||
function(value=col("value")),
|
||||
)
|
||||
assert create_job.wait(timeout=_job_timeout()) is None
|
||||
initial_reader = _connect().open_table(table_name)
|
||||
assert initial_reader.generated_column_status("derived") == "complete"
|
||||
|
||||
initial_reader.update(where="row_id = 0", values={"value": 10_000})
|
||||
incomplete_reader = _connect().open_table(table_name)
|
||||
assert incomplete_reader.generated_column_status("derived") == "incomplete"
|
||||
|
||||
refresh_job = incomplete_reader.refresh_generated_column("derived")
|
||||
refresh_job_id = refresh_job.id
|
||||
assert isinstance(refresh_job_id, str) and refresh_job_id
|
||||
deadline = time.monotonic() + _RUNNING_DEADLINE_SECONDS
|
||||
observed_states = []
|
||||
running_observations = 0
|
||||
while running_observations < 2:
|
||||
state = refresh_job.status()
|
||||
if not observed_states or observed_states[-1] != state:
|
||||
observed_states.append(state)
|
||||
if state == "running":
|
||||
running_observations += 1
|
||||
else:
|
||||
running_observations = 0
|
||||
assert state not in {"finished", "failed", "cancelled"}
|
||||
assert time.monotonic() < deadline
|
||||
if running_observations < 2:
|
||||
time.sleep(0.05)
|
||||
|
||||
concurrent_writer = _connect().open_table(table_name)
|
||||
concurrent_writer.update(where="row_id = 1", values={"value": 20_000})
|
||||
with pytest.raises(JobFailedError) as raised:
|
||||
refresh_job.wait(timeout=_job_timeout())
|
||||
assert raised.value.error_code == "stale_or_conflicting_input"
|
||||
|
||||
stale_job = _connect().get_job(refresh_job_id)
|
||||
assert stale_job is not None
|
||||
assert stale_job.job_id == refresh_job_id
|
||||
assert stale_job.state == "failed"
|
||||
assert stale_job.failure is not None
|
||||
assert stale_job.failure.error_code == raised.value.error_code
|
||||
if observed_states[-1] != stale_job.state:
|
||||
observed_states.append(stale_job.state)
|
||||
|
||||
stale_reader = _connect().open_table(table_name)
|
||||
stale_rows = _read_rows(stale_reader, ["row_id", "value"], row_count)
|
||||
assert len(stale_rows) == row_count
|
||||
for row_id, row in enumerate(stale_rows):
|
||||
expected_value = 10_000 if row_id == 0 else 20_000 if row_id == 1 else row_id
|
||||
assert (row["row_id"], row["value"]) == (row_id, expected_value)
|
||||
assert stale_reader.generated_column_status("derived") == "incomplete"
|
||||
with pytest.raises(FunctionError) as incomplete:
|
||||
(
|
||||
stale_reader.search()
|
||||
.select(["row_id", "derived"])
|
||||
.limit(row_count)
|
||||
.to_list(timeout=_query_timeout())
|
||||
)
|
||||
assert incomplete.value.code == "generated_column_incomplete"
|
||||
|
||||
resubmitted_job = stale_reader.refresh_generated_column("derived")
|
||||
resubmitted_job_id = resubmitted_job.id
|
||||
assert isinstance(resubmitted_job_id, str) and resubmitted_job_id
|
||||
assert resubmitted_job.wait(timeout=_job_timeout()) is None
|
||||
|
||||
final_reader = _connect().open_table(table_name)
|
||||
final_status = final_reader.generated_column_status("derived")
|
||||
assert final_status == "complete"
|
||||
final_rows = _read_rows(
|
||||
final_reader,
|
||||
["row_id", "value", "derived"],
|
||||
row_count,
|
||||
)
|
||||
assert len(final_rows) == row_count
|
||||
final_checksum = 0
|
||||
for row_id, row in enumerate(final_rows):
|
||||
expected_value = 10_000 if row_id == 0 else 20_000 if row_id == 1 else row_id
|
||||
assert (row["row_id"], row["value"], row["derived"]) == (
|
||||
row_id,
|
||||
expected_value,
|
||||
expected_value * 3,
|
||||
)
|
||||
final_checksum += row["derived"]
|
||||
|
||||
_emit_evidence(
|
||||
"concurrent_refresh_fencing",
|
||||
{
|
||||
"failure_code": stale_job.failure.error_code,
|
||||
"final_checksum": final_checksum,
|
||||
"final_status": final_status,
|
||||
"observed_states": observed_states,
|
||||
"resubmitted_job_id": resubmitted_job_id,
|
||||
"row_count": row_count,
|
||||
"stale_job_id": refresh_job_id,
|
||||
"table": table_name,
|
||||
},
|
||||
)
|
||||
@@ -0,0 +1,268 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract: Python projection of JobFailure.error_code / JobFailedError.error_code.
|
||||
|
||||
Public Function failures expose eight stable string categories. Asynchronous
|
||||
errors remain the unified JobFailedError and JobFailureInfo. Python must
|
||||
project the optional exact error_code string already supplied structurally by
|
||||
Rust: preserve a known code, preserve an unknown nonempty future code
|
||||
byte-for-byte, and return None for legacy failure payloads without error_code.
|
||||
Never infer or override a code from message, phase, retryable, HTTP status,
|
||||
job type, or state.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from datetime import timedelta
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import JobFailedError
|
||||
|
||||
_DESCRIBE_PATH = "/v1/jobs/describe"
|
||||
_KNOWN_CODE = "name_or_function_not_found"
|
||||
_CONFLICTING_STABLE_IN_MESSAGE = "definition_validation_failure"
|
||||
_UNKNOWN_CODE = "enterprise_future_category_xyz"
|
||||
_WAIT_KNOWN_CODE = "unsupported_runtime_or_capability"
|
||||
_WAIT_CONFLICTING_IN_MESSAGE = "revoked_function"
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _failed_describe_body(
|
||||
*,
|
||||
job_id: str,
|
||||
error_code: Optional[str] = None,
|
||||
include_error_code: bool = True,
|
||||
phase: str = "execute",
|
||||
message: str = "worker died",
|
||||
retryable: bool = False,
|
||||
job_type: str = "create_index",
|
||||
) -> dict[str, Any]:
|
||||
failure: dict[str, Any] = {
|
||||
"phase": phase,
|
||||
"message": message,
|
||||
"retryable": retryable,
|
||||
}
|
||||
if include_error_code:
|
||||
failure["error_code"] = error_code
|
||||
return {
|
||||
"job_id": job_id,
|
||||
"job_type": job_type,
|
||||
"job_state": "FAILED",
|
||||
"creation_ms": 1000,
|
||||
"spec": {},
|
||||
"failure": failure,
|
||||
}
|
||||
|
||||
|
||||
def _describe_handler(bodies_by_job_id: dict[str, dict[str, Any]]):
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _DESCRIBE_PATH
|
||||
payload = json.loads(_read_body(request).decode("utf-8") or "{}")
|
||||
job_id = payload["job_id"]
|
||||
body = bodies_by_job_id.get(job_id)
|
||||
if body is None:
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
return
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def test_get_job_failure_error_code_known_not_inferred_from_message():
|
||||
"""Structural error_code wins; conflicting message text must not override."""
|
||||
body = _failed_describe_body(
|
||||
job_id="job-known",
|
||||
error_code=_KNOWN_CODE,
|
||||
phase="validate",
|
||||
message=f"looks like {_CONFLICTING_STABLE_IN_MESSAGE}",
|
||||
retryable=False,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-known": body})) as db:
|
||||
description = db.get_job("job-known")
|
||||
assert description is not None
|
||||
failure = description.failure
|
||||
assert failure is not None
|
||||
assert failure.error_code == _KNOWN_CODE
|
||||
assert failure.error_code != _CONFLICTING_STABLE_IN_MESSAGE
|
||||
assert failure.phase == "validate"
|
||||
assert failure.message == f"looks like {_CONFLICTING_STABLE_IN_MESSAGE}"
|
||||
assert failure.retryable is False
|
||||
|
||||
|
||||
def test_get_job_failure_error_code_unknown_preserved_byte_for_byte():
|
||||
body = _failed_describe_body(
|
||||
job_id="job-unknown",
|
||||
error_code=_UNKNOWN_CODE,
|
||||
phase="execute",
|
||||
message=f"new category mentioning {_KNOWN_CODE}",
|
||||
retryable=True,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-unknown": body})) as db:
|
||||
failure = db.get_job("job-unknown").failure
|
||||
assert failure.error_code == _UNKNOWN_CODE
|
||||
assert failure.error_code != _KNOWN_CODE
|
||||
|
||||
|
||||
def test_get_job_failure_error_code_absent_is_none():
|
||||
"""Legacy describe payloads without error_code must not invent a category."""
|
||||
body = _failed_describe_body(
|
||||
job_id="job-legacy",
|
||||
include_error_code=False,
|
||||
phase="execute",
|
||||
message=f"{_KNOWN_CODE} in logs",
|
||||
retryable=True,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-legacy": body})) as db:
|
||||
failure = db.get_job("job-legacy").failure
|
||||
assert failure.error_code is None
|
||||
assert failure.phase == "execute"
|
||||
assert failure.retryable is True
|
||||
|
||||
|
||||
def test_sync_job_wait_job_failed_error_code_known_not_inferred():
|
||||
body = _failed_describe_body(
|
||||
job_id="job-wait-known",
|
||||
error_code=_WAIT_KNOWN_CODE,
|
||||
phase="dispatch",
|
||||
message=f"{_WAIT_CONFLICTING_IN_MESSAGE} in transport logs",
|
||||
retryable=False,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-wait-known": body})) as db:
|
||||
with pytest.raises(JobFailedError) as exc_info:
|
||||
db.job("job-wait-known").wait(timeout=timedelta(seconds=5))
|
||||
|
||||
err = exc_info.value
|
||||
assert isinstance(err, JobFailedError)
|
||||
assert err.error_code == _WAIT_KNOWN_CODE
|
||||
assert err.error_code != _WAIT_CONFLICTING_IN_MESSAGE
|
||||
|
||||
|
||||
def test_sync_job_wait_job_failed_error_code_absent_is_none():
|
||||
body = _failed_describe_body(
|
||||
job_id="job-wait-legacy",
|
||||
include_error_code=False,
|
||||
phase="execute",
|
||||
message=f"{_WAIT_KNOWN_CODE} mentioned only in message",
|
||||
retryable=True,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-wait-legacy": body})) as db:
|
||||
with pytest.raises(JobFailedError) as exc_info:
|
||||
db.job("job-wait-legacy").wait(timeout=timedelta(seconds=5))
|
||||
|
||||
assert exc_info.value.error_code is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_job_wait_job_failed_error_code_unknown_preserved():
|
||||
body = _failed_describe_body(
|
||||
job_id="job-wait-unknown",
|
||||
error_code=_UNKNOWN_CODE,
|
||||
phase="execute",
|
||||
message=f"future code with {_WAIT_KNOWN_CODE} in text",
|
||||
retryable=False,
|
||||
)
|
||||
async with _mock_remote_db_async(
|
||||
_describe_handler({"job-wait-unknown": body})
|
||||
) as db:
|
||||
with pytest.raises(JobFailedError) as exc_info:
|
||||
await db.job("job-wait-unknown").wait(timeout=timedelta(seconds=5))
|
||||
|
||||
err = exc_info.value
|
||||
assert err.error_code == _UNKNOWN_CODE
|
||||
assert err.error_code != _WAIT_KNOWN_CODE
|
||||
|
||||
|
||||
def test_job_failed_error_legacy_message_construction_error_code_is_none():
|
||||
err = JobFailedError("legacy construction with only a message")
|
||||
assert err.error_code is None
|
||||
|
||||
|
||||
def test_job_failed_error_error_code_is_read_only():
|
||||
err = JobFailedError("message")
|
||||
with pytest.raises(AttributeError):
|
||||
err.error_code = _KNOWN_CODE
|
||||
@@ -0,0 +1,634 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python first-class Function catalog lookup."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.remote.errors import HttpError
|
||||
|
||||
|
||||
def _function_error_cls(*, required: bool = True):
|
||||
"""Resolve FunctionError from the live module (records RED when absent)."""
|
||||
from lancedb import exceptions as exc_mod
|
||||
|
||||
cls = getattr(exc_mod, "FunctionError", None)
|
||||
if cls is None:
|
||||
if required:
|
||||
pytest.fail("lancedb.exceptions.FunctionError is missing")
|
||||
return type("MissingFunctionError", (), {})
|
||||
return cls
|
||||
|
||||
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_LOOKUP_CATALOG_NAME = "text.normalize.lookup-name"
|
||||
_LOOKUP_FUNCTION_ID = "fn.exact.lookup-handle"
|
||||
_LOOKUP_SERVER_MESSAGE_MARKER = (
|
||||
"SERVER_LOOKUP_DIAGNOSTIC_MARKER name=text.normalize.lookup-name "
|
||||
"id=fn.exact.lookup-handle"
|
||||
)
|
||||
_SENSITIVE_BODY_MARKER = "SENSITIVE_LOOKUP_BODY_MARKER"
|
||||
_UNKNOWN_CODE = "enterprise_future_lookup_category_xyz"
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as job-result
|
||||
# tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust FileWriter.
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_DELETED_LOOKUP_KEYWORDS = (
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"replace",
|
||||
"expected_current_function_id",
|
||||
"list",
|
||||
"alias",
|
||||
"lineage",
|
||||
"FunctionVersion",
|
||||
)
|
||||
|
||||
|
||||
def _sample_function_wire() -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": _LOOKUP_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _lookup_success_body(
|
||||
*,
|
||||
function: dict[str, Any] | None = None,
|
||||
extra_outer: dict[str, Any] | None = None,
|
||||
) -> bytes:
|
||||
body: dict[str, Any] = {"function": function or _sample_function_wire()}
|
||||
if extra_outer:
|
||||
body.update(extra_outer)
|
||||
return json.dumps(body).encode("utf-8")
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_chain_text(exc: BaseException) -> str:
|
||||
parts: list[str] = []
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(str(current))
|
||||
parts.append(repr(current))
|
||||
current = current.__cause__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _assert_payload_free(exc: BaseException) -> None:
|
||||
text = _exception_chain_text(exc)
|
||||
assert _LOOKUP_SERVER_MESSAGE_MARKER not in text
|
||||
assert _LOOKUP_CATALOG_NAME not in text
|
||||
assert _LOOKUP_FUNCTION_ID not in text
|
||||
assert _SENSITIVE_BODY_MARKER not in text
|
||||
|
||||
|
||||
def _assert_exact_lookup_function(function: object) -> None:
|
||||
assert type(function) is lancedb.Function
|
||||
assert function.id == _LOOKUP_FUNCTION_ID
|
||||
assert not hasattr(function, "name")
|
||||
assert function.parameters == (
|
||||
("text", pa.string()),
|
||||
("limit", pa.int32()),
|
||||
)
|
||||
assert function.output_type == pa.string()
|
||||
assert function.output_nullable is True
|
||||
assert _LOOKUP_CATALOG_NAME not in repr(function)
|
||||
assert _LOOKUP_CATALOG_NAME not in str(function)
|
||||
|
||||
|
||||
def _assert_name_request(raw: bytes, body: dict[str, Any]) -> None:
|
||||
assert raw
|
||||
assert body == {"name": _LOOKUP_CATALOG_NAME}
|
||||
assert "function_id" not in body
|
||||
|
||||
|
||||
def _assert_id_request(raw: bytes, body: dict[str, Any]) -> None:
|
||||
assert raw
|
||||
assert body == {"function_id": _LOOKUP_FUNCTION_ID}
|
||||
assert "name" not in body
|
||||
|
||||
|
||||
def _assert_native_lookup_methods_present() -> None:
|
||||
assert hasattr(_native.Connection, "_lookup_function_by_name")
|
||||
assert hasattr(_native.Connection, "_lookup_function_by_id")
|
||||
assert callable(getattr(_native.Connection, "_lookup_function_by_name"))
|
||||
assert callable(getattr(_native.Connection, "_lookup_function_by_id"))
|
||||
|
||||
|
||||
def test_native_connection_exposes_private_lookup_methods():
|
||||
_assert_native_lookup_methods_present()
|
||||
|
||||
|
||||
def test_sync_remote_get_by_name_exact_request_and_function_shape():
|
||||
_assert_native_lookup_methods_present()
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _LOOKUP_PATH
|
||||
raw = _read_body(request)
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
assert not hasattr(db, "get_function")
|
||||
function = db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
_assert_name_request(seen["raw"], seen["body"])
|
||||
_assert_exact_lookup_function(function)
|
||||
|
||||
|
||||
def test_sync_remote_get_by_id_exact_request_and_function_shape():
|
||||
_assert_native_lookup_methods_present()
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _LOOKUP_PATH
|
||||
raw = _read_body(request)
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
function = db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
_assert_id_request(seen["raw"], seen["body"])
|
||||
_assert_exact_lookup_function(function)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_get_by_name_and_id():
|
||||
_assert_native_lookup_methods_present()
|
||||
name_seen: dict[str, Any] = {}
|
||||
id_seen: dict[str, Any] = {}
|
||||
stage = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _LOOKUP_PATH
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
stage["n"] += 1
|
||||
if stage["n"] == 1:
|
||||
name_seen["raw"] = raw
|
||||
name_seen["body"] = body
|
||||
else:
|
||||
id_seen["raw"] = raw
|
||||
id_seen["body"] = body
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
|
||||
async with _mock_remote_db_async(handler) as db:
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
by_name = await db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
by_id = await db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
_assert_name_request(name_seen["raw"], name_seen["body"])
|
||||
_assert_id_request(id_seen["raw"], id_seen["body"])
|
||||
_assert_exact_lookup_function(by_name)
|
||||
_assert_exact_lookup_function(by_id)
|
||||
|
||||
|
||||
def test_sync_remote_get_accepts_additive_outer_success_fields():
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
_lookup_success_body(
|
||||
extra_outer={
|
||||
"server_extra": {"ok": True},
|
||||
"request_echo_name": _LOOKUP_CATALOG_NAME,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
function = db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
_assert_exact_lookup_function(function)
|
||||
|
||||
|
||||
def test_empty_name_and_id_reject_before_transport():
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
_read_body(request)
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"should not be reached")
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.get("")
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.get_by_id("")
|
||||
|
||||
assert received["n"] == 0
|
||||
|
||||
|
||||
def test_local_sync_lookup_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_lookup_methods_present()
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_lookup_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_lookup_methods_present()
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
|
||||
|
||||
def test_explicit_known_code_is_function_error_with_exact_code():
|
||||
body = {
|
||||
"error_code": "name_or_function_not_found",
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
"looks_like": "definition_validation_failure",
|
||||
}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
request.send_response(404)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
err = exc_info.value
|
||||
assert isinstance(err, function_error)
|
||||
assert err.code == "name_or_function_not_found"
|
||||
assert err.code != "definition_validation_failure"
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
def test_explicit_unknown_code_preserved_despite_status_and_message():
|
||||
body = {
|
||||
"error_code": _UNKNOWN_CODE,
|
||||
"message": (
|
||||
f"{_LOOKUP_SERVER_MESSAGE_MARKER} revoked_function "
|
||||
"name_or_function_not_found"
|
||||
),
|
||||
}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
raw = _read_body(request)
|
||||
assert json.loads(raw.decode("utf-8")) == {"function_id": _LOOKUP_FUNCTION_ID}
|
||||
request.send_response(409)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
err = exc_info.value
|
||||
assert err.code == _UNKNOWN_CODE
|
||||
assert err.code != "revoked_function"
|
||||
assert err.code != "name_or_function_not_found"
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,status,response_body",
|
||||
[
|
||||
(
|
||||
"missing_code_404",
|
||||
404,
|
||||
{
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"empty_code",
|
||||
400,
|
||||
{
|
||||
"error_code": "",
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"wrong_type_code",
|
||||
400,
|
||||
{
|
||||
"error_code": 123,
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"null_code",
|
||||
404,
|
||||
{
|
||||
"error_code": None,
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"non_json",
|
||||
404,
|
||||
f"not-json {_LOOKUP_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_invalid_or_missing_error_code_is_payload_free_http(
|
||||
label: str, status: int, response_body: object
|
||||
):
|
||||
del label # parametrize label for failure diagnosis only
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
request.send_response(status)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
if isinstance(response_body, str):
|
||||
request.wfile.write(response_body.encode("utf-8"))
|
||||
else:
|
||||
request.wfile.write(json.dumps(response_body).encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert not isinstance(err, _function_error_cls(required=False))
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,response_body",
|
||||
[
|
||||
(
|
||||
"missing_function",
|
||||
{
|
||||
"server_extra": True,
|
||||
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
},
|
||||
),
|
||||
(
|
||||
"null_function",
|
||||
{
|
||||
"function": None,
|
||||
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
},
|
||||
),
|
||||
(
|
||||
"wrong_type_function",
|
||||
{
|
||||
"function": "not-an-object",
|
||||
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
},
|
||||
),
|
||||
(
|
||||
"invalid_function_shape",
|
||||
{
|
||||
"function": {
|
||||
"format_version": 1,
|
||||
"id": _LOOKUP_FUNCTION_ID,
|
||||
# missing signature
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
}
|
||||
},
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_malformed_success_is_payload_free_http(label: str, response_body: dict):
|
||||
del label
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(response_body).encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert not isinstance(err, _function_error_cls(required=False))
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
def test_function_error_surface_omits_server_marker_name_and_id():
|
||||
body = {
|
||||
"error_code": "name_or_function_not_found",
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
"function_id": _LOOKUP_FUNCTION_ID,
|
||||
"name": _LOOKUP_CATALOG_NAME,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
_read_body(request)
|
||||
request.send_response(404)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
err = exc_info.value
|
||||
_assert_payload_free(err)
|
||||
assert getattr(err, "code", None) == "name_or_function_not_found"
|
||||
|
||||
|
||||
def test_no_direct_db_lookup_methods_and_no_deleted_keywords():
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
_read_body(request)
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"should not be reached")
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
assert not hasattr(db, "lookup_function")
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
assert not hasattr(db, "get_function")
|
||||
assert not hasattr(db.functions, "get_by_name")
|
||||
assert not hasattr(db.functions, "list")
|
||||
|
||||
for keyword in _DELETED_LOOKUP_KEYWORDS:
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME, **{keyword: True})
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.get_by_id(_LOOKUP_FUNCTION_ID, **{keyword: True})
|
||||
|
||||
assert received["n"] == 0
|
||||
|
||||
|
||||
def test_function_error_is_not_top_level_export():
|
||||
assert not hasattr(lancedb, "FunctionError")
|
||||
function_error = _function_error_cls()
|
||||
assert issubclass(function_error, RuntimeError)
|
||||
@@ -0,0 +1,398 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for Python first-class Function registration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from datetime import timedelta
|
||||
from typing import Any, Callable
|
||||
from unittest import mock
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
import lancedb._udf as _udf_mod
|
||||
import lancedb.job
|
||||
from lancedb import FunctionCapability, udf
|
||||
from lancedb.remote.errors import HttpError
|
||||
|
||||
_SOURCE_MARKER = "registration-source-marker-unique-xyz"
|
||||
_SECRET_REFERENCE = "secret://team/registration-redact-token-xyz"
|
||||
_SECRET_ENV = "REGISTER_API_TOKEN"
|
||||
_NETWORK_ORIGIN = "https://api.registration-example.com"
|
||||
_FUNCTION_NAME = "text.normalize"
|
||||
_FUNCTION_ID_RETRY = "fn.register-retry-1"
|
||||
_JOB_ID_RETRY = "job-register-retry-1"
|
||||
_JOB_ID_ASYNC = "job-register-async-1"
|
||||
_REGISTER_PATH = "/v1/functions/register"
|
||||
_DESCRIBE_PATH = "/v1/jobs/describe"
|
||||
|
||||
_DELETED_REGISTER_KEYWORDS = (
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"replace",
|
||||
"expected_current_function_id",
|
||||
)
|
||||
|
||||
_SPEC_KEYS = {
|
||||
"format_version",
|
||||
"name",
|
||||
"definition",
|
||||
"expected_current_function_id",
|
||||
}
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"text": pa.string(), "limit": pa.int32()},
|
||||
output=pa.string(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=True,
|
||||
capabilities=[
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
],
|
||||
)
|
||||
def packable_register_normalize(text, limit):
|
||||
"""registration-source-marker-unique-xyz."""
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def _definition_json(fn: object) -> dict[str, Any]:
|
||||
payload = _udf_mod._build_function_definition(fn)._to_json()
|
||||
if isinstance(payload, bytes):
|
||||
return json.loads(payload.decode("utf-8"))
|
||||
assert isinstance(payload, str)
|
||||
return json.loads(payload)
|
||||
|
||||
|
||||
def _expected_register_spec(name: str, fn: object) -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"name": name,
|
||||
"definition": _definition_json(fn),
|
||||
"expected_current_function_id": None,
|
||||
}
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_chain_text(exc: BaseException) -> str:
|
||||
parts: list[str] = []
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(str(current))
|
||||
parts.append(repr(current))
|
||||
current = current.__cause__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _assert_markers_absent_from_exception(exc: BaseException) -> None:
|
||||
text = _exception_chain_text(exc)
|
||||
assert _SOURCE_MARKER not in text
|
||||
assert _SECRET_REFERENCE not in text
|
||||
|
||||
|
||||
def _assert_exact_register_spec(body: dict[str, Any], expected: dict[str, Any]) -> None:
|
||||
assert set(body) == _SPEC_KEYS
|
||||
assert body == expected
|
||||
assert body["format_version"] == 1
|
||||
assert body["expected_current_function_id"] is None
|
||||
assert _SOURCE_MARKER in json.dumps(body["definition"])
|
||||
assert any(
|
||||
capability.get("reference") == _SECRET_REFERENCE
|
||||
for capability in body["definition"]["capabilities"]
|
||||
)
|
||||
|
||||
|
||||
def test_sync_remote_register_retries_exact_wire_and_returns_job():
|
||||
expected_spec = _expected_register_spec(_FUNCTION_NAME, packable_register_normalize)
|
||||
attempts: list[dict[str, Any]] = []
|
||||
describe_calls: list[dict[str, Any]] = []
|
||||
function_result_wire = {
|
||||
"kind": "function",
|
||||
"format_version": 1,
|
||||
"function": {
|
||||
"format_version": 1,
|
||||
"id": _FUNCTION_ID_RETRY,
|
||||
"signature": expected_spec["definition"]["signature"],
|
||||
},
|
||||
}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _REGISTER_PATH:
|
||||
request_id = request.headers.get("x-request-id")
|
||||
attempts.append(
|
||||
{
|
||||
"request_id": request_id,
|
||||
"raw": raw,
|
||||
"body": json.loads(raw.decode("utf-8")),
|
||||
}
|
||||
)
|
||||
if len(attempts) == 1:
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"transient register failure")
|
||||
return
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps({"job_id": _JOB_ID_RETRY}).encode("utf-8"))
|
||||
return
|
||||
|
||||
assert request.path == _DESCRIBE_PATH
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body["job_id"] == _JOB_ID_RETRY
|
||||
describe_calls.append(body)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
"job_id": _JOB_ID_RETRY,
|
||||
"job_state": "DONE",
|
||||
"job_type": "register_function",
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
"result": function_result_wire,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
package_calls = {"n": 0}
|
||||
original_package = _udf_mod._package_udf
|
||||
|
||||
def counting_package(fn: object):
|
||||
package_calls["n"] += 1
|
||||
return original_package(fn)
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
assert not hasattr(db, "register_function")
|
||||
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
|
||||
job = db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_RETRY
|
||||
waited = job.wait(timeout=timedelta(seconds=5))
|
||||
|
||||
assert package_calls["n"] == 1
|
||||
assert len(attempts) == 2
|
||||
first, second = attempts
|
||||
assert isinstance(first["request_id"], str) and first["request_id"]
|
||||
assert first["request_id"] == second["request_id"]
|
||||
assert first["raw"] == second["raw"]
|
||||
assert first["raw"]
|
||||
_assert_exact_register_spec(first["body"], expected_spec)
|
||||
_assert_exact_register_spec(second["body"], expected_spec)
|
||||
|
||||
assert len(describe_calls) == 1
|
||||
assert describe_calls[0]["job_id"] == _JOB_ID_RETRY
|
||||
assert type(waited) is lancedb.Function
|
||||
assert waited.id == _FUNCTION_ID_RETRY
|
||||
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
|
||||
assert waited.output_type == pa.string()
|
||||
assert waited.output_nullable is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_register_returns_async_job_with_exact_spec():
|
||||
expected_spec = _expected_register_spec(_FUNCTION_NAME, packable_register_normalize)
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _REGISTER_PATH
|
||||
raw = _read_body(request)
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
seen["request_id"] = request.headers.get("x-request-id")
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
|
||||
|
||||
async with _mock_remote_db_async(handler) as db:
|
||||
assert not hasattr(db, "register_function")
|
||||
job = await db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert seen.get("raw")
|
||||
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
|
||||
_assert_exact_register_spec(seen["body"], expected_spec)
|
||||
assert type(job) is lancedb.job.AsyncJob
|
||||
assert job.id == _JOB_ID_ASYNC
|
||||
|
||||
|
||||
def test_sync_remote_register_http_error_omits_source_and_secret_markers():
|
||||
echoed = f"register failed with {_SOURCE_MARKER} and {_SECRET_REFERENCE}"
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
assert request.path == _REGISTER_PATH
|
||||
_read_body(request)
|
||||
request.send_response(400)
|
||||
request.end_headers()
|
||||
request.wfile.write(echoed.encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert received["n"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert err.status_code == 400
|
||||
_assert_markers_absent_from_exception(err)
|
||||
|
||||
|
||||
def test_empty_name_rejects_before_http():
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
_read_body(request)
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"should not be reached")
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.register("", packable_register_normalize)
|
||||
|
||||
assert received["n"] == 0
|
||||
|
||||
|
||||
def test_local_sync_register_not_implemented_without_table_mutation(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "register_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_register_not_implemented_without_table_mutation(tmp_path):
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "register_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
|
||||
|
||||
@pytest.mark.parametrize("keyword", _DELETED_REGISTER_KEYWORDS)
|
||||
def test_register_rejects_deleted_overdesign_keywords_before_submission(keyword):
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
_read_body(request)
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"should not be reached")
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.register(
|
||||
_FUNCTION_NAME,
|
||||
packable_register_normalize,
|
||||
**{keyword: True},
|
||||
)
|
||||
|
||||
assert received["n"] == 0
|
||||
@@ -0,0 +1,719 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python conditional first-class Function name removal."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.remote.errors import HttpError
|
||||
|
||||
_REMOVE_PATH = "/v1/functions/remove"
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_REMOVE_CATALOG_NAME = "text.normalize.remove-name"
|
||||
_REMOVE_FUNCTION_ID = "fn.exact.remove-handle"
|
||||
_REMOVE_SERVER_MESSAGE_MARKER = (
|
||||
"SERVER_REMOVE_DIAGNOSTIC_MARKER name=text.normalize.remove-name "
|
||||
"id=fn.exact.remove-handle"
|
||||
)
|
||||
_SENSITIVE_BODY_MARKER = "SENSITIVE_REMOVE_BODY_MARKER"
|
||||
_CONFLICTING_MESSAGE_CODE = "revoked_function"
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
|
||||
# replace tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_DELETED_REMOVE_KEYWORDS = (
|
||||
"expected_current_function_id",
|
||||
"function_id",
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"version",
|
||||
"force",
|
||||
"if_exists",
|
||||
"revoke",
|
||||
"delete",
|
||||
)
|
||||
|
||||
|
||||
def _function_error_cls(*, required: bool = True):
|
||||
"""Resolve FunctionError from the live module (records RED when absent)."""
|
||||
from lancedb import exceptions as exc_mod
|
||||
|
||||
cls = getattr(exc_mod, "FunctionError", None)
|
||||
if cls is None:
|
||||
if required:
|
||||
pytest.fail("lancedb.exceptions.FunctionError is missing")
|
||||
return type("MissingFunctionError", (), {})
|
||||
return cls
|
||||
|
||||
|
||||
def _sample_function_wire() -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": _REMOVE_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _lookup_success_body() -> bytes:
|
||||
return json.dumps({"function": _sample_function_wire()}).encode("utf-8")
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
def _close_db(db: Any) -> None:
|
||||
with contextlib.suppress(Exception):
|
||||
inner = getattr(db, "_conn", None)
|
||||
if inner is not None:
|
||||
inner.close()
|
||||
return
|
||||
close = getattr(db, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
db = None
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
if db is not None:
|
||||
_close_db(db)
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
db = None
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
if db is not None:
|
||||
_close_db(db)
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_chain_text(exc: BaseException) -> str:
|
||||
parts: list[str] = []
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(str(current))
|
||||
parts.append(repr(current))
|
||||
current = current.__cause__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _assert_payload_free(exc: BaseException) -> None:
|
||||
text = _exception_chain_text(exc)
|
||||
assert _REMOVE_SERVER_MESSAGE_MARKER not in text
|
||||
assert _REMOVE_CATALOG_NAME not in text
|
||||
assert _REMOVE_FUNCTION_ID not in text
|
||||
assert _SENSITIVE_BODY_MARKER not in text
|
||||
|
||||
|
||||
def _assert_exact_remove_function(function: object) -> None:
|
||||
assert type(function) is lancedb.Function
|
||||
assert function.id == _REMOVE_FUNCTION_ID
|
||||
assert not hasattr(function, "name")
|
||||
assert function.parameters == (
|
||||
("text", pa.string()),
|
||||
("limit", pa.int32()),
|
||||
)
|
||||
assert function.output_type == pa.string()
|
||||
assert function.output_nullable is True
|
||||
assert _REMOVE_CATALOG_NAME not in repr(function)
|
||||
assert _REMOVE_CATALOG_NAME not in str(function)
|
||||
|
||||
|
||||
def _assert_exact_remove_request(
|
||||
request: http.server.BaseHTTPRequestHandler,
|
||||
raw: bytes,
|
||||
body: dict[str, Any],
|
||||
*,
|
||||
expected_id: str,
|
||||
) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _REMOVE_PATH
|
||||
assert "?" not in request.path
|
||||
assert raw
|
||||
assert body == {
|
||||
"name": _REMOVE_CATALOG_NAME,
|
||||
"expected_current_function_id": expected_id,
|
||||
}
|
||||
assert set(body) == {"name", "expected_current_function_id"}
|
||||
assert "format_version" not in body
|
||||
assert "function_id" not in body
|
||||
assert "function" not in body
|
||||
assert "signature" not in body
|
||||
assert "job_id" not in body
|
||||
assert "idempotency_key" not in body
|
||||
assert "user_version" not in body
|
||||
assert "force" not in body
|
||||
assert "if_exists" not in body
|
||||
request_id = request.headers.get("x-request-id")
|
||||
assert isinstance(request_id, str) and request_id
|
||||
|
||||
|
||||
def _assert_native_remove_method_present() -> None:
|
||||
assert hasattr(_native.Connection, "_remove_function_name")
|
||||
assert callable(getattr(_native.Connection, "_remove_function_name"))
|
||||
|
||||
|
||||
def _lookup_success_handler(
|
||||
counters: dict[str, int],
|
||||
*,
|
||||
after_lookup: (
|
||||
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
|
||||
) = None,
|
||||
):
|
||||
"""Serve exact name lookup; optionally continue for remove."""
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
counters["lookup"] = counters.get("lookup", 0) + 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body == {"name": _REMOVE_CATALOG_NAME}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
if after_lookup is not None:
|
||||
after_lookup(request, raw)
|
||||
return
|
||||
counters["remove"] = counters.get("remove", 0) + 1
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected remove")
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _observe_current(db) -> lancedb.Function:
|
||||
current = db.functions.get(_REMOVE_CATALOG_NAME)
|
||||
_assert_exact_remove_function(current)
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
return current
|
||||
|
||||
|
||||
async def _observe_current_async(db) -> lancedb.Function:
|
||||
current = await db.functions.get(_REMOVE_CATALOG_NAME)
|
||||
_assert_exact_remove_function(current)
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
return current
|
||||
|
||||
|
||||
def test_native_connection_exposes_private_remove_function_name():
|
||||
_assert_native_remove_method_present()
|
||||
|
||||
|
||||
def test_sync_remote_remove_exact_body_path_request_id_returns_none():
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
remove_attempts: list[dict[str, Any]] = []
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
body = json.loads(payload.decode("utf-8"))
|
||||
remove_attempts.append(
|
||||
{
|
||||
"request": request,
|
||||
"raw": payload,
|
||||
"body": body,
|
||||
"request_id": request.headers.get("x-request-id"),
|
||||
}
|
||||
)
|
||||
# Illegal body on 204 must be ignored; success is status-driven only.
|
||||
request.send_response(204)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
"message": _REMOVE_SERVER_MESSAGE_MARKER,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
|
||||
result = db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert result is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 1
|
||||
assert len(remove_attempts) == 1
|
||||
attempt = remove_attempts[0]
|
||||
_assert_exact_remove_request(
|
||||
attempt["request"],
|
||||
attempt["raw"],
|
||||
attempt["body"],
|
||||
expected_id=current.id,
|
||||
)
|
||||
assert attempt["body"]["expected_current_function_id"] == current.id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_remove_exact_body_returns_none():
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
seen["request"] = request
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
seen["request_id"] = request.headers.get("x-request-id")
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
async with _mock_remote_db_async(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
current = await _observe_current_async(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
result = await db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert result is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 1
|
||||
assert seen.get("raw")
|
||||
_assert_exact_remove_request(
|
||||
seen["request"],
|
||||
seen["raw"],
|
||||
seen["body"],
|
||||
expected_id=current.id,
|
||||
)
|
||||
|
||||
|
||||
def test_after_remove_name_lookup_not_found_id_lookup_same_function():
|
||||
"""Catalog-pointer SDK sequence via a stateful fixture; not server atomicity."""
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {
|
||||
"lookup_name": 0,
|
||||
"lookup_id": 0,
|
||||
"remove": 0,
|
||||
}
|
||||
removed = {"yes": False}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
if "name" in body:
|
||||
counters["lookup_name"] += 1
|
||||
assert body == {"name": _REMOVE_CATALOG_NAME}
|
||||
if removed["yes"]:
|
||||
request.send_response(404)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
"error_code": "name_or_function_not_found",
|
||||
"message": _REMOVE_SERVER_MESSAGE_MARKER,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
return
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
|
||||
counters["lookup_id"] += 1
|
||||
assert body == {"function_id": _REMOVE_FUNCTION_ID}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
_assert_exact_remove_request(
|
||||
request, raw, body, expected_id=_REMOVE_FUNCTION_ID
|
||||
)
|
||||
removed["yes"] = True
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(handler) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup_name"] == 1
|
||||
assert counters["lookup_id"] == 0
|
||||
assert counters["remove"] == 0
|
||||
|
||||
result = db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
assert result is None
|
||||
assert counters["lookup_name"] == 1
|
||||
assert counters["remove"] == 1
|
||||
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.get(_REMOVE_CATALOG_NAME)
|
||||
err = exc_info.value
|
||||
assert err.code == "name_or_function_not_found"
|
||||
_assert_payload_free(err)
|
||||
|
||||
by_id = db.functions.get_by_id(_REMOVE_FUNCTION_ID)
|
||||
|
||||
assert counters["lookup_name"] == 2
|
||||
assert counters["lookup_id"] == 1
|
||||
assert counters["remove"] == 1
|
||||
_assert_exact_remove_function(by_id)
|
||||
assert by_id.id == current.id
|
||||
assert by_id.parameters == current.parameters
|
||||
assert by_id.output_type == current.output_type
|
||||
assert by_id.output_nullable is current.output_nullable
|
||||
|
||||
|
||||
def test_explicit_name_conflict_is_function_error_payload_free():
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
body = {
|
||||
"error_code": "name_conflict",
|
||||
"message": (
|
||||
f"{_REMOVE_SERVER_MESSAGE_MARKER} looks_like {_CONFLICTING_MESSAGE_CODE}"
|
||||
),
|
||||
"name": _REMOVE_CATALOG_NAME,
|
||||
"function_id": _REMOVE_FUNCTION_ID,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
parsed = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_remove_request(
|
||||
request, payload, parsed, expected_id=_REMOVE_FUNCTION_ID
|
||||
)
|
||||
request.send_response(409)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, function_error)
|
||||
assert err.code == "name_conflict"
|
||||
assert err.code != _CONFLICTING_MESSAGE_CODE
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,status,response_body",
|
||||
[
|
||||
(
|
||||
"200_with_body",
|
||||
200,
|
||||
{
|
||||
"ok": True,
|
||||
"message": _REMOVE_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
"job_id": "must-not-infer-job",
|
||||
},
|
||||
),
|
||||
(
|
||||
"202_empty",
|
||||
202,
|
||||
f"{_REMOVE_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
|
||||
),
|
||||
("200_empty", 200, ""),
|
||||
],
|
||||
)
|
||||
def test_http_200_202_cannot_return_success(
|
||||
label: str, status: int, response_body: object
|
||||
):
|
||||
del label
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
parsed = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_remove_request(
|
||||
request, payload, parsed, expected_id=_REMOVE_FUNCTION_ID
|
||||
)
|
||||
request.send_response(status)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
if isinstance(response_body, str):
|
||||
request.wfile.write(response_body.encode("utf-8"))
|
||||
else:
|
||||
request.wfile.write(json.dumps(response_body).encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert not isinstance(err, _function_error_cls(required=False))
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
def test_empty_name_rejects_before_remove_transport():
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.remove("", current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_current",
|
||||
[
|
||||
_REMOVE_FUNCTION_ID,
|
||||
{"id": _REMOVE_FUNCTION_ID},
|
||||
object(),
|
||||
123,
|
||||
],
|
||||
)
|
||||
def test_raw_id_or_arbitrary_current_rejected_without_remove(bad_current):
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
# Observe a real handle separately so the bad-current path is isolated.
|
||||
_ = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.remove(_REMOVE_CATALOG_NAME, bad_current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
|
||||
|
||||
def test_local_sync_remove_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
_close_db(db)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_remove_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("keyword", _DELETED_REMOVE_KEYWORDS)
|
||||
def test_remove_rejects_deleted_cas_retry_version_kwargs(keyword):
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.remove(
|
||||
_REMOVE_CATALOG_NAME,
|
||||
current,
|
||||
**{keyword: True},
|
||||
)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
|
||||
|
||||
def test_no_direct_remove_methods_and_function_has_no_remove_facade_private():
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
assert callable(getattr(db.functions, "remove", None))
|
||||
assert not hasattr(lancedb, "_SyncFunctions")
|
||||
assert not hasattr(lancedb, "_AsyncFunctions")
|
||||
assert type(db.functions).__name__.startswith("_")
|
||||
assert "_SyncFunctions" not in getattr(lancedb, "__all__", [])
|
||||
assert "_AsyncFunctions" not in getattr(lancedb, "__all__", [])
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
@@ -0,0 +1,579 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python conditional first-class Function replacement."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from datetime import timedelta
|
||||
from typing import Any, Callable
|
||||
from unittest import mock
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
import lancedb._udf as _udf_mod
|
||||
import lancedb.job
|
||||
from lancedb import FunctionCapability, udf
|
||||
from lancedb.exceptions import JobFailedError
|
||||
|
||||
_SOURCE_MARKER = "replace-source-marker-unique-xyz"
|
||||
_SECRET_REFERENCE = "secret://team/replace-redact-token-xyz"
|
||||
_SECRET_ENV = "REPLACE_API_TOKEN"
|
||||
_NETWORK_ORIGIN = "https://api.replace-example.com"
|
||||
_FUNCTION_NAME = "text.normalize"
|
||||
_CURRENT_FUNCTION_ID = "fn.replace-current-1"
|
||||
_REPLACED_FUNCTION_ID = "fn.replace-result-1"
|
||||
_JOB_ID_SYNC = "job-replace-sync-1"
|
||||
_JOB_ID_ASYNC = "job-replace-async-1"
|
||||
_JOB_ID_CONFLICT = "job-replace-conflict-1"
|
||||
_REGISTER_PATH = "/v1/functions/register"
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_DESCRIBE_PATH = "/v1/jobs/describe"
|
||||
_CONFLICTING_MESSAGE_CODE = "definition_validation_failure"
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
|
||||
# job-result tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_DELETED_REPLACE_KEYWORDS = (
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"version",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"replace",
|
||||
"expected_current_function_id",
|
||||
"alias",
|
||||
"lineage",
|
||||
)
|
||||
|
||||
_SPEC_KEYS = {
|
||||
"format_version",
|
||||
"name",
|
||||
"definition",
|
||||
"expected_current_function_id",
|
||||
}
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"text": pa.string(), "limit": pa.int32()},
|
||||
output=pa.string(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=True,
|
||||
capabilities=[
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
],
|
||||
)
|
||||
def packable_replace_normalize(text, limit):
|
||||
"""replace-source-marker-unique-xyz."""
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def _current_function_wire() -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": _CURRENT_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _definition_json(fn: object) -> dict[str, Any]:
|
||||
payload = _udf_mod._build_function_definition(fn)._to_json()
|
||||
if isinstance(payload, bytes):
|
||||
return json.loads(payload.decode("utf-8"))
|
||||
assert isinstance(payload, str)
|
||||
return json.loads(payload)
|
||||
|
||||
|
||||
def _expected_replace_spec(name: str, current_id: str, fn: object) -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"name": name,
|
||||
"definition": _definition_json(fn),
|
||||
"expected_current_function_id": current_id,
|
||||
}
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _assert_exact_replace_spec(
|
||||
body: dict[str, Any], expected: dict[str, Any], current_id: str
|
||||
) -> None:
|
||||
assert set(body) == _SPEC_KEYS
|
||||
assert body == expected
|
||||
assert body["format_version"] == 1
|
||||
assert body["expected_current_function_id"] == current_id
|
||||
assert body["expected_current_function_id"] is not None
|
||||
assert _SOURCE_MARKER in json.dumps(body["definition"])
|
||||
assert any(
|
||||
capability.get("reference") == _SECRET_REFERENCE
|
||||
for capability in body["definition"]["capabilities"]
|
||||
)
|
||||
|
||||
|
||||
def _lookup_success_handler(
|
||||
counters: dict[str, int],
|
||||
*,
|
||||
after_lookup: (
|
||||
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
|
||||
) = None,
|
||||
):
|
||||
"""Serve exact lookup; optionally continue for register/describe."""
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
counters["lookup"] = counters.get("lookup", 0) + 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body == {"name": _FUNCTION_NAME}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps({"function": _current_function_wire()}).encode("utf-8")
|
||||
)
|
||||
return
|
||||
if after_lookup is not None:
|
||||
after_lookup(request, raw)
|
||||
return
|
||||
counters["register"] = counters.get("register", 0) + 1
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected register")
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _observe_current(db) -> lancedb.Function:
|
||||
current = db.functions.get(_FUNCTION_NAME)
|
||||
assert type(current) is lancedb.Function
|
||||
assert current.id == _CURRENT_FUNCTION_ID
|
||||
assert not hasattr(current, "name")
|
||||
assert not hasattr(current, "replace")
|
||||
return current
|
||||
|
||||
|
||||
async def _observe_current_async(db) -> lancedb.Function:
|
||||
current = await db.functions.get(_FUNCTION_NAME)
|
||||
assert type(current) is lancedb.Function
|
||||
assert current.id == _CURRENT_FUNCTION_ID
|
||||
assert not hasattr(current, "name")
|
||||
assert not hasattr(current, "replace")
|
||||
return current
|
||||
|
||||
|
||||
def test_sync_remote_replace_exact_body_one_package_job_and_function_result():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
|
||||
expected_spec = _expected_replace_spec(
|
||||
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
|
||||
)
|
||||
register_attempts: list[dict[str, Any]] = []
|
||||
function_result_wire = {
|
||||
"kind": "function",
|
||||
"format_version": 1,
|
||||
"function": {
|
||||
"format_version": 1,
|
||||
"id": _REPLACED_FUNCTION_ID,
|
||||
"signature": expected_spec["definition"]["signature"],
|
||||
},
|
||||
}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
if request.path == _REGISTER_PATH:
|
||||
counters["register"] += 1
|
||||
register_attempts.append(
|
||||
{
|
||||
"request_id": request.headers.get("x-request-id"),
|
||||
"raw": payload,
|
||||
"body": json.loads(payload.decode("utf-8")),
|
||||
}
|
||||
)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps({"job_id": _JOB_ID_SYNC}).encode("utf-8"))
|
||||
return
|
||||
|
||||
assert request.path == _DESCRIBE_PATH
|
||||
counters["describe"] += 1
|
||||
body = json.loads(payload.decode("utf-8"))
|
||||
assert body["job_id"] == _JOB_ID_SYNC
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
"job_id": _JOB_ID_SYNC,
|
||||
"job_state": "DONE",
|
||||
"job_type": "register_function",
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
"result": function_result_wire,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
package_calls = {"n": 0}
|
||||
original_package = _udf_mod._package_udf
|
||||
|
||||
def counting_package(fn: object):
|
||||
package_calls["n"] += 1
|
||||
return original_package(fn)
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "replace_function")
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
|
||||
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
|
||||
job = db.functions.replace(
|
||||
_FUNCTION_NAME, current, packable_replace_normalize
|
||||
)
|
||||
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_SYNC
|
||||
waited = job.wait(timeout=timedelta(seconds=5))
|
||||
|
||||
assert package_calls["n"] == 1
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 1
|
||||
assert counters["describe"] == 1
|
||||
assert len(register_attempts) == 1
|
||||
attempt = register_attempts[0]
|
||||
assert isinstance(attempt["request_id"], str) and attempt["request_id"]
|
||||
assert attempt["raw"]
|
||||
_assert_exact_replace_spec(attempt["body"], expected_spec, current.id)
|
||||
assert type(waited) is lancedb.Function
|
||||
assert waited.id == _REPLACED_FUNCTION_ID
|
||||
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
|
||||
assert waited.output_type == pa.string()
|
||||
assert waited.output_nullable is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_replace_exact_body_returns_async_job():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
expected_spec = _expected_replace_spec(
|
||||
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
|
||||
)
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
|
||||
assert request.path == _REGISTER_PATH
|
||||
counters["register"] += 1
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
seen["request_id"] = request.headers.get("x-request-id")
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
|
||||
|
||||
async with _mock_remote_db_async(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "replace_function")
|
||||
current = await _observe_current_async(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
job = await db.functions.replace(
|
||||
_FUNCTION_NAME, current, packable_replace_normalize
|
||||
)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 1
|
||||
assert seen.get("raw")
|
||||
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
|
||||
_assert_exact_replace_spec(seen["body"], expected_spec, current.id)
|
||||
assert type(job) is lancedb.job.AsyncJob
|
||||
assert job.id == _JOB_ID_ASYNC
|
||||
|
||||
|
||||
def test_sync_remote_replace_failed_name_conflict_raises_job_failed_error_code():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
|
||||
|
||||
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
|
||||
if request.path == _REGISTER_PATH:
|
||||
counters["register"] += 1
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps({"job_id": _JOB_ID_CONFLICT}).encode("utf-8")
|
||||
)
|
||||
return
|
||||
|
||||
assert request.path == _DESCRIBE_PATH
|
||||
counters["describe"] += 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body["job_id"] == _JOB_ID_CONFLICT
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
"job_id": _JOB_ID_CONFLICT,
|
||||
"job_type": "register_function",
|
||||
"job_state": "FAILED",
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
"failure": {
|
||||
"phase": "validate",
|
||||
"message": (
|
||||
f"looks like {_CONFLICTING_MESSAGE_CODE} during CAS"
|
||||
),
|
||||
"retryable": False,
|
||||
"error_code": "name_conflict",
|
||||
},
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
job = db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
|
||||
assert type(job) is lancedb.job.Job
|
||||
with pytest.raises(JobFailedError) as exc_info:
|
||||
job.wait(timeout=timedelta(seconds=5))
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 1
|
||||
assert counters["describe"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, JobFailedError)
|
||||
assert err.error_code == "name_conflict"
|
||||
assert err.error_code != _CONFLICTING_MESSAGE_CODE
|
||||
|
||||
|
||||
def test_empty_name_rejects_before_register_transport():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.replace("", current, packable_replace_normalize)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_current",
|
||||
[
|
||||
_CURRENT_FUNCTION_ID,
|
||||
{"id": _CURRENT_FUNCTION_ID},
|
||||
object(),
|
||||
123,
|
||||
],
|
||||
)
|
||||
def test_raw_id_or_arbitrary_current_rejected_without_register(bad_current):
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
# Observe a real handle separately so the bad-current path is isolated.
|
||||
_ = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.replace(
|
||||
_FUNCTION_NAME, bad_current, packable_replace_normalize
|
||||
)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
|
||||
|
||||
def test_local_sync_replace_not_implemented_without_table_mutation(tmp_path):
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "replace_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_replace_not_implemented_without_table_mutation(tmp_path):
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "replace_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
|
||||
|
||||
@pytest.mark.parametrize("keyword", _DELETED_REPLACE_KEYWORDS)
|
||||
def test_replace_rejects_deleted_cas_retry_version_kwargs(keyword):
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.replace(
|
||||
_FUNCTION_NAME,
|
||||
current,
|
||||
packable_replace_normalize,
|
||||
**{keyword: True},
|
||||
)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
|
||||
|
||||
def test_no_direct_replace_function_methods_and_function_has_no_replace():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert not hasattr(db, "replace_function")
|
||||
assert not hasattr(db, "register_function")
|
||||
assert not hasattr(current, "replace")
|
||||
assert not hasattr(current, "replace_function")
|
||||
assert callable(getattr(db.functions, "replace", None))
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
@@ -0,0 +1,728 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python exact first-class Function revocation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.remote.errors import HttpError
|
||||
|
||||
_REVOKE_PATH = "/v1/functions/revoke"
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_REVOKE_CATALOG_NAME = "text.normalize.revoke-name"
|
||||
_REVOKE_FUNCTION_ID = "fn.exact.revoke-handle"
|
||||
_REVOKE_SERVER_MESSAGE_MARKER = (
|
||||
"SERVER_REVOKE_DIAGNOSTIC_MARKER id=fn.exact.revoke-handle "
|
||||
"name=text.normalize.revoke-name"
|
||||
)
|
||||
_SENSITIVE_BODY_MARKER = "SENSITIVE_REVOKE_BODY_MARKER"
|
||||
_CONFLICTING_MESSAGE_CODE = "revoked_function"
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
|
||||
# remove tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_DELETED_REVOKE_KEYWORDS = (
|
||||
"function_id",
|
||||
"name",
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"version",
|
||||
"reason",
|
||||
"expiry",
|
||||
"force",
|
||||
"if_exists",
|
||||
"remove",
|
||||
"delete",
|
||||
)
|
||||
|
||||
|
||||
def _function_error_cls(*, required: bool = True):
|
||||
"""Resolve FunctionError from the live module (records RED when absent)."""
|
||||
from lancedb import exceptions as exc_mod
|
||||
|
||||
cls = getattr(exc_mod, "FunctionError", None)
|
||||
if cls is None:
|
||||
if required:
|
||||
pytest.fail("lancedb.exceptions.FunctionError is missing")
|
||||
return type("MissingFunctionError", (), {})
|
||||
return cls
|
||||
|
||||
|
||||
def _sample_function_wire() -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": _REVOKE_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _lookup_success_body() -> bytes:
|
||||
return json.dumps({"function": _sample_function_wire()}).encode("utf-8")
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
def _close_db(db: Any) -> None:
|
||||
with contextlib.suppress(Exception):
|
||||
inner = getattr(db, "_conn", None)
|
||||
if inner is not None:
|
||||
inner.close()
|
||||
return
|
||||
close = getattr(db, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
db = None
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
if db is not None:
|
||||
_close_db(db)
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
db = None
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
if db is not None:
|
||||
_close_db(db)
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_chain_text(exc: BaseException) -> str:
|
||||
parts: list[str] = []
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(str(current))
|
||||
parts.append(repr(current))
|
||||
current = current.__cause__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _assert_payload_free(exc: BaseException) -> None:
|
||||
text = _exception_chain_text(exc)
|
||||
assert _REVOKE_SERVER_MESSAGE_MARKER not in text
|
||||
assert _REVOKE_CATALOG_NAME not in text
|
||||
assert _REVOKE_FUNCTION_ID not in text
|
||||
assert _SENSITIVE_BODY_MARKER not in text
|
||||
|
||||
|
||||
def _assert_exact_revoke_function(function: object) -> None:
|
||||
assert type(function) is lancedb.Function
|
||||
assert function.id == _REVOKE_FUNCTION_ID
|
||||
assert not hasattr(function, "name")
|
||||
assert function.parameters == (
|
||||
("text", pa.string()),
|
||||
("limit", pa.int32()),
|
||||
)
|
||||
assert function.output_type == pa.string()
|
||||
assert function.output_nullable is True
|
||||
assert _REVOKE_CATALOG_NAME not in repr(function)
|
||||
assert _REVOKE_CATALOG_NAME not in str(function)
|
||||
|
||||
|
||||
def _assert_exact_revoke_request(
|
||||
request: http.server.BaseHTTPRequestHandler,
|
||||
raw: bytes,
|
||||
body: dict[str, Any],
|
||||
*,
|
||||
expected_id: str,
|
||||
) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _REVOKE_PATH
|
||||
assert "?" not in request.path
|
||||
assert "remove" not in request.path
|
||||
assert raw
|
||||
assert body == {"function_id": expected_id}
|
||||
assert set(body) == {"function_id"}
|
||||
assert "name" not in body
|
||||
assert "expected_current_function_id" not in body
|
||||
assert "format_version" not in body
|
||||
assert "function" not in body
|
||||
assert "signature" not in body
|
||||
assert "job_id" not in body
|
||||
assert "idempotency_key" not in body
|
||||
assert "user_version" not in body
|
||||
assert "reason" not in body
|
||||
assert "expiry" not in body
|
||||
assert "force" not in body
|
||||
request_id = request.headers.get("x-request-id")
|
||||
assert isinstance(request_id, str) and request_id
|
||||
|
||||
|
||||
def _assert_native_revoke_method_present() -> None:
|
||||
assert hasattr(_native.Connection, "_revoke_function")
|
||||
assert callable(getattr(_native.Connection, "_revoke_function"))
|
||||
|
||||
|
||||
def _lookup_success_handler(
|
||||
counters: dict[str, int],
|
||||
*,
|
||||
after_lookup: (
|
||||
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
|
||||
) = None,
|
||||
):
|
||||
"""Serve exact name lookup; optionally continue for revoke."""
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
counters["lookup"] = counters.get("lookup", 0) + 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body == {"name": _REVOKE_CATALOG_NAME}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
if after_lookup is not None:
|
||||
after_lookup(request, raw)
|
||||
return
|
||||
counters["revoke"] = counters.get("revoke", 0) + 1
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected revoke")
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _observe_current(db) -> lancedb.Function:
|
||||
current = db.functions.get(_REVOKE_CATALOG_NAME)
|
||||
_assert_exact_revoke_function(current)
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
return current
|
||||
|
||||
|
||||
async def _observe_current_async(db) -> lancedb.Function:
|
||||
current = await db.functions.get(_REVOKE_CATALOG_NAME)
|
||||
_assert_exact_revoke_function(current)
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
return current
|
||||
|
||||
|
||||
def test_native_connection_exposes_private_revoke_function():
|
||||
_assert_native_revoke_method_present()
|
||||
|
||||
|
||||
def test_sync_remote_revoke_exact_body_path_request_id_returns_none():
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
revoke_attempts: list[dict[str, Any]] = []
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
body = json.loads(payload.decode("utf-8"))
|
||||
revoke_attempts.append(
|
||||
{
|
||||
"request": request,
|
||||
"raw": payload,
|
||||
"body": body,
|
||||
"request_id": request.headers.get("x-request-id"),
|
||||
}
|
||||
)
|
||||
# Illegal body on 204 must be ignored; success is status-driven only.
|
||||
request.send_response(204)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
"message": _REVOKE_SERVER_MESSAGE_MARKER,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "revoke_function")
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
|
||||
result = db.functions.revoke(current)
|
||||
|
||||
assert result is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
assert len(revoke_attempts) == 1
|
||||
attempt = revoke_attempts[0]
|
||||
_assert_exact_revoke_request(
|
||||
attempt["request"],
|
||||
attempt["raw"],
|
||||
attempt["body"],
|
||||
expected_id=current.id,
|
||||
)
|
||||
assert attempt["body"]["function_id"] == current.id
|
||||
_assert_exact_revoke_function(current)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_revoke_exact_body_returns_none():
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
seen["request"] = request
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
seen["request_id"] = request.headers.get("x-request-id")
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
async with _mock_remote_db_async(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "revoke_function")
|
||||
current = await _observe_current_async(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
result = await db.functions.revoke(current)
|
||||
|
||||
assert result is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
assert seen.get("raw")
|
||||
_assert_exact_revoke_request(
|
||||
seen["request"],
|
||||
seen["raw"],
|
||||
seen["body"],
|
||||
expected_id=current.id,
|
||||
)
|
||||
|
||||
|
||||
def test_repeated_remote_revoke_204_both_return_none():
|
||||
"""Two logical calls each receiving 204 both succeed (Python outcome only)."""
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
revoke_request_ids: list[str] = []
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
body = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_revoke_request(
|
||||
request, payload, body, expected_id=_REVOKE_FUNCTION_ID
|
||||
)
|
||||
request_id = request.headers.get("x-request-id")
|
||||
assert isinstance(request_id, str) and request_id
|
||||
revoke_request_ids.append(request_id)
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
|
||||
first = db.functions.revoke(current)
|
||||
second = db.functions.revoke(current)
|
||||
|
||||
assert first is None
|
||||
assert second is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 2
|
||||
assert len(revoke_request_ids) == 2
|
||||
_assert_exact_revoke_function(current)
|
||||
|
||||
|
||||
def test_after_revoke_name_and_id_lookup_still_return_function():
|
||||
"""Revoke does not unlink names; SDK-visible sequence only, not Sophon proof."""
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {
|
||||
"lookup_name": 0,
|
||||
"lookup_id": 0,
|
||||
"revoke": 0,
|
||||
}
|
||||
revoked = {"yes": False}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
if "name" in body:
|
||||
counters["lookup_name"] += 1
|
||||
assert body == {"name": _REVOKE_CATALOG_NAME}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
|
||||
counters["lookup_id"] += 1
|
||||
assert body == {"function_id": _REVOKE_FUNCTION_ID}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
_assert_exact_revoke_request(
|
||||
request, raw, body, expected_id=_REVOKE_FUNCTION_ID
|
||||
)
|
||||
revoked["yes"] = True
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup_name"] == 1
|
||||
assert counters["lookup_id"] == 0
|
||||
assert counters["revoke"] == 0
|
||||
assert not revoked["yes"]
|
||||
|
||||
result = db.functions.revoke(current)
|
||||
assert result is None
|
||||
assert counters["lookup_name"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
assert revoked["yes"]
|
||||
|
||||
by_name = db.functions.get(_REVOKE_CATALOG_NAME)
|
||||
by_id = db.functions.get_by_id(_REVOKE_FUNCTION_ID)
|
||||
|
||||
assert counters["lookup_name"] == 2
|
||||
assert counters["lookup_id"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
_assert_exact_revoke_function(by_name)
|
||||
_assert_exact_revoke_function(by_id)
|
||||
assert by_name.id == current.id
|
||||
assert by_id.id == current.id
|
||||
assert by_name.parameters == current.parameters
|
||||
assert by_id.parameters == current.parameters
|
||||
assert by_name.output_type == current.output_type
|
||||
assert by_id.output_type == current.output_type
|
||||
assert by_name.output_nullable is current.output_nullable
|
||||
assert by_id.output_nullable is current.output_nullable
|
||||
|
||||
|
||||
def test_explicit_name_or_function_not_found_is_function_error_payload_free():
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
body = {
|
||||
"error_code": "name_or_function_not_found",
|
||||
"message": (
|
||||
f"{_REVOKE_SERVER_MESSAGE_MARKER} looks_like {_CONFLICTING_MESSAGE_CODE} "
|
||||
"name_conflict"
|
||||
),
|
||||
"name": _REVOKE_CATALOG_NAME,
|
||||
"function_id": _REVOKE_FUNCTION_ID,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
parsed = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_revoke_request(
|
||||
request, payload, parsed, expected_id=_REVOKE_FUNCTION_ID
|
||||
)
|
||||
request.send_response(404)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.revoke(current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, function_error)
|
||||
assert err.code == "name_or_function_not_found"
|
||||
assert err.code != _CONFLICTING_MESSAGE_CODE
|
||||
assert err.code != "name_conflict"
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,status,response_body",
|
||||
[
|
||||
(
|
||||
"200_with_body",
|
||||
200,
|
||||
{
|
||||
"ok": True,
|
||||
"message": _REVOKE_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
"job_id": "must-not-infer-job",
|
||||
},
|
||||
),
|
||||
(
|
||||
"202_empty",
|
||||
202,
|
||||
f"{_REVOKE_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
|
||||
),
|
||||
("200_empty", 200, ""),
|
||||
],
|
||||
)
|
||||
def test_http_200_202_cannot_return_success(
|
||||
label: str, status: int, response_body: object
|
||||
):
|
||||
del label
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
parsed = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_revoke_request(
|
||||
request, payload, parsed, expected_id=_REVOKE_FUNCTION_ID
|
||||
)
|
||||
request.send_response(status)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
if isinstance(response_body, str):
|
||||
request.wfile.write(response_body.encode("utf-8"))
|
||||
else:
|
||||
request.wfile.write(json.dumps(response_body).encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.revoke(current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert not isinstance(err, _function_error_cls(required=False))
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_function",
|
||||
[
|
||||
_REVOKE_FUNCTION_ID,
|
||||
{"id": _REVOKE_FUNCTION_ID},
|
||||
object(),
|
||||
123,
|
||||
],
|
||||
)
|
||||
def test_raw_id_or_arbitrary_function_rejected_without_revoke(bad_function):
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
# Observe a real handle separately so the bad-function path is isolated.
|
||||
_ = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.revoke(bad_function)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
|
||||
|
||||
def test_local_sync_revoke_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "revoke_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.revoke(current)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
_close_db(db)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_revoke_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "revoke_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.revoke(current)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("keyword", _DELETED_REVOKE_KEYWORDS)
|
||||
def test_revoke_rejects_overdesigned_kwargs(keyword):
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.revoke(current, **{keyword: True})
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
|
||||
|
||||
def test_no_direct_revoke_methods_and_function_has_no_revoke_facade_private():
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert not hasattr(db, "revoke_function")
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
assert callable(getattr(db.functions, "revoke", None))
|
||||
assert not hasattr(lancedb, "_SyncFunctions")
|
||||
assert not hasattr(lancedb, "_AsyncFunctions")
|
||||
assert type(db.functions).__name__.startswith("_")
|
||||
assert "_SyncFunctions" not in getattr(lancedb, "__all__", [])
|
||||
assert "_AsyncFunctions" not in getattr(lancedb, "__all__", [])
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,899 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python ``table.add_generated_column`` (FF-032).
|
||||
|
||||
Public user shape under test:
|
||||
|
||||
job = table.add_generated_column(
|
||||
"normalized_text",
|
||||
normalize(text=col("text")),
|
||||
)
|
||||
job.wait()
|
||||
|
||||
These tests exercise the live worktree PyO3 extension and public sync/async
|
||||
wrappers. While the public methods and hidden native bridge are absent they
|
||||
fail against that extension; once present they freeze the public contract
|
||||
below. They must not fake success paths.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import inspect
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
import lancedb.job
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.expr import col
|
||||
from lancedb.remote.table import RemoteTable
|
||||
from lancedb.table import AsyncTable, LanceTable, Table
|
||||
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_JOB_DESCRIBE_PATH = "/v1/jobs/describe"
|
||||
_TABLE_NAME = "articles"
|
||||
_DESCRIBE_PATH = f"/v1/table/{_TABLE_NAME}/describe/"
|
||||
_CREATE_PATH = f"/v1/table/{_TABLE_NAME}/generated_columns/create/"
|
||||
_BRANCHES_CREATE_PATH = f"/v1/table/{_TABLE_NAME}/branches/create/"
|
||||
_BRANCHES_LIST_PATH = f"/v1/table/{_TABLE_NAME}/branches/list/"
|
||||
|
||||
_CATALOG_NAME = "text.normalize"
|
||||
_FUNCTION_ID = "fn.exact.normalize.gen-col"
|
||||
_JOB_ID_SYNC = "job-create-gen-col-sync-1"
|
||||
_JOB_ID_ASYNC = "job-create-gen-col-async-1"
|
||||
_JOB_ID_BRANCH = "job-create-gen-col-branch-1"
|
||||
_SOURCE_TABLE_VERSION = 42
|
||||
_TEXT_FIELD_ID = 7
|
||||
_BRANCH_NAME = "exp"
|
||||
_BRANCH_SOURCE_VERSION = 9
|
||||
_BRANCH_TEXT_FIELD_ID = 11
|
||||
|
||||
_DESCRIBE_BODY_MARKER = "SENSITIVE_DESCRIBE_BODY_MARKER_gen_col_xyz"
|
||||
_CREATE_RESPONSE_MARKER = "SENSITIVE_CREATE_RESPONSE_MARKER_gen_col_xyz"
|
||||
_LITERAL_PAYLOAD_SENTINEL = "LITERAL_PAYLOAD_SENTINEL_gen_col_xyz"
|
||||
|
||||
# Pinned Rust-canonical schema-only Utf8 type IPC (base64), shared with FF-028.
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_FORBIDDEN_PUBLIC_NAMES = (
|
||||
"FunctionCall",
|
||||
"BoundFunctionCall",
|
||||
"AuthoredFunctionCall",
|
||||
"CreateGeneratedColumnRequest",
|
||||
"CreateGeneratedColumnJobSpec",
|
||||
"GeneratedColumnBindingSnapshot",
|
||||
"GeneratedColumnCreateRequest",
|
||||
"geneva",
|
||||
"GenevaFunction",
|
||||
"VirtualColumnDefinition",
|
||||
)
|
||||
|
||||
_FORBIDDEN_METHOD_KWARGS = (
|
||||
"source_table_version",
|
||||
"version",
|
||||
"field_id",
|
||||
"field_ids",
|
||||
"output",
|
||||
"output_type",
|
||||
"output_nullable",
|
||||
"nullable",
|
||||
"spec",
|
||||
"retry_key",
|
||||
"idempotency_key",
|
||||
"request",
|
||||
"envelope",
|
||||
"table_ref",
|
||||
"branch",
|
||||
)
|
||||
|
||||
|
||||
def _sample_function_wire(
|
||||
*,
|
||||
function_id: str = _FUNCTION_ID,
|
||||
parameters: list[dict[str, str]] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": function_id,
|
||||
"signature": {
|
||||
"parameters": parameters
|
||||
or [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _text_schema_fields(
|
||||
*, arrow_type: str = "string", nullable: bool = True
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"fields": [
|
||||
{
|
||||
"name": "text",
|
||||
"type": {"type": arrow_type},
|
||||
"nullable": nullable,
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def _describe_body(
|
||||
*,
|
||||
version: int = _SOURCE_TABLE_VERSION,
|
||||
field_ids: list[int] | None = None,
|
||||
arrow_type: str = "string",
|
||||
include_marker: bool = True,
|
||||
) -> dict[str, Any]:
|
||||
body: dict[str, Any] = {
|
||||
"version": version,
|
||||
"schema": _text_schema_fields(arrow_type=arrow_type),
|
||||
"field_ids": field_ids if field_ids is not None else [_TEXT_FIELD_ID],
|
||||
}
|
||||
if include_marker:
|
||||
body["server_diagnostic"] = _DESCRIBE_BODY_MARKER
|
||||
return body
|
||||
|
||||
|
||||
def _create_gen_column_done_body(job_id: str) -> dict[str, Any]:
|
||||
# DONE with omitted result: create_gen_column projects JobResult::None.
|
||||
return {
|
||||
"job_id": job_id,
|
||||
"job_state": "DONE",
|
||||
"job_type": "create_gen_column",
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
}
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_text(exc: BaseException) -> str:
|
||||
parts = [str(exc), repr(exc)]
|
||||
current: BaseException | None = exc
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(f"{type(current).__name__}: {current}")
|
||||
current = current.__cause__ or current.__context__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _json_response(
|
||||
request: http.server.BaseHTTPRequestHandler, body: dict[str, Any]
|
||||
) -> None:
|
||||
payload = json.dumps(body).encode("utf-8")
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(payload)
|
||||
|
||||
|
||||
def _lookup_function(db: Any) -> lancedb.Function:
|
||||
return db.functions.get(_CATALOG_NAME)
|
||||
|
||||
|
||||
class _RequestLog:
|
||||
"""Track lookup/describe/create after setup; setup traffic is excluded."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.lookup: list[dict[str, Any]] = []
|
||||
self.describe: list[dict[str, Any]] = []
|
||||
self.create: list[dict[str, Any]] = []
|
||||
self.other_table: list[str] = []
|
||||
self.recording = False
|
||||
|
||||
def start(self) -> None:
|
||||
# Drop setup's explicit Function lookup and open_table describe so
|
||||
# operation accounting cannot be polluted by fixture traffic.
|
||||
self.lookup.clear()
|
||||
self.describe.clear()
|
||||
self.create.clear()
|
||||
self.other_table.clear()
|
||||
self.recording = True
|
||||
|
||||
def note(self, path: str, body: dict[str, Any] | None = None) -> None:
|
||||
if not self.recording:
|
||||
return
|
||||
if path == _LOOKUP_PATH:
|
||||
self.lookup.append(body or {})
|
||||
elif path == _DESCRIBE_PATH:
|
||||
self.describe.append(body or {})
|
||||
elif path == _CREATE_PATH:
|
||||
self.create.append(body or {})
|
||||
elif path.startswith(f"/v1/table/{_TABLE_NAME}/"):
|
||||
self.other_table.append(path)
|
||||
|
||||
|
||||
def _assert_no_operation_traffic(log: _RequestLog) -> None:
|
||||
assert log.lookup == []
|
||||
assert log.describe == []
|
||||
assert log.create == []
|
||||
assert log.other_table == []
|
||||
|
||||
|
||||
def _assert_exact_public_signature(method: Any) -> None:
|
||||
"""Freeze ``(self, column_name, call)`` with no varargs/kwargs escape hatches."""
|
||||
params = list(inspect.signature(method).parameters.values())
|
||||
assert [p.name for p in params] == ["self", "column_name", "call"]
|
||||
for param in params:
|
||||
assert param.kind in (
|
||||
inspect.Parameter.POSITIONAL_ONLY,
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
)
|
||||
assert param.default is inspect.Parameter.empty
|
||||
assert param.kind is not inspect.Parameter.VAR_POSITIONAL
|
||||
assert param.kind is not inspect.Parameter.VAR_KEYWORD
|
||||
assert param.kind is not inspect.Parameter.KEYWORD_ONLY
|
||||
|
||||
|
||||
def _open_table_and_function(
|
||||
*,
|
||||
describe_body: dict[str, Any] | None = None,
|
||||
on_create: Callable[[dict[str, Any], http.server.BaseHTTPRequestHandler], None]
|
||||
| None = None,
|
||||
job_id: str = _JOB_ID_SYNC,
|
||||
support_branch_create: bool = False,
|
||||
function_wire: dict[str, Any] | None = None,
|
||||
):
|
||||
"""Open remote table + immutable Function; return (db, table, function, log, cm)."""
|
||||
log = _RequestLog()
|
||||
binding_describe = describe_body or _describe_body()
|
||||
open_describe = {
|
||||
"version": 1,
|
||||
"schema": _text_schema_fields(),
|
||||
}
|
||||
state = {"opened": False}
|
||||
wire = function_wire or _sample_function_wire()
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _LOOKUP_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"function": wire})
|
||||
return
|
||||
|
||||
if request.path == _JOB_DESCRIBE_PATH:
|
||||
assert body["job_id"] == job_id
|
||||
_json_response(request, _create_gen_column_done_body(job_id))
|
||||
return
|
||||
|
||||
if support_branch_create and request.path == _BRANCHES_CREATE_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {})
|
||||
return
|
||||
|
||||
if support_branch_create and request.path == _BRANCHES_LIST_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(
|
||||
request,
|
||||
{
|
||||
"branches": {
|
||||
_BRANCH_NAME: {
|
||||
"parentBranch": None,
|
||||
"parentVersion": 1,
|
||||
"createAt": 1,
|
||||
"manifestSize": 1,
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
# First describe seeds open_table; later ones are binding snapshots.
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding_describe)
|
||||
return
|
||||
|
||||
if request.path == _CREATE_PATH:
|
||||
log.note(request.path, body)
|
||||
if on_create is not None:
|
||||
on_create(body, request)
|
||||
return
|
||||
_json_response(
|
||||
request,
|
||||
{
|
||||
"job_id": job_id,
|
||||
"server_extra": {"marker": _CREATE_RESPONSE_MARKER},
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
if request.path.startswith(f"/v1/table/{_TABLE_NAME}/"):
|
||||
log.note(request.path, body)
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected path")
|
||||
|
||||
cm = _mock_remote_db(handler)
|
||||
db = cm.__enter__()
|
||||
function = _lookup_function(db)
|
||||
table = db.open_table(_TABLE_NAME)
|
||||
assert isinstance(table, RemoteTable)
|
||||
# open_table consumed the seed describe; binding/create accounting starts now.
|
||||
# Setup's one explicit lookup is cleared here and must not pollute counts.
|
||||
log.start()
|
||||
return db, table, function, log, cm
|
||||
|
||||
|
||||
def _assert_exact_create_envelope(
|
||||
body: dict[str, Any],
|
||||
*,
|
||||
source_table_version: int,
|
||||
column_name: str,
|
||||
field_id: int,
|
||||
branch: str | None = None,
|
||||
) -> None:
|
||||
expected_keys = {"source_table_version", "spec"}
|
||||
if branch is not None:
|
||||
expected_keys.add("branch")
|
||||
assert set(body) == expected_keys
|
||||
assert body["source_table_version"] == source_table_version
|
||||
assert "table_ref" not in body
|
||||
if branch is None:
|
||||
assert "branch" not in body
|
||||
else:
|
||||
assert body["branch"] == branch
|
||||
|
||||
spec = body["spec"]
|
||||
assert set(spec) == {"format_version", "column_name", "function_call"}
|
||||
assert spec["format_version"] == 1
|
||||
assert spec["column_name"] == column_name
|
||||
for forbidden in (
|
||||
"table_ref",
|
||||
"source_table_version",
|
||||
"version",
|
||||
"output",
|
||||
"output_type",
|
||||
"output_field_id",
|
||||
"dependency_epoch",
|
||||
"materialized_epoch",
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"name",
|
||||
"handle",
|
||||
"artifact",
|
||||
"geneva",
|
||||
):
|
||||
assert forbidden not in spec
|
||||
|
||||
call = spec["function_call"]
|
||||
assert set(call) == {"function_id", "arguments"}
|
||||
assert call["function_id"] == _FUNCTION_ID
|
||||
assert len(call["arguments"]) == 1
|
||||
binding = call["arguments"][0]
|
||||
assert binding["parameter"] == "text"
|
||||
value = binding["value"]
|
||||
assert value["kind"] == "field"
|
||||
assert value["field_id"] == field_id
|
||||
assert value["data_type_ipc"] == _UTF8_TYPE_IPC_B64
|
||||
assert "name" not in value
|
||||
assert "column_name" not in value
|
||||
assert "text" not in value
|
||||
# Serialized call must not late-bind by column name anywhere relevant.
|
||||
dumped = json.dumps(call)
|
||||
assert '"column_name"' not in dumped
|
||||
assert "normalized_text" not in dumped
|
||||
|
||||
|
||||
def test_public_and_native_add_generated_column_seams_must_exist():
|
||||
"""Public sync/async methods and the private native bridge must exist."""
|
||||
assert hasattr(_native.Table, "_add_generated_column"), (
|
||||
"native private bridge Table._add_generated_column is missing"
|
||||
)
|
||||
assert hasattr(AsyncTable, "add_generated_column"), (
|
||||
"AsyncTable.add_generated_column is missing"
|
||||
)
|
||||
assert hasattr(Table, "add_generated_column"), (
|
||||
"Table.add_generated_column is missing"
|
||||
)
|
||||
assert hasattr(LanceTable, "add_generated_column"), (
|
||||
"LanceTable.add_generated_column is missing"
|
||||
)
|
||||
assert hasattr(RemoteTable, "add_generated_column"), (
|
||||
"RemoteTable.add_generated_column is missing"
|
||||
)
|
||||
|
||||
# Once present, freeze the exact public positional surface.
|
||||
_assert_exact_public_signature(Table.add_generated_column)
|
||||
_assert_exact_public_signature(LanceTable.add_generated_column)
|
||||
_assert_exact_public_signature(RemoteTable.add_generated_column)
|
||||
_assert_exact_public_signature(AsyncTable.add_generated_column)
|
||||
|
||||
|
||||
def test_sync_remote_add_generated_column_returns_job_without_eager_wrapper_mutation():
|
||||
db, table, normalize, log, cm = _open_table_and_function(job_id=_JOB_ID_SYNC)
|
||||
try:
|
||||
# Capture public wrapper state before the operation window.
|
||||
schema_before = table.schema
|
||||
version_before = table.version
|
||||
log.start()
|
||||
|
||||
call = normalize(text=col("text"))
|
||||
# Exact public argument order from the frozen user example.
|
||||
job = table.add_generated_column(
|
||||
"normalized_text",
|
||||
call,
|
||||
)
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_SYNC
|
||||
|
||||
# Exact success path stops after submit: one binding describe, one create,
|
||||
# and no catalog re-lookup. Do not wait yet.
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.create) == 1
|
||||
assert log.other_table == []
|
||||
|
||||
# Public schema/version through the existing wrapper must still reflect
|
||||
# the pre-submit table: generated column is not published by Job accept.
|
||||
# Access both before wait so eager wrapper cache invalidation / refresh /
|
||||
# version advancement is observable.
|
||||
schema_after = table.schema
|
||||
assert "normalized_text" not in schema_after.names
|
||||
assert schema_after == schema_before
|
||||
# Schema must be served from the existing wrapper cache — no extra
|
||||
# describe beyond the one binding snapshot.
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.create) == 1
|
||||
|
||||
version_after = table.version
|
||||
assert version_after == version_before
|
||||
# Public Remote ``version`` always describes once by design; that probe
|
||||
# must not drag a schema-cache miss, create, or catalog lookup with it.
|
||||
assert len(log.describe) == 2
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.create) == 1
|
||||
assert log.other_table == []
|
||||
|
||||
waited = job.wait()
|
||||
assert waited is None
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.create) == 1
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_add_generated_column_returns_async_job_and_wait_none():
|
||||
log = _RequestLog()
|
||||
state = {"opened": False}
|
||||
binding_describe = _describe_body()
|
||||
open_describe = {"version": 1, "schema": _text_schema_fields()}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _LOOKUP_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"function": _sample_function_wire()})
|
||||
return
|
||||
if request.path == _JOB_DESCRIBE_PATH:
|
||||
assert body["job_id"] == _JOB_ID_ASYNC
|
||||
_json_response(request, _create_gen_column_done_body(_JOB_ID_ASYNC))
|
||||
return
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding_describe)
|
||||
return
|
||||
if request.path == _CREATE_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"job_id": _JOB_ID_ASYNC})
|
||||
return
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
|
||||
async with _mock_remote_db_async(handler) as db:
|
||||
normalize = await db.functions.get(_CATALOG_NAME)
|
||||
table = await db.open_table(_TABLE_NAME)
|
||||
log.start()
|
||||
call = normalize(text=col("text"))
|
||||
job = await table.add_generated_column("normalized_text", call)
|
||||
assert type(job) is lancedb.job.AsyncJob
|
||||
assert job.id == _JOB_ID_ASYNC
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.create) == 1
|
||||
waited = await job.wait()
|
||||
assert waited is None
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.create) == 1
|
||||
|
||||
|
||||
def test_remote_add_generated_column_one_describe_one_create_exact_envelope():
|
||||
db, table, normalize, log, cm = _open_table_and_function(job_id=_JOB_ID_SYNC)
|
||||
try:
|
||||
call = normalize(text=col("text"))
|
||||
job = table.add_generated_column("normalized_text", call)
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_SYNC
|
||||
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.create) == 1
|
||||
assert log.other_table == []
|
||||
_assert_exact_create_envelope(
|
||||
log.create[0],
|
||||
source_table_version=_SOURCE_TABLE_VERSION,
|
||||
column_name="normalized_text",
|
||||
field_id=_TEXT_FIELD_ID,
|
||||
)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_remote_branch_add_generated_column_includes_exact_branch_identity():
|
||||
branch_describe = _describe_body(
|
||||
version=_BRANCH_SOURCE_VERSION,
|
||||
field_ids=[_BRANCH_TEXT_FIELD_ID],
|
||||
)
|
||||
db, table, normalize, log, cm = _open_table_and_function(
|
||||
describe_body=branch_describe,
|
||||
job_id=_JOB_ID_BRANCH,
|
||||
support_branch_create=True,
|
||||
)
|
||||
try:
|
||||
branched = table.branches.create(_BRANCH_NAME)
|
||||
assert isinstance(branched, RemoteTable)
|
||||
assert branched.current_branch() == _BRANCH_NAME
|
||||
log.start()
|
||||
|
||||
call = normalize(text=col("text"))
|
||||
job = branched.add_generated_column("normalized_text", call)
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_BRANCH
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.create) == 1
|
||||
assert log.describe[0].get("branch") == _BRANCH_NAME
|
||||
_assert_exact_create_envelope(
|
||||
log.create[0],
|
||||
source_table_version=_BRANCH_SOURCE_VERSION,
|
||||
column_name="normalized_text",
|
||||
field_id=_BRANCH_TEXT_FIELD_ID,
|
||||
branch=_BRANCH_NAME,
|
||||
)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_empty_column_name_fails_locally_with_zero_table_requests():
|
||||
db, table, normalize, log, cm = _open_table_and_function()
|
||||
try:
|
||||
# Authored call owns a real literal so payload-free failure is not vacuous.
|
||||
call = normalize(text=_LITERAL_PAYLOAD_SENTINEL)
|
||||
with pytest.raises((ValueError, TypeError)) as raised:
|
||||
table.add_generated_column("", call)
|
||||
text = _exception_text(raised.value)
|
||||
lowered = text.lower()
|
||||
assert "column" in lowered or "empty" in lowered or "non-empty" in lowered
|
||||
assert _DESCRIBE_BODY_MARKER not in text
|
||||
assert _CREATE_RESPONSE_MARKER not in text
|
||||
assert _LITERAL_PAYLOAD_SENTINEL not in text
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("column_ref", "expected_token"),
|
||||
[
|
||||
("missing_text", "missing_text"),
|
||||
("Text", "Text"), # exact-case mismatch against schema field "text"
|
||||
],
|
||||
)
|
||||
def test_missing_or_case_mismatch_column_one_describe_zero_create(
|
||||
column_ref: str, expected_token: str
|
||||
):
|
||||
db, table, normalize, log, cm = _open_table_and_function()
|
||||
try:
|
||||
call = normalize(text=col(column_ref))
|
||||
with pytest.raises(ValueError) as raised:
|
||||
table.add_generated_column("normalized_text", call)
|
||||
text = _exception_text(raised.value)
|
||||
assert expected_token in text
|
||||
assert "text" in text # parameter name from the Function signature
|
||||
assert "missing" in text.lower() or "field" in text.lower()
|
||||
assert _DESCRIBE_BODY_MARKER not in text
|
||||
assert _CREATE_RESPONSE_MARKER not in text
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert log.create == []
|
||||
assert log.other_table == []
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_type_mismatch_one_describe_zero_create_identifies_parameter():
|
||||
db, table, normalize, log, cm = _open_table_and_function(
|
||||
describe_body=_describe_body(arrow_type="int32"),
|
||||
)
|
||||
try:
|
||||
call = normalize(text=col("text"))
|
||||
with pytest.raises(ValueError) as raised:
|
||||
table.add_generated_column("normalized_text", call)
|
||||
text = _exception_text(raised.value)
|
||||
assert "text" in text
|
||||
assert "type" in text.lower() or "mismatch" in text.lower()
|
||||
assert _DESCRIBE_BODY_MARKER not in text
|
||||
assert _CREATE_RESPONSE_MARKER not in text
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert log.create == []
|
||||
assert log.other_table == []
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_literal_payload_stays_out_of_field_binding_failure():
|
||||
"""Authored call owns a real literal; later field binding fails payload-free."""
|
||||
wire = _sample_function_wire(
|
||||
parameters=[
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "prefix", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
]
|
||||
)
|
||||
db, table, normalize, log, cm = _open_table_and_function(function_wire=wire)
|
||||
try:
|
||||
call = normalize(text=col("missing_text"), prefix=_LITERAL_PAYLOAD_SENTINEL)
|
||||
with pytest.raises(ValueError) as raised:
|
||||
table.add_generated_column("normalized_text", call)
|
||||
text = _exception_text(raised.value)
|
||||
assert "missing_text" in text
|
||||
assert _LITERAL_PAYLOAD_SENTINEL not in text
|
||||
assert _DESCRIBE_BODY_MARKER not in text
|
||||
assert _CREATE_RESPONSE_MARKER not in text
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert log.create == []
|
||||
assert log.other_table == []
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_closed_async_table_fails_with_zero_operation_requests():
|
||||
log = _RequestLog()
|
||||
state = {"opened": False}
|
||||
binding_describe = _describe_body()
|
||||
open_describe = {"version": 1, "schema": _text_schema_fields()}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _LOOKUP_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"function": _sample_function_wire()})
|
||||
return
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding_describe)
|
||||
return
|
||||
if request.path == _CREATE_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"job_id": _JOB_ID_ASYNC})
|
||||
return
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
|
||||
async with _mock_remote_db_async(handler) as db:
|
||||
normalize = await db.functions.get(_CATALOG_NAME)
|
||||
table = await db.open_table(_TABLE_NAME)
|
||||
call = normalize(text=col("text"))
|
||||
# Public close only — do not mutate private implementation fields.
|
||||
table.close()
|
||||
log.start()
|
||||
try:
|
||||
await table.add_generated_column("normalized_text", call)
|
||||
except AttributeError:
|
||||
# Method missing: re-raise so the failure names the public seam.
|
||||
raise
|
||||
except Exception as exc:
|
||||
text = _exception_text(exc)
|
||||
assert "closed" in text.lower()
|
||||
else:
|
||||
pytest.fail("closed AsyncTable must fail before transport")
|
||||
_assert_no_operation_traffic(log)
|
||||
|
||||
|
||||
def test_rejects_non_authored_call_before_any_operation_request():
|
||||
db, table, normalize, log, cm = _open_table_and_function()
|
||||
try:
|
||||
bad_values = (
|
||||
normalize, # exact Function handle itself
|
||||
{"text": "x"},
|
||||
col("text"), # direct query Expr
|
||||
object(),
|
||||
)
|
||||
for bad in bad_values:
|
||||
with pytest.raises(TypeError):
|
||||
table.add_generated_column("normalized_text", bad)
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_native_valid_call_returns_not_supported_without_mutation(tmp_path):
|
||||
# Immutable Function handle is connection-free; obtain it via remote lookup.
|
||||
def lookup_only(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
_json_response(request, {"function": _sample_function_wire()})
|
||||
|
||||
with _mock_remote_db(lookup_only) as remote_db:
|
||||
normalize = _lookup_function(remote_db)
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table(_TABLE_NAME, [{"text": "Hello"}, {"text": "World"}])
|
||||
assert isinstance(table, LanceTable)
|
||||
version_before = table.version
|
||||
schema_before = table.schema
|
||||
rows_before = table.to_arrow().to_pylist()
|
||||
call = normalize(text=col("text"))
|
||||
|
||||
with pytest.raises(NotImplementedError) as raised:
|
||||
table.add_generated_column("normalized_text", call)
|
||||
text = _exception_text(raised.value)
|
||||
assert "not supported" in text.lower() or "submit_create_generated_column" in text
|
||||
assert "add_columns" not in text.lower()
|
||||
|
||||
assert table.version == version_before
|
||||
assert table.schema == schema_before
|
||||
assert "normalized_text" not in table.schema.names
|
||||
assert table.to_arrow().to_pylist() == rows_before
|
||||
|
||||
|
||||
def test_public_surface_is_minimal_and_private_call_stays_opaque():
|
||||
for name in _FORBIDDEN_PUBLIC_NAMES:
|
||||
assert name not in getattr(lancedb, "__all__", [])
|
||||
assert not hasattr(lancedb, name)
|
||||
|
||||
assert not hasattr(lancedb, "_FunctionCall")
|
||||
authored_type = getattr(_native, "_FunctionCall", None)
|
||||
assert authored_type is not None
|
||||
with pytest.raises(TypeError):
|
||||
authored_type()
|
||||
|
||||
# When the public method exists, reject overdesign kwargs and keep the frozen
|
||||
# positional surface: (self, column_name, call).
|
||||
if hasattr(Table, "add_generated_column"):
|
||||
_assert_exact_public_signature(Table.add_generated_column)
|
||||
for keyword in _FORBIDDEN_METHOD_KWARGS:
|
||||
assert (
|
||||
keyword not in inspect.signature(Table.add_generated_column).parameters
|
||||
)
|
||||
db, table, normalize, log, cm = _open_table_and_function()
|
||||
try:
|
||||
call = normalize(text=col("text"))
|
||||
for keyword in _FORBIDDEN_METHOD_KWARGS:
|
||||
with pytest.raises(TypeError):
|
||||
table.add_generated_column(
|
||||
"normalized_text",
|
||||
call,
|
||||
**{keyword: object()},
|
||||
)
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
if hasattr(LanceTable, "add_generated_column"):
|
||||
_assert_exact_public_signature(LanceTable.add_generated_column)
|
||||
if hasattr(RemoteTable, "add_generated_column"):
|
||||
_assert_exact_public_signature(RemoteTable.add_generated_column)
|
||||
if hasattr(AsyncTable, "add_generated_column"):
|
||||
_assert_exact_public_signature(AsyncTable.add_generated_column)
|
||||
for keyword in _FORBIDDEN_METHOD_KWARGS:
|
||||
assert (
|
||||
keyword
|
||||
not in inspect.signature(AsyncTable.add_generated_column).parameters
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,672 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python ``table.generated_column_status`` (B3d2).
|
||||
|
||||
Public user shape under test:
|
||||
|
||||
status = table.generated_column_status("complete_col") # "complete" | "incomplete"
|
||||
|
||||
These tests exercise the live worktree PyO3 extension and public sync/async
|
||||
wrappers. While the public methods and hidden native bridge are absent they
|
||||
fail against that extension; once present they freeze the public contract
|
||||
below. They must not fake success paths.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import inspect
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable, Literal, get_type_hints
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
import lancedb.table
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.remote.table import RemoteTable
|
||||
from lancedb.table import AsyncTable, LanceTable, Table
|
||||
|
||||
_TABLE_NAME = "articles"
|
||||
_DESCRIBE_PATH = f"/v1/table/{_TABLE_NAME}/describe/"
|
||||
|
||||
_ORDINARY_FIELD_ID = 1
|
||||
_COMPLETE_FIELD_ID = 5
|
||||
_INCOMPLETE_FIELD_ID = 7
|
||||
_STABLE_FIELD_IDS = [_ORDINARY_FIELD_ID, _COMPLETE_FIELD_ID, _INCOMPLETE_FIELD_ID]
|
||||
|
||||
_STATUS_FUNCTION_ID = "fn.exact.status.projection"
|
||||
_METADATA_KEY = "lancedb::generated_column"
|
||||
_RAW_METADATA_MARKER = "SENSITIVE_STATUS_METADATA_MARKER_b3d2_py_9f2e"
|
||||
|
||||
# Pinned Rust-canonical schema-only Utf8 type IPC (base64), shared with FF-028.
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_EXPECTED_RETURN = Literal["complete", "incomplete"]
|
||||
|
||||
_FORBIDDEN_PUBLIC_NAMES = (
|
||||
"GeneratedColumnStatus",
|
||||
"GeneratedColumnDefinition",
|
||||
"GeneratedColumnBindingSnapshot",
|
||||
"GeneratedColumnBindingEntry",
|
||||
)
|
||||
|
||||
_FORBIDDEN_BRIDGE_KWARGS = (
|
||||
"epoch",
|
||||
"dependency_epoch",
|
||||
"materialized_epoch",
|
||||
"function_id",
|
||||
"field_id",
|
||||
"field_ids",
|
||||
"version",
|
||||
"branch",
|
||||
"wait",
|
||||
"job",
|
||||
"request",
|
||||
"backend",
|
||||
)
|
||||
|
||||
|
||||
def _definition_metadata_json(
|
||||
output_field_id: int,
|
||||
dependency_epoch: int,
|
||||
materialized_epoch: int,
|
||||
*,
|
||||
text_field_id: int = _ORDINARY_FIELD_ID,
|
||||
) -> str:
|
||||
"""Exact JSON stored under Arrow field metadata ``lancedb::generated_column``."""
|
||||
return json.dumps(
|
||||
{
|
||||
"format_version": 1,
|
||||
"output_field_id": output_field_id,
|
||||
"function_call": {
|
||||
"function_id": _STATUS_FUNCTION_ID,
|
||||
"arguments": [
|
||||
{
|
||||
"parameter": "text",
|
||||
"value": {
|
||||
"kind": "field",
|
||||
"field_id": text_field_id,
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
"dependency_epoch": dependency_epoch,
|
||||
"materialized_epoch": materialized_epoch,
|
||||
},
|
||||
separators=(",", ":"),
|
||||
)
|
||||
|
||||
|
||||
def _field(
|
||||
name: str,
|
||||
*,
|
||||
arrow_type: str = "string",
|
||||
nullable: bool = True,
|
||||
metadata: dict[str, str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
body: dict[str, Any] = {
|
||||
"name": name,
|
||||
"type": {"type": arrow_type},
|
||||
"nullable": nullable,
|
||||
}
|
||||
if metadata is not None:
|
||||
body["metadata"] = metadata
|
||||
return body
|
||||
|
||||
|
||||
def _status_schema_fields(
|
||||
*,
|
||||
complete_meta: str | None = None,
|
||||
incomplete_meta: str | None = None,
|
||||
bad_name: str | None = None,
|
||||
bad_meta: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
fields = [
|
||||
_field("ordinary", arrow_type="string"),
|
||||
_field(
|
||||
"complete_col",
|
||||
arrow_type="int32",
|
||||
metadata={
|
||||
_METADATA_KEY: complete_meta
|
||||
if complete_meta is not None
|
||||
else _definition_metadata_json(_COMPLETE_FIELD_ID, 3, 3)
|
||||
},
|
||||
),
|
||||
_field(
|
||||
"incomplete_col",
|
||||
arrow_type="int32",
|
||||
metadata={
|
||||
_METADATA_KEY: incomplete_meta
|
||||
if incomplete_meta is not None
|
||||
else _definition_metadata_json(_INCOMPLETE_FIELD_ID, 4, 1)
|
||||
},
|
||||
),
|
||||
]
|
||||
if bad_name is not None and bad_meta is not None:
|
||||
fields.append(
|
||||
_field(
|
||||
bad_name,
|
||||
arrow_type="int32",
|
||||
metadata={_METADATA_KEY: bad_meta},
|
||||
)
|
||||
)
|
||||
return {"fields": fields}
|
||||
|
||||
|
||||
def _describe_body(
|
||||
*,
|
||||
version: int = 11,
|
||||
field_ids: list[int] | None = _STABLE_FIELD_IDS,
|
||||
schema: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
body: dict[str, Any] = {
|
||||
"version": version,
|
||||
"schema": schema if schema is not None else _status_schema_fields(),
|
||||
}
|
||||
if field_ids is not None:
|
||||
body["field_ids"] = field_ids
|
||||
return body
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_text(exc: BaseException) -> str:
|
||||
parts = [str(exc), repr(exc)]
|
||||
current: BaseException | None = exc
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(f"{type(current).__name__}: {current}")
|
||||
current = current.__cause__ or current.__context__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _json_response(
|
||||
request: http.server.BaseHTTPRequestHandler, body: dict[str, Any]
|
||||
) -> None:
|
||||
payload = json.dumps(body).encode("utf-8")
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(payload)
|
||||
|
||||
|
||||
class _RequestLog:
|
||||
"""Track post-open describe and any non-describe operation traffic."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.describe: list[dict[str, Any]] = []
|
||||
self.other: list[str] = []
|
||||
self.recording = False
|
||||
|
||||
def start(self) -> None:
|
||||
self.describe.clear()
|
||||
self.other.clear()
|
||||
self.recording = True
|
||||
|
||||
def note(self, path: str, body: dict[str, Any] | None = None) -> None:
|
||||
if not self.recording:
|
||||
return
|
||||
if path == _DESCRIBE_PATH:
|
||||
self.describe.append(body or {})
|
||||
else:
|
||||
self.other.append(path)
|
||||
|
||||
|
||||
def _assert_no_operation_traffic(log: _RequestLog) -> None:
|
||||
assert log.describe == []
|
||||
assert log.other == []
|
||||
|
||||
|
||||
def _assert_one_status_describe(log: _RequestLog) -> None:
|
||||
assert len(log.describe) == 1, f"expected one status describe, got {log.describe!r}"
|
||||
assert log.other == [], f"unexpected non-describe traffic: {log.other!r}"
|
||||
|
||||
|
||||
def _assert_exact_public_signature(method: Any) -> None:
|
||||
"""Freeze ``(self, column_name)`` with no varargs/kwargs/keyword-only escape."""
|
||||
params = list(inspect.signature(method).parameters.values())
|
||||
assert [p.name for p in params] == ["self", "column_name"]
|
||||
for param in params:
|
||||
assert param.kind in (
|
||||
inspect.Parameter.POSITIONAL_ONLY,
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
)
|
||||
assert param.default is inspect.Parameter.empty
|
||||
assert param.kind is not inspect.Parameter.VAR_POSITIONAL
|
||||
assert param.kind is not inspect.Parameter.VAR_KEYWORD
|
||||
assert param.kind is not inspect.Parameter.KEYWORD_ONLY
|
||||
|
||||
|
||||
def _assert_status_string(value: Any, expected: str) -> None:
|
||||
assert value == expected
|
||||
assert type(value) is str
|
||||
assert value in ("complete", "incomplete")
|
||||
|
||||
|
||||
def _open_remote_table(
|
||||
*,
|
||||
status_describe: dict[str, Any] | None = None,
|
||||
):
|
||||
"""Open sync RemoteTable; return (table, log, cm)."""
|
||||
log = _RequestLog()
|
||||
binding = status_describe if status_describe is not None else _describe_body()
|
||||
open_describe = {
|
||||
"version": 1,
|
||||
"schema": {"fields": [_field("ordinary", arrow_type="string")]},
|
||||
}
|
||||
state = {"opened": False}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding)
|
||||
return
|
||||
|
||||
log.note(request.path, body)
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected path")
|
||||
|
||||
cm = _mock_remote_db(handler)
|
||||
db = cm.__enter__()
|
||||
table = db.open_table(_TABLE_NAME)
|
||||
assert isinstance(table, RemoteTable)
|
||||
log.start()
|
||||
return table, log, cm
|
||||
|
||||
|
||||
async def _open_remote_table_async(
|
||||
*,
|
||||
status_describe: dict[str, Any] | None = None,
|
||||
):
|
||||
"""Open async table under a live mock server; return (table, log, cm)."""
|
||||
log = _RequestLog()
|
||||
binding = status_describe if status_describe is not None else _describe_body()
|
||||
open_describe = {
|
||||
"version": 1,
|
||||
"schema": {"fields": [_field("ordinary", arrow_type="string")]},
|
||||
}
|
||||
state = {"opened": False}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding)
|
||||
return
|
||||
|
||||
log.note(request.path, body)
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected path")
|
||||
|
||||
cm = _mock_remote_db_async(handler)
|
||||
db = await cm.__aenter__()
|
||||
table = await db.open_table(_TABLE_NAME)
|
||||
assert isinstance(table, AsyncTable)
|
||||
log.start()
|
||||
return table, log, cm
|
||||
|
||||
|
||||
def test_no_public_generated_column_status_resource_exported():
|
||||
"""Baseline: no public status class/enum/resource is exported."""
|
||||
for mod in (lancedb, lancedb.table, _native):
|
||||
for name in _FORBIDDEN_PUBLIC_NAMES:
|
||||
assert not hasattr(mod, name), f"{mod.__name__}.{name} must not be public"
|
||||
|
||||
|
||||
def test_public_surface_signatures_annotations_and_hidden_bridge():
|
||||
"""Four public methods + hidden native bridge must exist with frozen shape."""
|
||||
assert hasattr(_native.Table, "_generated_column_status"), (
|
||||
"native private bridge Table._generated_column_status is missing"
|
||||
)
|
||||
assert hasattr(Table, "generated_column_status"), (
|
||||
"Table.generated_column_status is missing"
|
||||
)
|
||||
assert hasattr(LanceTable, "generated_column_status"), (
|
||||
"LanceTable.generated_column_status is missing"
|
||||
)
|
||||
assert hasattr(RemoteTable, "generated_column_status"), (
|
||||
"RemoteTable.generated_column_status is missing"
|
||||
)
|
||||
assert hasattr(AsyncTable, "generated_column_status"), (
|
||||
"AsyncTable.generated_column_status is missing"
|
||||
)
|
||||
|
||||
bridge = _native.Table._generated_column_status
|
||||
_assert_exact_public_signature(bridge)
|
||||
for keyword in _FORBIDDEN_BRIDGE_KWARGS:
|
||||
assert keyword not in inspect.signature(bridge).parameters
|
||||
|
||||
for method in (
|
||||
Table.generated_column_status,
|
||||
LanceTable.generated_column_status,
|
||||
RemoteTable.generated_column_status,
|
||||
):
|
||||
_assert_exact_public_signature(method)
|
||||
assert not inspect.iscoroutinefunction(method)
|
||||
assert get_type_hints(method)["return"] == _EXPECTED_RETURN
|
||||
|
||||
async_method = AsyncTable.generated_column_status
|
||||
_assert_exact_public_signature(async_method)
|
||||
assert inspect.iscoroutinefunction(async_method)
|
||||
assert get_type_hints(async_method)["return"] == _EXPECTED_RETURN
|
||||
|
||||
|
||||
def test_sync_remote_complete_and_incomplete_one_describe_each():
|
||||
table, log, cm = _open_remote_table()
|
||||
try:
|
||||
complete = table.generated_column_status("complete_col")
|
||||
_assert_status_string(complete, "complete")
|
||||
_assert_one_status_describe(log)
|
||||
|
||||
log.start()
|
||||
incomplete = table.generated_column_status("incomplete_col")
|
||||
_assert_status_string(incomplete, "incomplete")
|
||||
_assert_one_status_describe(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_complete_and_incomplete_one_describe_each():
|
||||
table, log, cm = await _open_remote_table_async()
|
||||
try:
|
||||
complete = await table.generated_column_status("complete_col")
|
||||
_assert_status_string(complete, "complete")
|
||||
_assert_one_status_describe(log)
|
||||
|
||||
log.start()
|
||||
incomplete = await table.generated_column_status("incomplete_col")
|
||||
_assert_status_string(incomplete, "incomplete")
|
||||
_assert_one_status_describe(log)
|
||||
finally:
|
||||
await cm.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("column_name", "status_describe", "expected_exc"),
|
||||
[
|
||||
(
|
||||
"missing",
|
||||
_describe_body(),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"Complete_Col",
|
||||
_describe_body(),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"ordinary",
|
||||
_describe_body(),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"complete_col",
|
||||
_describe_body(
|
||||
schema=_status_schema_fields(
|
||||
complete_meta=_definition_metadata_json(
|
||||
_COMPLETE_FIELD_ID + 1, 3, 3
|
||||
)
|
||||
)
|
||||
),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"gen_bad",
|
||||
_describe_body(
|
||||
field_ids=[*_STABLE_FIELD_IDS, 9],
|
||||
schema=_status_schema_fields(
|
||||
bad_name="gen_bad",
|
||||
bad_meta=(
|
||||
'{"format_version":1,"output_field_id":9,'
|
||||
f'"function_call":{_RAW_METADATA_MARKER},'
|
||||
'"dependency_epoch":1,"materialized_epoch":1}'
|
||||
),
|
||||
),
|
||||
),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"complete_col",
|
||||
_describe_body(
|
||||
schema=_status_schema_fields(
|
||||
complete_meta=_definition_metadata_json(
|
||||
_COMPLETE_FIELD_ID, 1, 1
|
||||
).replace('"format_version":1', '"format_version":2')
|
||||
)
|
||||
),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"incomplete_col",
|
||||
_describe_body(
|
||||
schema=_status_schema_fields(
|
||||
incomplete_meta=_definition_metadata_json(
|
||||
_INCOMPLETE_FIELD_ID, 1, 2
|
||||
)
|
||||
)
|
||||
),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"complete_col",
|
||||
_describe_body(field_ids=None),
|
||||
NotImplementedError,
|
||||
),
|
||||
],
|
||||
ids=[
|
||||
"missing",
|
||||
"case_mismatch",
|
||||
"ordinary",
|
||||
"output_id_mismatch",
|
||||
"malformed_metadata",
|
||||
"unknown_format_version",
|
||||
"reversed_epochs",
|
||||
"old_server_missing_field_ids",
|
||||
],
|
||||
)
|
||||
def test_remote_fail_closed_matrix_one_describe(
|
||||
column_name: str,
|
||||
status_describe: dict[str, Any],
|
||||
expected_exc: type[BaseException],
|
||||
):
|
||||
table, log, cm = _open_remote_table(status_describe=status_describe)
|
||||
try:
|
||||
with pytest.raises(expected_exc) as raised:
|
||||
table.generated_column_status(column_name)
|
||||
text = _exception_text(raised.value)
|
||||
assert _RAW_METADATA_MARKER not in text
|
||||
_assert_one_status_describe(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_sync_empty_name_zero_post_open_requests():
|
||||
table, log, cm = _open_remote_table()
|
||||
try:
|
||||
with pytest.raises(ValueError):
|
||||
table.generated_column_status("")
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_empty_name_zero_post_open_requests():
|
||||
table, log, cm = await _open_remote_table_async()
|
||||
try:
|
||||
with pytest.raises(ValueError):
|
||||
await table.generated_column_status("")
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
await cm.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_closed_status_empty_validation_wins_and_nonempty_closed():
|
||||
"""Publicly closed AsyncTable: empty validates first; nonempty is closed."""
|
||||
table, log, cm = await _open_remote_table_async()
|
||||
try:
|
||||
table.close()
|
||||
|
||||
log.start()
|
||||
try:
|
||||
await table.generated_column_status("complete_col")
|
||||
except AttributeError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
text = _exception_text(exc)
|
||||
assert "closed" in text.lower()
|
||||
else:
|
||||
pytest.fail("closed AsyncTable must fail before transport")
|
||||
_assert_no_operation_traffic(log)
|
||||
|
||||
log.start()
|
||||
with pytest.raises(ValueError) as raised:
|
||||
await table.generated_column_status("")
|
||||
text = _exception_text(raised.value)
|
||||
assert "closed" not in text.lower()
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
await cm.__aexit__(None, None, None)
|
||||
|
||||
|
||||
def test_local_sync_ordinary_column_fails_without_side_effects(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table(
|
||||
"ordinary_only",
|
||||
[{"ordinary": "alpha", "id": 1}, {"ordinary": "beta", "id": 2}],
|
||||
)
|
||||
assert isinstance(table, LanceTable)
|
||||
version_before = table.version
|
||||
schema_before = table.schema
|
||||
data_before = table.to_arrow()
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
table.generated_column_status("ordinary")
|
||||
|
||||
assert table.version == version_before
|
||||
assert table.schema == schema_before
|
||||
assert table.to_arrow().equals(data_before)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_ordinary_column_fails_without_side_effects(tmp_path):
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
table = await db.create_table(
|
||||
"ordinary_only_async",
|
||||
[{"ordinary": "alpha", "id": 1}, {"ordinary": "beta", "id": 2}],
|
||||
)
|
||||
assert isinstance(table, AsyncTable)
|
||||
version_before = await table.version()
|
||||
schema_before = await table.schema()
|
||||
data_before = await table.to_arrow()
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
await table.generated_column_status("ordinary")
|
||||
|
||||
assert await table.version() == version_before
|
||||
assert await table.schema() == schema_before
|
||||
assert (await table.to_arrow()).equals(data_before)
|
||||
@@ -0,0 +1,291 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for the local @udf declaration surface."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import inspect
|
||||
import types
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import Function, Job, udf
|
||||
from lancedb._udf import _get_udf_config
|
||||
|
||||
_REMOVED_AUTHORING_KNOBS = (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_handling",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"error_policy",
|
||||
"FunctionVersion",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
)
|
||||
|
||||
|
||||
def _decorate(fn, **overrides):
|
||||
kwargs = {
|
||||
"inputs": {"x": pa.int32()},
|
||||
"output": pa.int64(),
|
||||
"python": "3.12",
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return udf(**kwargs)(fn)
|
||||
|
||||
|
||||
def test_udf_top_level_export_and_identity_metadata_behavior():
|
||||
assert "udf" in lancedb.__all__
|
||||
assert udf is lancedb.udf
|
||||
assert isinstance(importlib.import_module("lancedb._udf"), types.ModuleType)
|
||||
assert not isinstance(lancedb.udf, types.ModuleType)
|
||||
|
||||
def add(x, y=1):
|
||||
"""Add locally."""
|
||||
return x + y
|
||||
|
||||
original = add
|
||||
decorated = _decorate(
|
||||
add,
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
)
|
||||
|
||||
assert decorated is original
|
||||
assert decorated.__name__ == "add"
|
||||
assert decorated.__doc__ == "Add locally."
|
||||
assert str(inspect.signature(decorated)) == "(x, y=1)"
|
||||
assert decorated(2) == 3
|
||||
assert decorated(2, 5) == 7
|
||||
assert decorated(x=4, y=6) == 10
|
||||
|
||||
|
||||
def test_udf_config_snapshot_order_defaults_and_immutability():
|
||||
inputs = {"z": pa.string(), "a": pa.int32()}
|
||||
packages = ["pkg-b==2", "pkg-a==1"]
|
||||
|
||||
def combine(z, a):
|
||||
return f"{z}:{a}"
|
||||
|
||||
decorated = udf(
|
||||
inputs=inputs,
|
||||
output=pa.string(),
|
||||
python="3.11",
|
||||
packages=packages,
|
||||
output_nullable=False,
|
||||
)(combine)
|
||||
|
||||
config = _get_udf_config(decorated)
|
||||
assert config.inputs == (("z", pa.string()), ("a", pa.int32()))
|
||||
assert isinstance(config.inputs, tuple)
|
||||
assert config.output == pa.string()
|
||||
assert config.output_nullable is False
|
||||
assert config.python == "3.11"
|
||||
assert config.packages == ("pkg-b==2", "pkg-a==1")
|
||||
assert isinstance(config.packages, tuple)
|
||||
|
||||
inputs["extra"] = pa.bool_()
|
||||
del inputs["z"]
|
||||
packages.append("pkg-c==3")
|
||||
packages[0] = "mutated==0"
|
||||
assert config.inputs == (("z", pa.string()), ("a", pa.int32()))
|
||||
assert config.packages == ("pkg-b==2", "pkg-a==1")
|
||||
|
||||
for attr in ("inputs", "output", "output_nullable", "python", "packages"):
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(config, attr, None)
|
||||
|
||||
def defaults_only(x):
|
||||
return x
|
||||
|
||||
defaulted = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
)(defaults_only)
|
||||
default_config = _get_udf_config(defaulted)
|
||||
assert default_config.packages == ()
|
||||
assert default_config.output_nullable is True
|
||||
|
||||
|
||||
def test_udf_accepts_lambda_and_closure_for_local_declaration():
|
||||
ambient = "ambient-secret-value-xyz"
|
||||
|
||||
lam = udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(lambda n: n + 1)
|
||||
assert lam(3) == 4
|
||||
assert _get_udf_config(lam).inputs == (("n", pa.int32()),)
|
||||
|
||||
def factory(offset):
|
||||
@udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
packages=["demo==0.1"],
|
||||
)
|
||||
def closed(n):
|
||||
return n + offset + len(ambient)
|
||||
|
||||
return closed
|
||||
|
||||
closed = factory(10)
|
||||
assert closed(2) == 12 + len(ambient)
|
||||
assert _get_udf_config(closed).packages == ("demo==0.1",)
|
||||
|
||||
|
||||
def test_udf_declaration_defers_signature_and_implementation_packaging():
|
||||
"""Declaration must not validate callable signature or embed implementation."""
|
||||
|
||||
def local_add(left, right=1):
|
||||
return left + right
|
||||
|
||||
decorated = udf(
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(local_add)
|
||||
|
||||
assert decorated is local_add
|
||||
assert str(inspect.signature(decorated)) == "(left, right=1)"
|
||||
assert decorated(2) == 3
|
||||
assert decorated(2, 5) == 7
|
||||
|
||||
config = _get_udf_config(decorated)
|
||||
assert config.inputs == (("x", pa.int32()), ("y", pa.int32()))
|
||||
for attr in (
|
||||
"source",
|
||||
"module",
|
||||
"callable",
|
||||
"function",
|
||||
"implementation",
|
||||
"bundle",
|
||||
"artifact",
|
||||
"digest",
|
||||
):
|
||||
assert not hasattr(config, attr)
|
||||
|
||||
|
||||
def test_udf_lookup_double_decoration_and_non_function_target():
|
||||
def plain(x):
|
||||
return x
|
||||
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
_get_udf_config(plain)
|
||||
|
||||
decorated = _decorate(plain)
|
||||
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
_decorate(decorated)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(object())
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(42)
|
||||
|
||||
|
||||
def test_udf_config_validation_errors():
|
||||
def target(x):
|
||||
return x
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
udf({"x": pa.int32()}, pa.int32(), "3.12")(target)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, inputs=[("x", pa.int32())])
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, inputs={1: pa.int32()})
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_decorate(target, inputs={"": pa.int32()})
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, inputs={"x": "int32"})
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, output="int64")
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, python=3.12)
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_decorate(target, python="")
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, packages="pkg==1")
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_decorate(target, packages=["pkg==1", ""])
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_decorate(target, packages=["pkg==1", "pkg==1"])
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, packages=["pkg==1", 2])
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, output_nullable=1)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, output_nullable="true")
|
||||
|
||||
|
||||
def test_udf_rejects_removed_overdesign_and_has_no_durable_side_effects():
|
||||
params = inspect.signature(udf).parameters
|
||||
for name in _REMOVED_AUTHORING_KNOBS:
|
||||
assert name not in params
|
||||
|
||||
def score(x):
|
||||
"""score body marker unique-xyz."""
|
||||
ambient = "ambient-secret-value-xyz"
|
||||
return f"{ambient}:{x}"
|
||||
|
||||
decorated = _decorate(
|
||||
score,
|
||||
packages=["score==1.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
config = _get_udf_config(decorated)
|
||||
text = repr(config).lower()
|
||||
|
||||
assert "score body marker unique-xyz" not in text
|
||||
assert "ambient-secret-value-xyz" not in text
|
||||
for token in (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"functionversion",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
):
|
||||
assert token not in text
|
||||
|
||||
for attr in _REMOVED_AUTHORING_KNOBS:
|
||||
assert not hasattr(config, attr)
|
||||
|
||||
assert not isinstance(decorated, Function)
|
||||
assert not isinstance(decorated, Job)
|
||||
for attr in ("id", "function_id", "job", "job_id", "registration"):
|
||||
assert not hasattr(decorated, attr)
|
||||
@@ -0,0 +1,490 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for local FunctionCapability authoring and @udf capabilities."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import Function, FunctionCapability, Job, udf
|
||||
from lancedb._udf import _get_udf_config, _package_udf
|
||||
|
||||
_SECRET_REFERENCE = "secret://team/capability-redact-token-xyz"
|
||||
_SECRET_ENV = "API_TOKEN"
|
||||
_NETWORK_ORIGIN = "https://api.example.com"
|
||||
|
||||
_OVERDESIGN_ATTRS = (
|
||||
"id",
|
||||
"function_id",
|
||||
"FunctionVersion",
|
||||
"function_version",
|
||||
"artifact",
|
||||
"digest",
|
||||
"authorization",
|
||||
"authorized",
|
||||
"value",
|
||||
"plaintext",
|
||||
"plaintext_secret",
|
||||
"secret_value",
|
||||
"job",
|
||||
"job_id",
|
||||
"catalog",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_handling",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"error_policy",
|
||||
"geneva",
|
||||
)
|
||||
|
||||
|
||||
def _decorate(fn, **overrides):
|
||||
kwargs = {
|
||||
"inputs": {"x": pa.int32()},
|
||||
"output": pa.int64(),
|
||||
"python": "3.12",
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return udf(**kwargs)(fn)
|
||||
|
||||
|
||||
def _network(origin: str = _NETWORK_ORIGIN) -> FunctionCapability:
|
||||
return FunctionCapability.network(origin)
|
||||
|
||||
|
||||
def _secret(
|
||||
reference: str = _SECRET_REFERENCE,
|
||||
*,
|
||||
environment_variable: str = _SECRET_ENV,
|
||||
) -> FunctionCapability:
|
||||
return FunctionCapability.secret(
|
||||
reference,
|
||||
environment_variable=environment_variable,
|
||||
)
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=False,
|
||||
)
|
||||
def packable_without_capabilities(x):
|
||||
return x + 1
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=False,
|
||||
capabilities=[
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
],
|
||||
)
|
||||
def packable_with_capabilities(x):
|
||||
return x + 1
|
||||
|
||||
|
||||
def test_function_capability_export_factories_projection_equality_immutability():
|
||||
assert "FunctionCapability" in lancedb.__all__
|
||||
assert FunctionCapability is lancedb.FunctionCapability
|
||||
|
||||
network = _network()
|
||||
secret = _secret()
|
||||
|
||||
assert network.kind == "network"
|
||||
assert network.origin == _NETWORK_ORIGIN
|
||||
assert network.reference is None
|
||||
assert network.environment_variable is None
|
||||
|
||||
assert secret.kind == "secret"
|
||||
assert secret.reference == _SECRET_REFERENCE
|
||||
assert secret.environment_variable == _SECRET_ENV
|
||||
assert secret.origin is None
|
||||
|
||||
assert network == FunctionCapability.network(_NETWORK_ORIGIN)
|
||||
assert secret == FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
)
|
||||
assert network != secret
|
||||
assert network != FunctionCapability.network("https://other.example.com")
|
||||
assert secret != FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable="OTHER_TOKEN",
|
||||
)
|
||||
|
||||
public_attrs = ("kind", "origin", "reference", "environment_variable")
|
||||
internal_slots = ("_kind", "_origin", "_reference", "_environment_variable")
|
||||
immutable_attrs = public_attrs + internal_slots
|
||||
|
||||
for attr in public_attrs:
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(network, attr, None)
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(secret, attr, None)
|
||||
|
||||
for attr in immutable_attrs:
|
||||
# Fresh instances per attempt so a RED slot mutation cannot corrupt
|
||||
# shared fixtures used by later assertions in this test.
|
||||
fresh_network = _network("https://fresh-immutability.example.com")
|
||||
fresh_secret = _secret(
|
||||
"secret://team/fresh-immutability-token",
|
||||
environment_variable="FRESH_IMMUTABILITY_TOKEN",
|
||||
)
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(fresh_network, attr, None)
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(fresh_secret, attr, None)
|
||||
with pytest.raises(AttributeError):
|
||||
delattr(fresh_network, attr)
|
||||
with pytest.raises(AttributeError):
|
||||
delattr(fresh_secret, attr)
|
||||
|
||||
retained_origin = "https://config-retain.example.com"
|
||||
retained_reference = "secret://team/config-retain-token"
|
||||
retained_env = "CONFIG_RETAIN_TOKEN"
|
||||
retained_network = FunctionCapability.network(retained_origin)
|
||||
retained_secret = FunctionCapability.secret(
|
||||
retained_reference,
|
||||
environment_variable=retained_env,
|
||||
)
|
||||
expected_capabilities = (
|
||||
FunctionCapability.network(retained_origin),
|
||||
FunctionCapability.secret(
|
||||
retained_reference,
|
||||
environment_variable=retained_env,
|
||||
),
|
||||
)
|
||||
|
||||
def retain_target(x):
|
||||
return x
|
||||
|
||||
retained = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
capabilities=[retained_network, retained_secret],
|
||||
)(retain_target)
|
||||
retained_config = _get_udf_config(retained)
|
||||
assert retained_config.capabilities == expected_capabilities
|
||||
|
||||
for attr in immutable_attrs:
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(retained_network, attr, "mutated")
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(retained_secret, attr, "mutated")
|
||||
with pytest.raises(AttributeError):
|
||||
delattr(retained_network, attr)
|
||||
with pytest.raises(AttributeError):
|
||||
delattr(retained_secret, attr)
|
||||
|
||||
assert retained_config.capabilities == expected_capabilities
|
||||
assert retained_config.capabilities[0] is retained_network
|
||||
assert retained_config.capabilities[1] is retained_secret
|
||||
assert retained_config.capabilities[0].kind == "network"
|
||||
assert retained_config.capabilities[0].origin == retained_origin
|
||||
assert retained_config.capabilities[0].reference is None
|
||||
assert retained_config.capabilities[0].environment_variable is None
|
||||
assert retained_config.capabilities[1].kind == "secret"
|
||||
assert retained_config.capabilities[1].reference == retained_reference
|
||||
assert retained_config.capabilities[1].environment_variable == retained_env
|
||||
assert retained_config.capabilities[1].origin is None
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability()
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability( # type: ignore[call-arg]
|
||||
kind="network",
|
||||
origin=_NETWORK_ORIGIN,
|
||||
)
|
||||
|
||||
assert not isinstance(network, Function)
|
||||
assert not isinstance(secret, Function)
|
||||
assert not isinstance(network, Job)
|
||||
assert not isinstance(secret, Job)
|
||||
for attr in _OVERDESIGN_ATTRS:
|
||||
assert not hasattr(network, attr)
|
||||
assert not hasattr(secret, attr)
|
||||
|
||||
|
||||
def test_function_capability_validation_and_secret_redaction():
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.network(None) # type: ignore[arg-type]
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.network(123) # type: ignore[arg-type]
|
||||
with pytest.raises(ValueError):
|
||||
FunctionCapability.network("")
|
||||
|
||||
# Backend authorization owns URL/scheme policy; non-empty is enough here.
|
||||
loose = FunctionCapability.network("example.com")
|
||||
assert loose.kind == "network"
|
||||
assert loose.origin == "example.com"
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[misc]
|
||||
_SECRET_REFERENCE,
|
||||
_SECRET_ENV,
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret(None, environment_variable=_SECRET_ENV) # type: ignore[arg-type]
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret(123, environment_variable=_SECRET_ENV) # type: ignore[arg-type]
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable=None) # type: ignore[arg-type]
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable=1) # type: ignore[arg-type]
|
||||
|
||||
with pytest.raises(ValueError) as empty_ref:
|
||||
FunctionCapability.secret("", environment_variable=_SECRET_ENV)
|
||||
assert _SECRET_REFERENCE not in str(empty_ref.value)
|
||||
assert _SECRET_REFERENCE not in repr(empty_ref.value)
|
||||
|
||||
with pytest.raises(ValueError) as empty_env:
|
||||
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable="")
|
||||
assert _SECRET_REFERENCE not in str(empty_env.value)
|
||||
assert _SECRET_REFERENCE not in repr(empty_env.value)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[call-arg]
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
value="super-secret",
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[call-arg]
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
plaintext_secret="super-secret",
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[call-arg]
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
environment={_SECRET_ENV: "super-secret"},
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[call-arg]
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
headers={"Authorization": "Bearer super-secret"},
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.network( # type: ignore[call-arg]
|
||||
_NETWORK_ORIGIN,
|
||||
headers={"X-Trace": "1"},
|
||||
)
|
||||
|
||||
secret = _secret()
|
||||
assert not hasattr(secret, "value")
|
||||
assert not hasattr(secret, "plaintext")
|
||||
assert not hasattr(secret, "plaintext_secret")
|
||||
assert not hasattr(secret, "secret_value")
|
||||
|
||||
secret_text = repr(secret)
|
||||
assert "secret" in secret_text.lower()
|
||||
assert _SECRET_ENV in secret_text
|
||||
assert _SECRET_REFERENCE not in secret_text
|
||||
assert "super-secret" not in secret_text
|
||||
|
||||
network_text = repr(_network())
|
||||
assert "network" in network_text.lower()
|
||||
assert _NETWORK_ORIGIN in network_text
|
||||
|
||||
|
||||
def test_udf_capabilities_ordered_immutable_config_default_and_validation():
|
||||
params = inspect.signature(udf).parameters
|
||||
assert "capabilities" in params
|
||||
assert params["capabilities"].kind is inspect.Parameter.KEYWORD_ONLY
|
||||
assert params["capabilities"].default == ()
|
||||
|
||||
def identity_target(x):
|
||||
"""capabilities identity marker."""
|
||||
return x + 1
|
||||
|
||||
original = identity_target
|
||||
decorated = _decorate(identity_target)
|
||||
assert decorated is original
|
||||
assert decorated.__name__ == "identity_target"
|
||||
assert decorated.__doc__ == "capabilities identity marker."
|
||||
assert decorated(2) == 3
|
||||
assert _get_udf_config(decorated).capabilities == ()
|
||||
|
||||
first = _network("https://b.example.com")
|
||||
second = _network("https://a.example.com")
|
||||
third = _network("https://b.example.com")
|
||||
secret = _secret()
|
||||
capabilities = [first, second, third, secret]
|
||||
|
||||
def combine(x):
|
||||
return x
|
||||
|
||||
with_caps = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
packages=["pkg-b==2", "pkg-a==1"],
|
||||
capabilities=capabilities,
|
||||
)(combine)
|
||||
config = _get_udf_config(with_caps)
|
||||
assert config.capabilities == (first, second, third, secret)
|
||||
assert isinstance(config.capabilities, tuple)
|
||||
assert config.packages == ("pkg-b==2", "pkg-a==1")
|
||||
assert config.inputs == (("x", pa.int32()),)
|
||||
|
||||
capabilities.append(_network("https://mutated.example.com"))
|
||||
capabilities[0] = _network("https://replaced.example.com")
|
||||
assert config.capabilities == (first, second, third, secret)
|
||||
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(config, "capabilities", ())
|
||||
|
||||
def target(x):
|
||||
return x
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, capabilities="https://api.example.com")
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, capabilities=b"https://api.example.com")
|
||||
|
||||
class _BadCapability:
|
||||
def __repr__(self) -> str:
|
||||
return "unique-bad-capability-repr-xyz"
|
||||
|
||||
with pytest.raises(TypeError) as bad_item:
|
||||
_decorate(target, capabilities=[_BadCapability()])
|
||||
assert "unique-bad-capability-repr-xyz" not in str(bad_item.value)
|
||||
assert "unique-bad-capability-repr-xyz" not in repr(bad_item.value)
|
||||
|
||||
with pytest.raises(TypeError) as bad_mixed:
|
||||
_decorate(
|
||||
target,
|
||||
capabilities=[_network(), "unique-bad-capability-string-xyz"],
|
||||
)
|
||||
assert "unique-bad-capability-string-xyz" not in str(bad_mixed.value)
|
||||
assert "unique-bad-capability-string-xyz" not in repr(bad_mixed.value)
|
||||
|
||||
|
||||
def test_udf_capabilities_rejects_function_capability_subclass_before_property_access():
|
||||
marker = "unique-hostile-capability-subclass-marker-xyz"
|
||||
|
||||
class _HostileFunctionCapability(FunctionCapability):
|
||||
@property
|
||||
def kind(self) -> str:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def origin(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def reference(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def environment_variable(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
hostile = object.__new__(_HostileFunctionCapability)
|
||||
assert isinstance(hostile, FunctionCapability)
|
||||
assert type(hostile) is not FunctionCapability
|
||||
|
||||
def target(x):
|
||||
return x
|
||||
|
||||
with pytest.raises(TypeError) as exc_info:
|
||||
_decorate(target, capabilities=[hostile])
|
||||
assert marker not in str(exc_info.value)
|
||||
assert marker not in repr(exc_info.value)
|
||||
assert _SECRET_REFERENCE not in str(exc_info.value)
|
||||
assert _SECRET_REFERENCE not in repr(exc_info.value)
|
||||
assert exc_info.value.__cause__ is None
|
||||
assert exc_info.value.__context__ is None
|
||||
|
||||
|
||||
def test_package_udf_preserves_capabilities_and_redacts_secret_reference():
|
||||
packaged = _package_udf(packable_with_capabilities)
|
||||
config = packaged.config
|
||||
|
||||
assert packaged.config is _get_udf_config(packable_with_capabilities)
|
||||
assert config.capabilities == (
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
)
|
||||
assert config.capabilities[0].kind == "network"
|
||||
assert config.capabilities[0].origin == _NETWORK_ORIGIN
|
||||
assert config.capabilities[1].kind == "secret"
|
||||
assert config.capabilities[1].reference == _SECRET_REFERENCE
|
||||
assert config.capabilities[1].environment_variable == _SECRET_ENV
|
||||
assert config.packages == ("pkg-a==1",)
|
||||
assert config.python == "3.12"
|
||||
assert config.output_nullable is False
|
||||
|
||||
nested = (
|
||||
f"{packaged!r}\n{config!r}\n{config.capabilities!r}\n{config.capabilities[1]!r}"
|
||||
)
|
||||
assert _SECRET_REFERENCE not in nested
|
||||
assert _SECRET_ENV in repr(config.capabilities[1])
|
||||
|
||||
|
||||
def test_capabilities_are_additive_to_existing_declaration_and_packaging():
|
||||
def score(x):
|
||||
return x
|
||||
|
||||
decorated = _decorate(
|
||||
score,
|
||||
packages=["score==1.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
config = _get_udf_config(decorated)
|
||||
assert config.inputs == (("x", pa.int32()),)
|
||||
assert config.output == pa.int64()
|
||||
assert config.output_nullable is True
|
||||
assert config.python == "3.12"
|
||||
assert config.packages == ("score==1.0",)
|
||||
assert config.capabilities == ()
|
||||
assert decorated is score
|
||||
assert decorated(4) == 4
|
||||
|
||||
packaged = _package_udf(packable_without_capabilities)
|
||||
assert packaged.config is _get_udf_config(packable_without_capabilities)
|
||||
assert packaged.callable_name == "packable_without_capabilities"
|
||||
assert packaged.config.capabilities == ()
|
||||
assert packaged.config.packages == ("pkg-a==1",)
|
||||
assert packaged.config.output_nullable is False
|
||||
assert packable_without_capabilities(1) == 2
|
||||
|
||||
params = inspect.signature(udf).parameters
|
||||
for name in (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_handling",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"error_policy",
|
||||
"FunctionVersion",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
):
|
||||
assert name not in params
|
||||
@@ -0,0 +1,506 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for the private UDF -> FunctionDefinition bridge."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import io
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import FunctionCapability, udf
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb import _udf as _udf_mod
|
||||
|
||||
_SOURCE_MARKER = "bridge-source-marker-unique-xyz"
|
||||
_SECRET_REFERENCE = "secret://team/bridge-redact-token-xyz"
|
||||
_SECRET_ENV = "BRIDGE_API_TOKEN"
|
||||
_NETWORK_ORIGIN = "https://api.bridge-example.com"
|
||||
_NETWORK_ORIGIN_B = "https://other.bridge-example.com"
|
||||
|
||||
_FORBIDDEN_WIRE_KEYS = (
|
||||
"id",
|
||||
"function_id",
|
||||
"FunctionId",
|
||||
"catalog",
|
||||
"catalog_name",
|
||||
"version",
|
||||
"function_version",
|
||||
"FunctionVersion",
|
||||
"lineage",
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"digest",
|
||||
"artifact",
|
||||
"artifact_digest",
|
||||
"storage",
|
||||
"storage_location",
|
||||
"location",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"nullPolicy",
|
||||
"timestamp",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"worker",
|
||||
"scheduler",
|
||||
"attempt",
|
||||
"attempt_id",
|
||||
"replica",
|
||||
"placement",
|
||||
"job",
|
||||
"job_id",
|
||||
"retry_key",
|
||||
"registration",
|
||||
)
|
||||
|
||||
_OVERDESIGN_ATTRS = (
|
||||
"id",
|
||||
"function_id",
|
||||
"FunctionVersion",
|
||||
"function_version",
|
||||
"artifact",
|
||||
"digest",
|
||||
"job",
|
||||
"job_id",
|
||||
"catalog",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"null_handling",
|
||||
)
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"text": pa.string(), "limit": pa.int32()},
|
||||
output=pa.string(),
|
||||
python="3.12",
|
||||
packages=["pkg-b==2", "pkg-a==1"],
|
||||
output_nullable=True,
|
||||
capabilities=[
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
FunctionCapability.network(_NETWORK_ORIGIN_B),
|
||||
],
|
||||
)
|
||||
def packable_bridge_normalize(text, limit):
|
||||
"""bridge-source-marker-unique-xyz."""
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def _build_function_definition(fn: object):
|
||||
return _udf_mod._build_function_definition(fn)
|
||||
|
||||
|
||||
def _function_definition_type():
|
||||
return _native._FunctionDefinition
|
||||
|
||||
|
||||
def _new_function_definition(**kwargs):
|
||||
return _native._new_function_definition(**kwargs)
|
||||
|
||||
|
||||
def _json_bytes(definition) -> bytes:
|
||||
payload = definition._to_json()
|
||||
if isinstance(payload, bytes):
|
||||
return payload
|
||||
assert isinstance(payload, str)
|
||||
return payload.encode("utf-8")
|
||||
|
||||
|
||||
def _decode_type_ipc(encoded: str) -> pa.DataType:
|
||||
raw = base64.b64decode(encoded)
|
||||
reader = pa.ipc.open_file(io.BytesIO(raw))
|
||||
assert reader.num_record_batches == 0
|
||||
assert len(reader.schema) == 1
|
||||
return reader.schema.field(0).type
|
||||
|
||||
|
||||
def _assert_exact_object_keys(value: dict, expected: set[str], *, context: str) -> None:
|
||||
assert isinstance(value, dict), f"{context} must be an object"
|
||||
assert set(value) == expected, f"{context} keys must match exactly: {set(value)!r}"
|
||||
|
||||
|
||||
def _assert_forbidden_keys_absent(value: object, *, context: str) -> None:
|
||||
if isinstance(value, dict):
|
||||
for key in value:
|
||||
assert key not in _FORBIDDEN_WIRE_KEYS, (
|
||||
f"forbidden key {key!r} at {context}: {value!r}"
|
||||
)
|
||||
if key == "name" and context in {
|
||||
"definition",
|
||||
"signature",
|
||||
"signature.output",
|
||||
"implementation",
|
||||
}:
|
||||
raise AssertionError(
|
||||
f"catalog/function identity key `name` must be absent at {context}"
|
||||
)
|
||||
child_context = f"{context}.{key}"
|
||||
if key == "parameters" and context == "signature":
|
||||
child_context = "signature.parameters"
|
||||
_assert_forbidden_keys_absent(value[key], context=child_context)
|
||||
elif isinstance(value, list):
|
||||
for idx, item in enumerate(value):
|
||||
item_context = (
|
||||
f"signature.parameters[{idx}]"
|
||||
if context == "signature.parameters"
|
||||
else f"{context}[{idx}]"
|
||||
)
|
||||
if context == "signature.parameters":
|
||||
assert isinstance(item, dict)
|
||||
assert "name" in item
|
||||
for key in item:
|
||||
assert key not in _FORBIDDEN_WIRE_KEYS
|
||||
assert key != "catalog_name"
|
||||
_assert_forbidden_keys_absent(
|
||||
{k: v for k, v in item.items() if k != "name"},
|
||||
context=item_context,
|
||||
)
|
||||
else:
|
||||
_assert_forbidden_keys_absent(item, context=item_context)
|
||||
|
||||
|
||||
def _assert_sanitized_text(*parts: object) -> None:
|
||||
combined = "\n".join(str(part) for part in parts)
|
||||
lowered = combined.lower()
|
||||
assert _SOURCE_MARKER.lower() not in lowered
|
||||
assert _SECRET_REFERENCE.lower() not in lowered
|
||||
assert str(Path(__file__).resolve()).lower() not in lowered
|
||||
assert Path(__file__).resolve().as_posix().lower() not in lowered
|
||||
|
||||
|
||||
def _assert_clean_validation_error(exc_info) -> None:
|
||||
_assert_sanitized_text(exc_info.value, repr(exc_info.value))
|
||||
assert exc_info.value.__cause__ is None
|
||||
assert exc_info.value.__context__ is None
|
||||
|
||||
|
||||
def _valid_builder_kwargs(**overrides):
|
||||
kwargs = {
|
||||
"parameters": [("text", pa.string()), ("limit", pa.int32())],
|
||||
"output_type": pa.string(),
|
||||
"output_nullable": True,
|
||||
"module": "bridge_mod",
|
||||
"callable_name": "normalize",
|
||||
"source": (
|
||||
"def normalize(text, limit):\n"
|
||||
f" # {_SOURCE_MARKER}\n"
|
||||
" return text[:limit]\n"
|
||||
),
|
||||
"python": "3.12",
|
||||
"packages": ["pkg-b==2", "pkg-a==1"],
|
||||
"capabilities": [
|
||||
("network", _NETWORK_ORIGIN, None),
|
||||
("secret", _SECRET_REFERENCE, _SECRET_ENV),
|
||||
("network", _NETWORK_ORIGIN_B, None),
|
||||
],
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return kwargs
|
||||
|
||||
|
||||
def test_build_function_definition_private_native_immutability_and_export_surface():
|
||||
assert "_build_function_definition" not in getattr(lancedb, "__all__", [])
|
||||
assert "_FunctionDefinition" not in lancedb.__all__
|
||||
assert not hasattr(lancedb, "_FunctionDefinition")
|
||||
assert not hasattr(lancedb, "_build_function_definition")
|
||||
assert not hasattr(lancedb, "_new_function_definition")
|
||||
|
||||
definition = _build_function_definition(packable_bridge_normalize)
|
||||
definition_type = _function_definition_type()
|
||||
assert type(definition) is definition_type
|
||||
assert definition_type.__module__ == "lancedb._lancedb"
|
||||
assert definition_type.__name__ == "_FunctionDefinition"
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
definition_type()
|
||||
|
||||
for attr in _OVERDESIGN_ATTRS:
|
||||
assert not hasattr(definition, attr)
|
||||
|
||||
for attr in ("signature", "module", "source", "capabilities"):
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(definition, attr, None)
|
||||
|
||||
|
||||
def test_build_function_definition_json_wire_ordered_contract_without_identity():
|
||||
definition = _build_function_definition(packable_bridge_normalize)
|
||||
encoded_a = _json_bytes(definition)
|
||||
encoded_b = _json_bytes(definition)
|
||||
assert encoded_a == encoded_b
|
||||
|
||||
wire = json.loads(encoded_a.decode("utf-8"))
|
||||
_assert_exact_object_keys(
|
||||
wire,
|
||||
{"format_version", "signature", "implementation", "capabilities"},
|
||||
context="definition",
|
||||
)
|
||||
assert wire["format_version"] == 1
|
||||
_assert_forbidden_keys_absent(wire, context="definition")
|
||||
|
||||
signature = wire["signature"]
|
||||
_assert_exact_object_keys(signature, {"parameters", "output"}, context="signature")
|
||||
parameters = signature["parameters"]
|
||||
assert [parameter["name"] for parameter in parameters] == ["text", "limit"]
|
||||
for parameter in parameters:
|
||||
_assert_exact_object_keys(
|
||||
parameter, {"name", "data_type_ipc"}, context="parameter"
|
||||
)
|
||||
assert isinstance(parameter["data_type_ipc"], str)
|
||||
assert parameter["data_type_ipc"]
|
||||
assert _decode_type_ipc(parameters[0]["data_type_ipc"]) == pa.string()
|
||||
assert _decode_type_ipc(parameters[1]["data_type_ipc"]) == pa.int32()
|
||||
|
||||
output = signature["output"]
|
||||
_assert_exact_object_keys(
|
||||
output, {"data_type_ipc", "nullable"}, context="signature.output"
|
||||
)
|
||||
assert output["nullable"] is True
|
||||
assert _decode_type_ipc(output["data_type_ipc"]) == pa.string()
|
||||
|
||||
implementation = wire["implementation"]
|
||||
_assert_exact_object_keys(
|
||||
implementation,
|
||||
{"kind", "module", "callable", "source", "python", "packages"},
|
||||
context="implementation",
|
||||
)
|
||||
assert implementation["kind"] == "python"
|
||||
assert implementation["module"] == __name__
|
||||
assert implementation["callable"] == "packable_bridge_normalize"
|
||||
assert implementation["source"] == Path(__file__).read_text(encoding="utf-8")
|
||||
assert _SOURCE_MARKER in implementation["source"]
|
||||
assert implementation["python"] == "3.12"
|
||||
assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"]
|
||||
|
||||
capabilities = wire["capabilities"]
|
||||
assert capabilities == [
|
||||
{"kind": "network", "origin": _NETWORK_ORIGIN},
|
||||
{
|
||||
"kind": "secret",
|
||||
"reference": _SECRET_REFERENCE,
|
||||
"environment_variable": _SECRET_ENV,
|
||||
},
|
||||
{"kind": "network", "origin": _NETWORK_ORIGIN_B},
|
||||
]
|
||||
for capability in capabilities:
|
||||
assert "value" not in capability
|
||||
assert "plaintext" not in capability
|
||||
assert "plaintext_secret" not in capability
|
||||
assert "secret_value" not in capability
|
||||
|
||||
|
||||
def test_native_definition_repr_includes_safe_structure_and_redacts_sensitive_text():
|
||||
definition = _build_function_definition(packable_bridge_normalize)
|
||||
rendered = repr(definition)
|
||||
assert "_FunctionDefinition" in rendered or "FunctionDefinition" in rendered
|
||||
assert __name__ in rendered
|
||||
assert "packable_bridge_normalize" in rendered
|
||||
assert "3.12" in rendered
|
||||
_assert_sanitized_text(rendered)
|
||||
|
||||
|
||||
def test_new_function_definition_builder_preserves_normalized_wire():
|
||||
definition = _new_function_definition(**_valid_builder_kwargs())
|
||||
assert type(definition) is _function_definition_type()
|
||||
|
||||
encoded_a = _json_bytes(definition)
|
||||
encoded_b = _json_bytes(definition)
|
||||
assert encoded_a == encoded_b
|
||||
|
||||
wire = json.loads(encoded_a.decode("utf-8"))
|
||||
assert wire["format_version"] == 1
|
||||
assert [parameter["name"] for parameter in wire["signature"]["parameters"]] == [
|
||||
"text",
|
||||
"limit",
|
||||
]
|
||||
assert _decode_type_ipc(wire["signature"]["parameters"][0]["data_type_ipc"]) == (
|
||||
pa.string()
|
||||
)
|
||||
assert _decode_type_ipc(wire["signature"]["parameters"][1]["data_type_ipc"]) == (
|
||||
pa.int32()
|
||||
)
|
||||
assert wire["signature"]["output"]["nullable"] is True
|
||||
assert _decode_type_ipc(wire["signature"]["output"]["data_type_ipc"]) == pa.string()
|
||||
|
||||
implementation = wire["implementation"]
|
||||
assert implementation["kind"] == "python"
|
||||
assert implementation["module"] == "bridge_mod"
|
||||
assert implementation["callable"] == "normalize"
|
||||
assert implementation["source"] == _valid_builder_kwargs()["source"]
|
||||
assert _SOURCE_MARKER in implementation["source"]
|
||||
assert implementation["python"] == "3.12"
|
||||
assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"]
|
||||
assert wire["capabilities"] == [
|
||||
{"kind": "network", "origin": _NETWORK_ORIGIN},
|
||||
{
|
||||
"kind": "secret",
|
||||
"reference": _SECRET_REFERENCE,
|
||||
"environment_variable": _SECRET_ENV,
|
||||
},
|
||||
{"kind": "network", "origin": _NETWORK_ORIGIN_B},
|
||||
]
|
||||
_assert_forbidden_keys_absent(wire, context="definition")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides",),
|
||||
[
|
||||
({"parameters": [("text", pa.string()), ("text", pa.int32())]},),
|
||||
({"parameters": [("", pa.string())]},),
|
||||
({"module": ""},),
|
||||
({"callable_name": ""},),
|
||||
({"source": ""},),
|
||||
({"python": ""},),
|
||||
({"packages": ["pkg-a==1", ""]},),
|
||||
({"packages": ["pkg-a==1", "pkg-a==1"]},),
|
||||
({"capabilities": [("filesystem", _NETWORK_ORIGIN, None)]},),
|
||||
({"capabilities": [("network", _NETWORK_ORIGIN, _SECRET_ENV)]},),
|
||||
({"capabilities": [("secret", _SECRET_REFERENCE, None)]},),
|
||||
({"capabilities": [("secret", _SECRET_REFERENCE, "")]},),
|
||||
({"capabilities": [("network", "", None)]},),
|
||||
({"capabilities": [("secret", "", _SECRET_ENV)]},),
|
||||
],
|
||||
)
|
||||
def test_new_function_definition_strict_validation_rejections(overrides):
|
||||
kwargs = _valid_builder_kwargs(**overrides)
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_new_function_definition(**kwargs)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
|
||||
def test_new_function_definition_validation_does_not_echo_secret_or_source_marker():
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_new_function_definition(**_valid_builder_kwargs(module=""))
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_new_function_definition(
|
||||
**_valid_builder_kwargs(packages=["pkg-a==1", "pkg-a==1"])
|
||||
)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_new_function_definition(
|
||||
**_valid_builder_kwargs(
|
||||
capabilities=[("secret", _SECRET_REFERENCE, None)],
|
||||
)
|
||||
)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides",),
|
||||
[
|
||||
({"parameters": [("text", "not-a-datatype")]},),
|
||||
({"parameters": [(123, pa.string())]},),
|
||||
({"output_type": "not-a-datatype"},),
|
||||
({"output_type": None},),
|
||||
({"output_nullable": "yes"},),
|
||||
({"packages": "pkg-a==1"},),
|
||||
({"capabilities": "network"},),
|
||||
({"capabilities": [("network", _NETWORK_ORIGIN)]},),
|
||||
({"capabilities": [("network", _NETWORK_ORIGIN, None, "extra")]},),
|
||||
],
|
||||
)
|
||||
def test_new_function_definition_wrong_pyarrow_and_shape_values_fail_closed(overrides):
|
||||
kwargs = _valid_builder_kwargs(**overrides)
|
||||
with pytest.raises((TypeError, ValueError)) as exc_info:
|
||||
_new_function_definition(**kwargs)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
|
||||
class _HostileRaisingIterable:
|
||||
def __iter__(self):
|
||||
raise RuntimeError(f"{_SECRET_REFERENCE} {_SOURCE_MARKER}")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides",),
|
||||
[
|
||||
({"parameters": _HostileRaisingIterable()},),
|
||||
({"packages": _HostileRaisingIterable()},),
|
||||
({"capabilities": _HostileRaisingIterable()},),
|
||||
(
|
||||
{
|
||||
"capabilities": [
|
||||
("network", _NETWORK_ORIGIN, None),
|
||||
_HostileRaisingIterable(),
|
||||
("network", _NETWORK_ORIGIN_B, None),
|
||||
]
|
||||
},
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_new_function_definition_hostile_iterable_iter_raises_fail_closed(overrides):
|
||||
kwargs = _valid_builder_kwargs(**overrides)
|
||||
with pytest.raises((TypeError, ValueError)) as exc_info:
|
||||
_new_function_definition(**kwargs)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=False,
|
||||
)
|
||||
def packable_bridge_capability_exact_type(x):
|
||||
return x + 1
|
||||
|
||||
|
||||
def test_build_function_definition_rejects_forged_function_capability_subclass():
|
||||
marker = f"{_SECRET_REFERENCE} {_SOURCE_MARKER}"
|
||||
|
||||
class _HostileFunctionCapability(FunctionCapability):
|
||||
@property
|
||||
def kind(self) -> str:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def origin(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def reference(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def environment_variable(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
hostile = object.__new__(_HostileFunctionCapability)
|
||||
assert isinstance(hostile, FunctionCapability)
|
||||
assert type(hostile) is not FunctionCapability
|
||||
|
||||
config_attr = _udf_mod._CONFIG_ATTR
|
||||
original = getattr(packable_bridge_capability_exact_type, config_attr)
|
||||
forged = _udf_mod._UdfConfig(
|
||||
inputs=original.inputs,
|
||||
output=original.output,
|
||||
output_nullable=original.output_nullable,
|
||||
python=original.python,
|
||||
packages=original.packages,
|
||||
capabilities=(hostile,),
|
||||
)
|
||||
setattr(packable_bridge_capability_exact_type, config_attr, forged)
|
||||
try:
|
||||
with pytest.raises((TypeError, ValueError)) as exc_info:
|
||||
_build_function_definition(packable_bridge_capability_exact_type)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
assert marker not in str(exc_info.value)
|
||||
assert marker not in repr(exc_info.value)
|
||||
finally:
|
||||
setattr(packable_bridge_capability_exact_type, config_attr, original)
|
||||
@@ -0,0 +1,486 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for private UDF packaging validation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import inspect
|
||||
import json
|
||||
import sys
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
from lancedb import Function, Job, udf
|
||||
from lancedb._udf import _get_udf_config, _package_udf
|
||||
|
||||
_BODY_MARKER = "packaging body marker unique-xyz"
|
||||
_AMBIENT_SECRET = "ambient-secret-value-xyz"
|
||||
_BUILTIN_SHADOW_SECRET = "builtin-shadow-secret-xyz"
|
||||
_SOURCE_MISMATCH_SECRET = "source-mismatch-secret-xyz"
|
||||
_INVALID_UTF8_SECRET = "invalid-utf8-secret-xyz"
|
||||
|
||||
_OVERDESIGN_ATTRS = (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_handling",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"error_policy",
|
||||
"FunctionVersion",
|
||||
"function_version",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
"id",
|
||||
"function_id",
|
||||
"job",
|
||||
"job_id",
|
||||
"registration",
|
||||
"catalog",
|
||||
"retry_key",
|
||||
"source_path",
|
||||
"path",
|
||||
"function",
|
||||
)
|
||||
|
||||
_PACKAGING_CONSTANT = 41
|
||||
|
||||
|
||||
def _packaging_helper(value: int) -> int:
|
||||
return value + _PACKAGING_CONSTANT
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=False,
|
||||
)
|
||||
def packable_add(x):
|
||||
"""packaging body marker unique-xyz."""
|
||||
return _packaging_helper(x) + len(json.dumps({"k": 1}))
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def packable_kwonly(x, *, y=2):
|
||||
return x + y
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def packable_rebind_target(x):
|
||||
return x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def uses_injected_ambient(x):
|
||||
return x + len(INJECTED_AMBIENT_GLOBAL) # noqa: F821
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def uses_shadowed_builtin_len(x):
|
||||
return x + len((1, 2, 3))
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def mismatch_names(left, right):
|
||||
return left + right
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"y": pa.int32(), "x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def mismatch_order(x, y):
|
||||
return x + y
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def positional_only(x, /, y):
|
||||
return x + y
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def varargs_fn(x, *args):
|
||||
return x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def kwargs_fn(x, **kwargs):
|
||||
return x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
async def async_fn(x):
|
||||
return x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
async def async_gen_fn(x):
|
||||
yield x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def generator_fn(x):
|
||||
yield x
|
||||
|
||||
|
||||
def _assert_sanitized_text(*parts: object, secret: str = _AMBIENT_SECRET) -> None:
|
||||
combined = "\n".join(str(part) for part in parts)
|
||||
lowered = combined.lower()
|
||||
assert _BODY_MARKER.lower() not in lowered
|
||||
assert secret.lower() not in lowered
|
||||
assert str(Path(__file__).resolve()).lower() not in lowered
|
||||
assert Path(__file__).resolve().as_posix().lower() not in lowered
|
||||
|
||||
|
||||
def _assert_packaging_rejection(exc_info, *, secret: str = _AMBIENT_SECRET) -> None:
|
||||
_assert_sanitized_text(exc_info.value, repr(exc_info.value), secret=secret)
|
||||
assert exc_info.value.__cause__ is None
|
||||
assert exc_info.value.__context__ is None
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _temporary_imported_module(
|
||||
directory: Path, module_name: str, source: str
|
||||
) -> Iterator[tuple[Path, object]]:
|
||||
path = directory / f"{module_name}.py"
|
||||
path.write_text(source, encoding="utf-8")
|
||||
inserted = str(directory)
|
||||
sys.path.insert(0, inserted)
|
||||
try:
|
||||
sys.modules.pop(module_name, None)
|
||||
module = importlib.import_module(module_name)
|
||||
yield path, module
|
||||
finally:
|
||||
sys.modules.pop(module_name, None)
|
||||
try:
|
||||
sys.path.remove(inserted)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
|
||||
def _temp_udf_module_source(*, body: str, secret: str | None = None) -> str:
|
||||
secret_line = f"_SECRET = {secret!r}\n" if secret is not None else ""
|
||||
return (
|
||||
"import pyarrow as pa\n"
|
||||
"from lancedb import udf\n"
|
||||
f"{secret_line}\n"
|
||||
"@udf(\n"
|
||||
' inputs={"x": pa.int32()},\n'
|
||||
" output=pa.int32(),\n"
|
||||
' python="3.12",\n'
|
||||
")\n"
|
||||
"def temp_pack_target(x):\n"
|
||||
f" {body}\n"
|
||||
)
|
||||
|
||||
|
||||
def test_package_udf_success_snapshot_source_module_callable_config_and_repr():
|
||||
packaged = _package_udf(packable_add)
|
||||
source = Path(__file__).read_text(encoding="utf-8")
|
||||
|
||||
assert packaged.source == source
|
||||
assert packaged.module == __name__
|
||||
assert packaged.module != "__main__"
|
||||
assert packaged.callable_name == "packable_add"
|
||||
assert packable_add.__qualname__ == "packable_add"
|
||||
assert packaged.config is _get_udf_config(packable_add)
|
||||
assert packaged.config.inputs == (("x", pa.int32()),)
|
||||
assert packaged.config.output == pa.int64()
|
||||
assert packaged.config.output_nullable is False
|
||||
assert packaged.config.python == "3.12"
|
||||
assert packaged.config.packages == ("pkg-a==1",)
|
||||
|
||||
for attr in ("source", "module", "callable_name", "config"):
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(packaged, attr, None)
|
||||
|
||||
text = repr(packaged)
|
||||
_assert_sanitized_text(text)
|
||||
assert _BODY_MARKER not in text
|
||||
|
||||
|
||||
def test_package_udf_allows_source_bound_import_constant_and_helper():
|
||||
packaged = _package_udf(packable_add)
|
||||
assert packaged.callable_name == "packable_add"
|
||||
assert "import json" in packaged.source
|
||||
assert "_PACKAGING_CONSTANT" in packaged.source
|
||||
assert "_packaging_helper" in packaged.source
|
||||
assert packable_add(1) == _packaging_helper(1) + len(json.dumps({"k": 1}))
|
||||
|
||||
|
||||
def test_package_udf_accepts_positional_or_keyword_and_keyword_only_defaults():
|
||||
packaged = _package_udf(packable_kwonly)
|
||||
assert packaged.callable_name == "packable_kwonly"
|
||||
assert packaged.config.inputs == (("x", pa.int32()), ("y", pa.int32()))
|
||||
assert str(inspect.signature(packable_kwonly)) == "(x, *, y=2)"
|
||||
assert packable_kwonly(3) == 5
|
||||
assert packable_kwonly(3, y=7) == 10
|
||||
|
||||
|
||||
def test_package_udf_rejects_lambda_and_closure():
|
||||
lam = udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(lambda n: n + 1)
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(lam)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
ambient = _AMBIENT_SECRET
|
||||
|
||||
def factory(offset):
|
||||
@udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def closed(n):
|
||||
return n + offset + len(ambient)
|
||||
|
||||
return closed
|
||||
|
||||
closed = factory(10)
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(closed)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
def outer():
|
||||
total = 0
|
||||
|
||||
@udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def nested(n):
|
||||
nonlocal total
|
||||
total += n
|
||||
return total
|
||||
|
||||
return nested
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(outer())
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_signature_mismatches_and_unsupported_parameter_kinds():
|
||||
for target in (
|
||||
mismatch_names,
|
||||
mismatch_order,
|
||||
positional_only,
|
||||
varargs_fn,
|
||||
kwargs_fn,
|
||||
):
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(target)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_async_and_generator_functions():
|
||||
for target in (async_fn, async_gen_fn, generator_fn):
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(target)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_dynamic_exec_source():
|
||||
namespace: dict[str, object] = {}
|
||||
exec(
|
||||
"def dynamic_pack_target(x):\n return x + 1\n",
|
||||
namespace,
|
||||
)
|
||||
dynamic = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(namespace["dynamic_pack_target"])
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(dynamic)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_undecorated_and_wrong_input_types():
|
||||
def plain(x):
|
||||
return x
|
||||
|
||||
with pytest.raises(TypeError) as exc_info:
|
||||
_package_udf(plain)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
with pytest.raises(TypeError) as exc_info:
|
||||
_package_udf(object())
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
with pytest.raises(TypeError) as exc_info:
|
||||
_package_udf(42)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_rebound_module_attribute():
|
||||
module = sys.modules[__name__]
|
||||
original = module.packable_rebind_target
|
||||
replacement = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(lambda x: x)
|
||||
module.packable_rebind_target = replacement
|
||||
try:
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(original)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
finally:
|
||||
module.packable_rebind_target = original
|
||||
|
||||
|
||||
def test_package_udf_rejects_injected_ambient_global():
|
||||
module = sys.modules[__name__]
|
||||
secret = _AMBIENT_SECRET
|
||||
module.INJECTED_AMBIENT_GLOBAL = secret
|
||||
try:
|
||||
assert uses_injected_ambient(3) == 3 + len(secret)
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(uses_injected_ambient)
|
||||
_assert_packaging_rejection(exc_info, secret=secret)
|
||||
finally:
|
||||
delattr(module, "INJECTED_AMBIENT_GLOBAL")
|
||||
|
||||
|
||||
def test_package_udf_rejects_builtin_shadow_injection():
|
||||
module = sys.modules[__name__]
|
||||
secret = _BUILTIN_SHADOW_SECRET
|
||||
assert not hasattr(module, "len")
|
||||
module.len = secret
|
||||
try:
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(uses_shadowed_builtin_len)
|
||||
_assert_packaging_rejection(exc_info, secret=secret)
|
||||
finally:
|
||||
delattr(module, "len")
|
||||
|
||||
|
||||
def test_package_udf_rejects_loaded_code_source_mismatch(tmp_path: Path):
|
||||
secret = _SOURCE_MISMATCH_SECRET
|
||||
module_name = "udf_pkg_source_mismatch_mod"
|
||||
original = _temp_udf_module_source(body="return x + 1")
|
||||
replacement = _temp_udf_module_source(
|
||||
body=f"return x + 99 # {secret}",
|
||||
secret=secret,
|
||||
)
|
||||
with _temporary_imported_module(tmp_path, module_name, original) as (
|
||||
path,
|
||||
module,
|
||||
):
|
||||
target = module.temp_pack_target
|
||||
assert target(1) == 2
|
||||
path.write_text(replacement, encoding="utf-8")
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(target)
|
||||
_assert_packaging_rejection(exc_info, secret=secret)
|
||||
err_text = f"{exc_info.value}\n{exc_info.value!r}"
|
||||
assert str(path.resolve()) not in err_text
|
||||
assert path.resolve().as_posix() not in err_text
|
||||
|
||||
|
||||
def test_package_udf_rejects_invalid_utf8_after_import(tmp_path: Path):
|
||||
secret = _INVALID_UTF8_SECRET
|
||||
module_name = "udf_pkg_invalid_utf8_mod"
|
||||
original = _temp_udf_module_source(body="return x + 1")
|
||||
with _temporary_imported_module(tmp_path, module_name, original) as (
|
||||
path,
|
||||
module,
|
||||
):
|
||||
target = module.temp_pack_target
|
||||
assert target(1) == 2
|
||||
path.write_bytes(secret.encode("utf-8") + b"\xff\xfe invalid-bytes")
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(target)
|
||||
assert type(exc_info.value) is ValueError
|
||||
assert exc_info.value.__cause__ is None
|
||||
assert exc_info.value.__context__ is None
|
||||
err_text = f"{exc_info.value}\n{exc_info.value!r}"
|
||||
assert secret not in err_text
|
||||
assert "b'" not in err_text
|
||||
assert r"\xff" not in err_text
|
||||
assert str(path.resolve()) not in err_text
|
||||
assert path.resolve().as_posix() not in err_text
|
||||
|
||||
|
||||
def test_package_udf_snapshot_has_no_durable_overdesign_and_is_not_function_or_job():
|
||||
packaged = _package_udf(packable_add)
|
||||
assert not isinstance(packaged, Function)
|
||||
assert not isinstance(packaged, Job)
|
||||
for attr in _OVERDESIGN_ATTRS:
|
||||
assert not hasattr(packaged, attr)
|
||||
|
||||
text = repr(packaged).lower()
|
||||
for token in (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"functionversion",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
"retry_key",
|
||||
):
|
||||
assert token not in text
|
||||
_assert_sanitized_text(text)
|
||||
@@ -12,7 +12,7 @@ import pyarrow.compute as pc
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
from lancedb.index import FTS
|
||||
from lancedb.index import BTree, FTS, IvfPq
|
||||
from lancedb.table import AsyncTable, Table
|
||||
|
||||
|
||||
@@ -99,6 +99,86 @@ async def test_async_hybrid_query_filters(table: AsyncTable):
|
||||
assert result["text"].to_pylist() == ["cat", "b"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hybrid_query_with_stale_fixed_size_binary_prefilter(
|
||||
tmpdir_factory,
|
||||
):
|
||||
tmp_path = str(tmpdir_factory.mktemp("stale_scalar_prefilter"))
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
|
||||
def fixed_size_binary(value: int) -> bytes:
|
||||
return value.to_bytes(16, byteorder="big")
|
||||
|
||||
num_rows = 1000
|
||||
data = pa.table(
|
||||
{
|
||||
"space_id": pa.array(
|
||||
[fixed_size_binary(i) for i in range(num_rows)],
|
||||
type=pa.binary(16),
|
||||
),
|
||||
"text": ["book"] * num_rows,
|
||||
"vector": pa.array(
|
||||
[[float(i), float(i)] for i in range(num_rows)],
|
||||
type=pa.list_(pa.float32(), 2),
|
||||
),
|
||||
}
|
||||
)
|
||||
table = await db.create_table("test", data)
|
||||
await table.create_index(
|
||||
"vector", config=IvfPq(num_partitions=4, num_sub_vectors=2)
|
||||
)
|
||||
await table.create_index("space_id", config=BTree())
|
||||
await table.create_index("text", config=FTS(with_position=False))
|
||||
|
||||
# Advance the search indices without advancing the scalar index. This is the
|
||||
# state that previously let hybrid search use an incomplete scalar prefilter.
|
||||
await table.add(data)
|
||||
lance_dataset = await table.to_lance()
|
||||
lance_dataset.optimize.optimize_indices(index_names=["vector_idx", "text_idx"])
|
||||
await table.checkout_latest()
|
||||
|
||||
scalar_stats = await table.index_stats("space_id_idx")
|
||||
assert scalar_stats is not None
|
||||
assert scalar_stats.num_indexed_rows == num_rows
|
||||
assert scalar_stats.num_unindexed_rows == num_rows
|
||||
|
||||
for index_name in ["vector_idx", "text_idx"]:
|
||||
search_stats = await table.index_stats(index_name)
|
||||
assert search_stats is not None
|
||||
assert search_stats.num_indexed_rows == num_rows * 2
|
||||
assert search_stats.num_unindexed_rows == 0
|
||||
|
||||
matching_ids = [5, 10, 15, 20, 25, 30]
|
||||
literals = [
|
||||
f"arrow_cast(0x{fixed_size_binary(i).hex()}, 'FixedSizeBinary(16)')"
|
||||
for i in matching_ids
|
||||
]
|
||||
predicate = f"space_id IN ({', '.join(literals)})"
|
||||
expected_ids = sorted(fixed_size_binary(i) for i in matching_ids for _ in range(2))
|
||||
|
||||
vector_query = (
|
||||
table.query().where(predicate).nearest_to([5.0, 5.0]).limit(num_rows * 2)
|
||||
)
|
||||
vector_results = await vector_query.to_arrow()
|
||||
assert sorted(vector_results["space_id"].to_pylist()) == expected_ids
|
||||
|
||||
fts_query = (
|
||||
table.query().where(predicate).nearest_to_text("book").limit(num_rows * 2)
|
||||
)
|
||||
fts_results = await fts_query.to_arrow()
|
||||
assert sorted(fts_results["space_id"].to_pylist()) == expected_ids
|
||||
|
||||
hybrid_results = await (
|
||||
table.query()
|
||||
.where(predicate)
|
||||
.nearest_to([5.0, 5.0])
|
||||
.nearest_to_text("book")
|
||||
.limit(num_rows * 2)
|
||||
.to_arrow()
|
||||
)
|
||||
assert sorted(hybrid_results["space_id"].to_pylist()) == expected_ids
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_hybrid_query_default_limit(table: AsyncTable):
|
||||
# add 10 new rows
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
import lancedb._lancedb as _lancedb
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.platform != "linux", reason="ldd is Linux-specific")
|
||||
def test_native_extension_does_not_link_openssl():
|
||||
"""OpenSSL-linked wheels abort when imported on RHEL hosts in FIPS mode."""
|
||||
ldd = shutil.which("ldd")
|
||||
if ldd is None:
|
||||
pytest.skip("ldd is not installed")
|
||||
|
||||
result = subprocess.run(
|
||||
[ldd, _lancedb.__file__],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
openssl_libraries = re.findall(
|
||||
r"^\s*(lib(?:crypto|ssl)\S*)\s+=>", result.stdout, flags=re.MULTILINE
|
||||
)
|
||||
|
||||
assert not openssl_libraries, (
|
||||
"the LanceDB native extension must use rustls instead of linking OpenSSL: "
|
||||
f"{openssl_libraries}"
|
||||
)
|
||||
@@ -372,6 +372,31 @@ async def test_create_vector_index(some_table: AsyncTable):
|
||||
assert stats.num_indices == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_ivf_index_reports_unsplittable_partitions(db_async):
|
||||
dim = 8
|
||||
num_partitions = 300 # More than 256 selects hierarchical k-means.
|
||||
base_vectors = [[float(row == column) for column in range(dim)] for row in range(5)]
|
||||
vectors = pa.array(base_vectors * 200, pa.list_(pa.float32(), dim))
|
||||
table = await db_async.create_table(
|
||||
"unsplittable_partitions",
|
||||
pa.table({"vector": vectors}),
|
||||
)
|
||||
|
||||
error_pattern = (
|
||||
rf"Cannot create {num_partitions} IVF partitions: k-means could only form"
|
||||
)
|
||||
with pytest.raises(RuntimeError, match=error_pattern):
|
||||
await table.create_index(
|
||||
"vector",
|
||||
config=IvfFlat(
|
||||
distance_type="dot",
|
||||
num_partitions=num_partitions,
|
||||
max_iterations=10,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_4bit_ivfpq_index(some_table: AsyncTable):
|
||||
# Can create
|
||||
|
||||
@@ -83,7 +83,9 @@ def test_lsm_write_spec_repr():
|
||||
assert s.spec_type == "bucket"
|
||||
assert s.column == "id"
|
||||
assert s.num_buckets == 4
|
||||
assert s.maintained_indexes == []
|
||||
# A fresh spec defers its maintained set to install time.
|
||||
assert s.maintained_indexes is None
|
||||
assert s.with_maintained_indexes([]).maintained_indexes == []
|
||||
assert "bucket" in repr(s)
|
||||
assert "id" in repr(s)
|
||||
assert "4" in repr(s)
|
||||
@@ -169,18 +171,23 @@ def test_get_lsm_write_spec(tmp_path):
|
||||
table.unset_lsm_write_spec()
|
||||
assert table.get_lsm_write_spec() is None
|
||||
|
||||
# Identity round-trips (column recovered from the schema).
|
||||
# Identity round-trips (column recovered from the schema). Leaving the
|
||||
# maintained set to be inferred picks up the index on the table, so the
|
||||
# spec reads back naming it rather than as "infer".
|
||||
table.set_lsm_write_spec(LsmWriteSpec.identity("id"))
|
||||
spec = table.get_lsm_write_spec()
|
||||
assert spec.spec_type == "identity"
|
||||
assert spec.column == "id"
|
||||
assert spec.maintained_indexes == [idx_name]
|
||||
table.unset_lsm_write_spec()
|
||||
|
||||
# Unsharded round-trips (no routing column).
|
||||
table.set_lsm_write_spec(LsmWriteSpec.unsharded())
|
||||
# Unsharded round-trips (no routing column). Opting out is distinct from
|
||||
# the inferred default.
|
||||
table.set_lsm_write_spec(LsmWriteSpec.unsharded().with_maintained_indexes([]))
|
||||
spec = table.get_lsm_write_spec()
|
||||
assert spec.spec_type == "unsharded"
|
||||
assert spec.column is None
|
||||
assert spec.maintained_indexes == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -544,7 +544,7 @@ def test_lsm_read_fts_unmaintained_index_errors(tmp_path):
|
||||
table.create_index("text", config=FTS())
|
||||
# No maintained indexes: the active memtable FTS arm cannot serve un-compacted
|
||||
# docs, so the search would silently omit them — reject instead.
|
||||
table.set_lsm_write_spec(LsmWriteSpec.unsharded())
|
||||
table.set_lsm_write_spec(LsmWriteSpec.unsharded().with_maintained_indexes([]))
|
||||
with pytest.raises(Exception, match="maintained"):
|
||||
table.search("fox", query_type="fts", fts_columns="text").to_arrow()
|
||||
|
||||
@@ -631,7 +631,7 @@ def test_lsm_read_vector_unmaintained_index_errors(tmp_path):
|
||||
)
|
||||
# Spec with NO maintained indexes: the base vector index's catch-up is untracked,
|
||||
# so the scanner rejects rather than risk dropping compacted-but-unindexed rows.
|
||||
table.set_lsm_write_spec(LsmWriteSpec.unsharded())
|
||||
table.set_lsm_write_spec(LsmWriteSpec.unsharded().with_maintained_indexes([]))
|
||||
with pytest.raises(Exception, match="maintained"):
|
||||
table.search([1.0] * VECTOR_DIM).to_arrow()
|
||||
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
import importlib
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def test_pyo3_abi_matches_minimum_supported_python():
|
||||
project_dir = Path(__file__).parents[2]
|
||||
pyproject = (project_dir / "pyproject.toml").read_text()
|
||||
cargo_manifest = (project_dir / "Cargo.toml").read_text()
|
||||
|
||||
minimum_python = re.search(
|
||||
r'^requires-python\s*=\s*">=(\d+)\.(\d+)"$', pyproject, re.MULTILINE
|
||||
)
|
||||
assert minimum_python is not None
|
||||
|
||||
major, minor = minimum_python.groups()
|
||||
expected_abi = f"abi3-py{major}{minor}"
|
||||
configured_abis = re.findall(r'"(abi3-py\d+)"', cargo_manifest)
|
||||
|
||||
assert configured_abis == [expected_abi, expected_abi], (
|
||||
"the pyo3 runtime and build ABI features must both match requires-python"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.platform != "win32", reason="Windows wheel regression test")
|
||||
def test_windows_wheel_tag_and_native_import():
|
||||
project_dir = Path(__file__).parents[2]
|
||||
wheels = list((project_dir.parent / "target" / "wheels").glob("lancedb-*.whl"))
|
||||
if not wheels:
|
||||
pytest.skip("no wheel artifact is available in this development environment")
|
||||
|
||||
assert len(wheels) == 1
|
||||
assert wheels[0].name.endswith("-cp310-abi3-win_amd64.whl")
|
||||
|
||||
native_module = importlib.import_module("lancedb._lancedb")
|
||||
assert Path(native_module.__file__).suffix == ".pyd"
|
||||
@@ -415,6 +415,17 @@ def test_nullable_vector():
|
||||
assert schema == pa.schema([pa.field("vec", pa.list_(pa.float32(), 16), True)])
|
||||
|
||||
|
||||
def test_bare_vector_raises_clear_error():
|
||||
namespace = {
|
||||
"__name__": "test_model_without_pyarrow",
|
||||
"LanceModel": LanceModel,
|
||||
"Vector": Vector,
|
||||
}
|
||||
|
||||
with pytest.raises(TypeError, match=r"Vector must be parameterized.*Vector\(128\)"):
|
||||
exec("class TestModel(LanceModel):\n vector: Vector", namespace)
|
||||
|
||||
|
||||
def test_fixed_size_list_field():
|
||||
class TestModel(pydantic.BaseModel):
|
||||
vec: Vector(16)
|
||||
|
||||
@@ -570,6 +570,15 @@ def test_query_builder(table):
|
||||
assert all(np.array(rs[0]["vector"]) == [1, 2])
|
||||
|
||||
|
||||
def test_query_multiple_vectors(table):
|
||||
results = table.search([np.array([1, 2]), np.array([4, 5])]).limit(1).to_list()
|
||||
|
||||
assert len(results) == 2
|
||||
results_by_query = {result["query_index"]: result for result in results}
|
||||
assert results_by_query[0]["id"] == 1
|
||||
assert results_by_query[1]["id"] == 2
|
||||
|
||||
|
||||
def test_with_row_id(table: lancedb.table.Table):
|
||||
rs = table.search().with_row_id(True).to_arrow()
|
||||
assert "_rowid" in rs.column_names
|
||||
|
||||
@@ -35,6 +35,12 @@ def make_mock_http_handler(handler):
|
||||
return MockLanceDBHandler
|
||||
|
||||
|
||||
@pytest.mark.parametrize("db_name", ["a" * 64, "invalid..database"])
|
||||
def test_connect_rejects_invalid_cloud_dns_hostname(db_name):
|
||||
with pytest.raises(ValueError, match="DNS labels must contain 1 to 63 bytes"):
|
||||
lancedb.connect(f"db://{db_name}", api_key="fake")
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def mock_lancedb_connection(handler):
|
||||
with http.server.HTTPServer(
|
||||
@@ -2300,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
|
||||
|
||||
@@ -2,10 +2,13 @@
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
|
||||
import ctypes
|
||||
import gc
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
import warnings
|
||||
import weakref
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import date, datetime, timedelta
|
||||
from time import sleep
|
||||
@@ -99,6 +102,30 @@ def test_basic(mem_db: DBConnection):
|
||||
assert table.to_arrow() == expected_data
|
||||
|
||||
|
||||
def test_search_preserves_nulls_from_sliced_arrow_table(mem_db: DBConnection):
|
||||
data = pa.table(
|
||||
{
|
||||
"id": [0, 1, 2, 3, 4],
|
||||
"score_cn": [None, 22, None, 5, 8],
|
||||
"score_mt": [None, 42, None, 5, 8],
|
||||
"vector": [
|
||||
[20, 19, -1, -1],
|
||||
[41, 38, 22, 42],
|
||||
[10, 10, -1, -1],
|
||||
[5, 5, 5, 5],
|
||||
[8, 8, 8, 8],
|
||||
],
|
||||
}
|
||||
).slice(1)
|
||||
|
||||
table = mem_db.create_table("sliced_nullable", data=data)
|
||||
result = table.search([41, 38, 22, 42]).limit(1).to_arrow()
|
||||
|
||||
assert result["id"].to_pylist() == [1]
|
||||
assert result["score_cn"].to_pylist() == [22]
|
||||
assert result["score_mt"].to_pylist() == [42]
|
||||
|
||||
|
||||
def test_table_to_pandas_default_matches_arrow(tmp_db: DBConnection):
|
||||
pd = pytest.importorskip("pandas")
|
||||
data = pa.table({"id": [1, 2], "text": ["one", "two"]})
|
||||
@@ -435,6 +462,38 @@ def test_add(mem_db: DBConnection):
|
||||
_add(table, schema)
|
||||
|
||||
|
||||
def test_add_releases_arrow_buffers_without_gc(mem_db: DBConnection):
|
||||
"""Regression test for https://github.com/lancedb/lancedb/issues/2512."""
|
||||
schema = pa.schema([pa.field("x", pa.int64())])
|
||||
table = mem_db.create_table("test_add_releases_arrow_buffers", schema=schema)
|
||||
|
||||
class BufferOwner:
|
||||
def __init__(self, size: int):
|
||||
self.memory = ctypes.create_string_buffer(size)
|
||||
|
||||
owner_refs = []
|
||||
gc_was_enabled = gc.isenabled()
|
||||
gc.disable()
|
||||
try:
|
||||
for _ in range(3):
|
||||
size = 8 * 1024
|
||||
owner = BufferOwner(size)
|
||||
arrow_buffer = pa.foreign_buffer(
|
||||
ctypes.addressof(owner.memory), size, owner
|
||||
)
|
||||
array = pa.Array.from_buffers(pa.int64(), 1024, [None, arrow_buffer])
|
||||
batch = pa.RecordBatch.from_arrays([array], schema=schema)
|
||||
owner_refs.append(weakref.ref(owner))
|
||||
|
||||
table.add(batch)
|
||||
del batch, array, arrow_buffer, owner
|
||||
|
||||
assert all(owner_ref() is None for owner_ref in owner_refs)
|
||||
finally:
|
||||
if gc_was_enabled:
|
||||
gc.enable()
|
||||
|
||||
|
||||
def test_add_write_parallelism(mem_db: DBConnection):
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = mem_db.create_table("test", schema=schema)
|
||||
@@ -870,6 +929,7 @@ def test_polars(mem_db: DBConnection):
|
||||
|
||||
# enter table to polars dataframe
|
||||
result = table.to_polars()
|
||||
assert isinstance(result, pl.LazyFrame)
|
||||
assert np.allclose(result.collect()["vector"].to_list(), data["vector"])
|
||||
|
||||
# make sure filtering isn't broken
|
||||
@@ -1786,6 +1846,27 @@ def test_add_with_empty_fixed_size_list_drops_bad_rows(mem_db: DBConnection):
|
||||
assert np.allclose(data["embedding"].to_pylist()[0], np.array([0.1] * 16))
|
||||
|
||||
|
||||
def test_add_nullable_fixed_size_list_with_none(mem_db: DBConnection):
|
||||
"""Regression test for issue #2340."""
|
||||
table = mem_db.create_table(
|
||||
"test_nullable_fixed_size_list",
|
||||
schema=pa.schema(
|
||||
[
|
||||
pa.field("id", pa.string()),
|
||||
pa.field("feature", pa.list_(pa.float32(), 256)),
|
||||
pa.field("tags", pa.list_(pa.string())),
|
||||
]
|
||||
),
|
||||
)
|
||||
|
||||
table.add([{"id": "1", "feature": None, "tags": ["tag1", "tag2"]}])
|
||||
|
||||
result = table.to_arrow()
|
||||
assert result.to_pylist() == [
|
||||
{"id": "1", "feature": None, "tags": ["tag1", "tag2"]}
|
||||
]
|
||||
|
||||
|
||||
def test_add_nullable_struct_with_none(mem_db: DBConnection):
|
||||
"""Regression test for issue #2654: a nullable struct column whose
|
||||
first batch contains only None values must not crash in
|
||||
@@ -1825,6 +1906,33 @@ def test_add_nullable_struct_with_none(mem_db: DBConnection):
|
||||
assert result.column("data").to_pylist() == [{"x": 1.0}, None]
|
||||
|
||||
|
||||
def test_read_mostly_null_list_v2_2_page_boundary(tmp_path):
|
||||
# Regression test for #3194. This row/value count crosses a v2.2 structural
|
||||
# encoding page boundary where Lance 3.0.0 sliced repetition/definition
|
||||
# levels by row offset and decoded child arrays at different lengths.
|
||||
num_rows = 64_885
|
||||
num_values = 217
|
||||
list_type = pa.list_(pa.float32())
|
||||
source = pa.table(
|
||||
{
|
||||
"id": np.arange(num_rows, dtype=np.int64),
|
||||
"coords": pa.array(
|
||||
[[1.0, 2.0, 3.0, 4.0]] * num_values + [None] * (num_rows - num_values),
|
||||
type=list_type,
|
||||
),
|
||||
}
|
||||
)
|
||||
db = lancedb.connect(
|
||||
tmp_path,
|
||||
storage_options={"new_table_data_storage_version": "2.2"},
|
||||
)
|
||||
table = db.create_table("test_sparse_nullable_list", data=source)
|
||||
|
||||
result = table.search().select(["id", "coords"]).limit(num_rows).to_arrow()
|
||||
|
||||
assert result.equals(source)
|
||||
|
||||
|
||||
def test_add_with_integer_embeddings_preserves_casting(mem_db: DBConnection):
|
||||
class Schema(LanceModel):
|
||||
text: str
|
||||
@@ -2110,6 +2218,45 @@ def test_merge(tmp_db: DBConnection, tmp_path):
|
||||
table.merge(other_dataset, left_on="id")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("storage_version", ["legacy", "stable"])
|
||||
def test_search_after_merge(tmp_path, storage_version):
|
||||
pytest.importorskip("lance")
|
||||
pd = pytest.importorskip("pandas")
|
||||
|
||||
db = lancedb.connect(
|
||||
tmp_path,
|
||||
storage_options={"new_table_data_storage_version": storage_version},
|
||||
)
|
||||
rng = np.random.default_rng(42)
|
||||
row_count = 512
|
||||
vectors = rng.standard_normal((row_count, 8)).astype(np.float32)
|
||||
table = db.create_table(
|
||||
"search_after_merge",
|
||||
data=pd.DataFrame(
|
||||
{
|
||||
"id": [str(i) for i in range(row_count)],
|
||||
"vector": list(vectors),
|
||||
}
|
||||
),
|
||||
)
|
||||
table.create_index("vector", config=IvfPq(num_partitions=1, num_sub_vectors=2))
|
||||
|
||||
links = pd.DataFrame(
|
||||
{
|
||||
"id": [str(i) for i in range(row_count // 2)],
|
||||
"link": [f"https://example.com/{i}" for i in range(row_count // 2)],
|
||||
}
|
||||
)
|
||||
table.merge(links, left_on="id")
|
||||
|
||||
query = table.search(vectors[-1]).refine_factor(50).limit(10)
|
||||
assert "ANN" in query.explain_plan(verbose=True)
|
||||
|
||||
result = query.to_arrow()
|
||||
links_by_id = dict(zip(result["id"].to_pylist(), result["link"].to_pylist()))
|
||||
assert links_by_id[str(row_count - 1)] is None
|
||||
|
||||
|
||||
def test_delete(mem_db: DBConnection):
|
||||
table = mem_db.create_table(
|
||||
"my_table",
|
||||
@@ -2196,6 +2343,20 @@ def test_update(mem_db: DBConnection):
|
||||
assert np.allclose(v, np.array([[1.2, 1.9], [1.1, 1.1]]))
|
||||
|
||||
|
||||
def test_update_with_arrow_scalar(mem_db: DBConnection):
|
||||
schema = pa.schema({"id": pa.int64(), "vector": pa.list_(pa.float32(), 4)})
|
||||
table = mem_db.create_table("my_table", schema=schema)
|
||||
table.add([{"id": 1, "vector": [1.0, 2.0, 3.0, 4.0]}])
|
||||
|
||||
value = table.search().select(["vector"]).limit(1).to_arrow()["vector"][0]
|
||||
assert isinstance(value, pa.FixedSizeListScalar)
|
||||
|
||||
result = table.update(where="id == 1", values={"vector": value})
|
||||
|
||||
assert result.rows_updated == 1
|
||||
assert table.to_arrow()["vector"].to_pylist() == [[1.0, 2.0, 3.0, 4.0]]
|
||||
|
||||
|
||||
def test_update_types(mem_db: DBConnection):
|
||||
table = mem_db.create_table(
|
||||
"my_table",
|
||||
@@ -2363,6 +2524,55 @@ def test_merge_insert(mem_db: DBConnection):
|
||||
)
|
||||
|
||||
|
||||
def test_merge_insert_nullable_pandas_into_pydantic_schema(mem_db: DBConnection):
|
||||
# Regression test for https://github.com/lancedb/lancedb/issues/2366
|
||||
pd = pytest.importorskip("pandas")
|
||||
|
||||
class Document(LanceModel):
|
||||
id: int
|
||||
title: str
|
||||
content: str
|
||||
|
||||
table = mem_db.create_table("documents", schema=Document)
|
||||
table.add(
|
||||
pd.DataFrame(
|
||||
{
|
||||
"title": ["Old title", "Unchanged"],
|
||||
"id": [2, 3],
|
||||
"content": ["Old content", "Keep this"],
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
# Pandas produces nullable Arrow fields, in an order that differs from the
|
||||
# non-nullable Pydantic schema. This is valid as long as the data has no nulls.
|
||||
new_data = pd.DataFrame(
|
||||
{
|
||||
"title": ["Inserted", "Updated"],
|
||||
"id": [1, 2],
|
||||
"content": ["New row", "New content"],
|
||||
}
|
||||
)
|
||||
result = (
|
||||
table.merge_insert("id")
|
||||
.when_matched_update_all()
|
||||
.when_not_matched_insert_all()
|
||||
.execute(new_data)
|
||||
)
|
||||
|
||||
assert result.num_inserted_rows == 1
|
||||
assert result.num_updated_rows == 1
|
||||
expected = pa.Table.from_pylist(
|
||||
[
|
||||
{"id": 1, "title": "Inserted", "content": "New row"},
|
||||
{"id": 2, "title": "Updated", "content": "New content"},
|
||||
{"id": 3, "title": "Unchanged", "content": "Keep this"},
|
||||
],
|
||||
schema=Document.to_arrow_schema(),
|
||||
)
|
||||
assert table.to_arrow().sort_by("id") == expected
|
||||
|
||||
|
||||
def test_merge_insert_by_source_delete_expr(mem_db: DBConnection):
|
||||
table = mem_db.create_table(
|
||||
"my_table",
|
||||
@@ -2463,6 +2673,36 @@ def test_merge_insert_subschema(mem_db: DBConnection, data_format):
|
||||
assert table.to_arrow().sort_by("id") == expected
|
||||
|
||||
|
||||
def test_repeated_partial_merge_insert_with_scalar_index(mem_db: DBConnection):
|
||||
def make_batch(start: int) -> pa.Table:
|
||||
return pa.table(
|
||||
{
|
||||
"id": [f"id-{i:04}" for i in range(start, start + 100)],
|
||||
"category": ["A"] * 100,
|
||||
"value_a": [float(i) for i in range(start, start + 100)],
|
||||
"value_b": [float(i) / 10 for i in range(100)],
|
||||
}
|
||||
)
|
||||
|
||||
table = mem_db.create_table("my_table", data=make_batch(0))
|
||||
table.add(make_batch(100))
|
||||
table.add(make_batch(200))
|
||||
table.create_index("id", config=BTree())
|
||||
|
||||
ids = [f"id-{i:04}" for i in range(100, 200)]
|
||||
for value in (999.0, 888.0):
|
||||
result = (
|
||||
table.merge_insert("id")
|
||||
.when_matched_update_all()
|
||||
.execute(pa.table({"id": ids, "value_a": [value] * 100}))
|
||||
)
|
||||
assert result.num_updated_rows == 100
|
||||
|
||||
actual = table.to_arrow().sort_by("id")
|
||||
assert actual.num_rows == 300
|
||||
assert actual["value_a"].to_pylist()[100:200] == [888.0] * 100
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_merge_insert_async(mem_db_async: AsyncConnection):
|
||||
data = pa.table({"a": [1, 2, 3], "b": ["a", "b", "c"]})
|
||||
@@ -2559,15 +2799,40 @@ def test_create_with_embedding_function(mem_db: DBConnection):
|
||||
assert actual == expected
|
||||
|
||||
|
||||
def test_create_f16_table_from_arrow_data(mem_db: DBConnection):
|
||||
dimension = 32
|
||||
num_rows = 512
|
||||
values = pa.array(
|
||||
np.random.default_rng(42)
|
||||
.standard_normal(num_rows * dimension)
|
||||
.astype(np.float16)
|
||||
)
|
||||
df = pa.table(
|
||||
{
|
||||
"text": [f"s-{i}" for i in range(num_rows)],
|
||||
"vector": pa.FixedSizeListArray.from_arrays(values, dimension),
|
||||
}
|
||||
)
|
||||
table = mem_db.create_table("f16_tbl", data=df)
|
||||
assert table.schema.field("vector").type == pa.list_(pa.float16(), dimension)
|
||||
table.create_index(num_partitions=2, num_sub_vectors=2)
|
||||
|
||||
query = df["vector"][2].as_py()
|
||||
expected = table.search(query).limit(2).to_arrow()
|
||||
|
||||
assert "s-2" in expected["text"].to_pylist()
|
||||
|
||||
|
||||
def test_create_f16_table(mem_db: DBConnection):
|
||||
class MyTable(LanceModel):
|
||||
text: str
|
||||
vector: Vector(32, value_type=pa.float16())
|
||||
|
||||
rng = np.random.default_rng(42)
|
||||
df = pa.table(
|
||||
{
|
||||
"text": [f"s-{i}" for i in range(512)],
|
||||
"vector": [np.random.randn(32).astype(np.float16) for _ in range(512)],
|
||||
"vector": [rng.standard_normal(32).astype(np.float16) for _ in range(512)],
|
||||
}
|
||||
)
|
||||
table = mem_db.create_table(
|
||||
@@ -3448,7 +3713,8 @@ def test_stats(mem_db: DBConnection):
|
||||
stats = table.stats()
|
||||
print(f"{stats=}")
|
||||
assert stats == {
|
||||
"total_bytes": 60,
|
||||
# Full on-disk size of the data file, footer and metadata included.
|
||||
"total_bytes": 633,
|
||||
"num_rows": 2,
|
||||
"num_indices": 0,
|
||||
"fragment_stats": {
|
||||
@@ -3466,6 +3732,13 @@ def test_stats(mem_db: DBConnection):
|
||||
},
|
||||
}
|
||||
|
||||
# Index files count toward total_bytes too (only deletion files and
|
||||
# manifests are excluded).
|
||||
table.create_index("id", config=BTree())
|
||||
stats_with_index = table.stats()
|
||||
assert stats_with_index["num_indices"] == 1
|
||||
assert stats_with_index["total_bytes"] > stats["total_bytes"]
|
||||
|
||||
|
||||
def test_create_table_empty_list_with_schema(mem_db: DBConnection):
|
||||
"""Test creating table with empty list data and schema
|
||||
@@ -3489,8 +3762,8 @@ def test_create_table_empty_list_no_schema_error(mem_db: DBConnection):
|
||||
mem_db.create_table("test_empty_no_schema", data=[])
|
||||
|
||||
|
||||
def test_add_table_with_empty_embeddings(tmp_path):
|
||||
"""Test exact scenario from issue #1968
|
||||
def test_create_table_without_data_with_vector_schema(tmp_path):
|
||||
"""Test exact scenario from issue #1968.
|
||||
|
||||
Regression test for issue #1968:
|
||||
https://github.com/lancedb/lancedb/issues/1968
|
||||
@@ -3502,6 +3775,9 @@ def test_add_table_with_empty_embeddings(tmp_path):
|
||||
embedding: Vector(16)
|
||||
|
||||
table = db.create_table("test", schema=MySchema)
|
||||
assert table.count_rows() == 0
|
||||
assert table.schema == MySchema.to_arrow_schema()
|
||||
|
||||
table.add(
|
||||
[{"text": "bar", "embedding": [0.1] * 16}],
|
||||
on_bad_vectors="drop",
|
||||
|
||||
@@ -75,6 +75,22 @@ class TestVoyageAIModelRegistration:
|
||||
with pytest.raises(ValueError, match="not supported"):
|
||||
func.ndims()
|
||||
|
||||
def test_voyage3_source_embeddings_use_text_api(self, mock_voyageai_client):
|
||||
"""Regression test for text table data being sent to the multimodal API."""
|
||||
mock_voyageai_client.tokenize.return_value = [["hello", "world"]]
|
||||
mock_voyageai_client.embed.return_value.embeddings = [[0.1] * 1024]
|
||||
|
||||
registry = get_registry()
|
||||
func = registry.get("voyageai").create(name="voyage-3")
|
||||
|
||||
embeddings = func.compute_source_embeddings("hello world")
|
||||
|
||||
assert embeddings == [[0.1] * 1024]
|
||||
mock_voyageai_client.embed.assert_called_once_with(
|
||||
texts=["hello world"], model="voyage-3", input_type="document"
|
||||
)
|
||||
mock_voyageai_client.multimodal_embed.assert_not_called()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
[
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
from typing import assert_type
|
||||
|
||||
import lancedb
|
||||
from lancedb import AsyncConnection, DBConnection
|
||||
|
||||
|
||||
def check_connect_type() -> None:
|
||||
assert_type(lancedb.connect("memory://"), DBConnection)
|
||||
|
||||
|
||||
async def check_connect_async_type() -> None:
|
||||
assert_type(await lancedb.connect_async("memory://"), AsyncConnection)
|
||||
@@ -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, ¤t)
|
||||
.await
|
||||
.infer_error()?;
|
||||
// `()` maps to an empty Python tuple via IntoPyObject; return Option
|
||||
// so the async bridge yields exact Python None.
|
||||
Ok(None::<()>)
|
||||
})
|
||||
}
|
||||
|
||||
/// Revoke an exact immutable Function by administrator set-bit.
|
||||
///
|
||||
/// Clones the observed native [`crate::function::Function`] once and
|
||||
/// delegates to Rust [`lancedb::Connection::revoke_function`].
|
||||
pub fn _revoke_function<'py>(
|
||||
self_: PyRef<'py, Self>,
|
||||
function: Bound<'_, crate::function::Function>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
let function = function.get().inner().clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
inner.revoke_function(&function).await.infer_error()?;
|
||||
// `()` maps to an empty Python tuple via IntoPyObject; return Option
|
||||
// so the async bridge yields exact Python None.
|
||||
Ok(None::<()>)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
|
||||
+15
-2
@@ -102,11 +102,14 @@ impl<T> PythonErrorExt<T> for std::result::Result<T, LanceError> {
|
||||
err.setattr(intern!(py, "__cause__"), cause_err)?;
|
||||
Err(PyErr::from_value(err))
|
||||
}),
|
||||
LanceError::JobFailed { .. } => Python::attach(|py| {
|
||||
LanceError::JobFailed { failure, .. } => Python::attach(|py| {
|
||||
let cls = py
|
||||
.import(intern!(py, "lancedb.exceptions"))?
|
||||
.getattr(intern!(py, "JobFailedError"))?;
|
||||
Err(PyErr::from_value(cls.call1((err.to_string(),))?))
|
||||
// Structural projection only: failure.error_code.as_str().
|
||||
// Never infer a code from message, phase, retryable, or source.
|
||||
let error_code = failure.error_code.as_ref().map(|code| code.as_str());
|
||||
Err(PyErr::from_value(cls.call1((err.to_string(), error_code))?))
|
||||
}),
|
||||
LanceError::JobCancelled { .. } => Python::attach(|py| {
|
||||
let cls = py
|
||||
@@ -114,6 +117,16 @@ impl<T> PythonErrorExt<T> for std::result::Result<T, LanceError> {
|
||||
.getattr(intern!(py, "JobCancelledError"))?;
|
||||
Err(PyErr::from_value(cls.call1((err.to_string(),))?))
|
||||
}),
|
||||
LanceError::Function { code, message } => Python::attach(|py| {
|
||||
let cls = py
|
||||
.import(intern!(py, "lancedb.exceptions"))?
|
||||
.getattr(intern!(py, "FunctionError"))?;
|
||||
// Structural projection only: code.as_str() + sanitized message.
|
||||
// Never infer a code from HTTP status or diagnostic text.
|
||||
Err(PyErr::from_value(
|
||||
cls.call1((message.as_str(), code.as_str()))?,
|
||||
))
|
||||
}),
|
||||
_ => self.runtime_error(),
|
||||
},
|
||||
}
|
||||
|
||||
+28
-1
@@ -10,7 +10,7 @@
|
||||
use std::ops::{Add, Div, Mul, Not, Sub};
|
||||
|
||||
use arrow::{datatypes::DataType, pyarrow::PyArrowType};
|
||||
use datafusion_common::ScalarValue;
|
||||
use datafusion_common::{Column, ScalarValue};
|
||||
use lancedb::expr::{
|
||||
DfExpr, col as ldb_col, contains, expr_cast, is_in, lit as df_lit, lower, upper,
|
||||
};
|
||||
@@ -27,6 +27,33 @@ use pyo3::{Bound, PyAny, PyResult, exceptions::PyValueError, prelude::*, pyfunct
|
||||
#[derive(Clone)]
|
||||
pub struct PyExpr(pub DfExpr);
|
||||
|
||||
/// Crate-private inspection result for Function call authoring (FF-028).
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) enum DirectExprView<'a> {
|
||||
/// Direct unqualified DataFusion Column; name is case-sensitive.
|
||||
UnqualifiedColumn(&'a str),
|
||||
/// Direct Literal scalar; Arrow type is owned by the scalar value.
|
||||
Literal(&'a ScalarValue),
|
||||
}
|
||||
|
||||
impl PyExpr {
|
||||
/// Inspect a direct Column/Literal node for Function call authoring.
|
||||
///
|
||||
/// Returns `None` for every other expression shape (arithmetic, cast,
|
||||
/// scalar function, predicate, alias, qualified column, etc.).
|
||||
pub(crate) fn as_direct_column_or_literal(&self) -> Option<DirectExprView<'_>> {
|
||||
match &self.0 {
|
||||
DfExpr::Column(Column {
|
||||
relation: None,
|
||||
name,
|
||||
..
|
||||
}) => Some(DirectExprView::UnqualifiedColumn(name.as_str())),
|
||||
DfExpr::Literal(value, _) => Some(DirectExprView::Literal(value)),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl PyExpr {
|
||||
// ── comparisons ──────────────────────────────────────────────────────────
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+1
-1
@@ -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
@@ -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),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -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>,
|
||||
|
||||
+280
-18
@@ -26,13 +26,74 @@ 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},
|
||||
types::{IntoPyDict, PyAnyMethods, PyBytes, PyDict, PyDictMethods, PyList, PyListMethods},
|
||||
};
|
||||
|
||||
mod scannable;
|
||||
|
||||
/// Convert `LsmStats` to a Python dict, preserving the per-bucket list.
|
||||
///
|
||||
/// Deliberately not flattened to a table-level summary: a table is N
|
||||
/// buckets on one node, and the per-bucket detail is the reason the
|
||||
/// endpoint exists — flattening hides the single hot bucket someone opened
|
||||
/// it to find.
|
||||
fn lsm_stats_to_py(py: Python<'_>, stats: &lancedb::table::LsmStats) -> PyResult<Py<PyDict>> {
|
||||
let out = PyDict::new(py);
|
||||
let buckets = PyList::empty(py);
|
||||
for b in &stats.buckets {
|
||||
let e = PyDict::new(py);
|
||||
e.set_item("shard_id", &b.shard_id)?;
|
||||
e.set_item("status", &b.status)?;
|
||||
e.set_item("writer_epoch", b.writer_epoch)?;
|
||||
e.set_item("manifest_version", b.manifest_version)?;
|
||||
e.set_item("current_generation", b.current_generation)?;
|
||||
e.set_item(
|
||||
"replay_after_wal_entry_position",
|
||||
b.replay_after_wal_entry_position,
|
||||
)?;
|
||||
e.set_item(
|
||||
"wal_entry_position_last_seen",
|
||||
b.wal_entry_position_last_seen,
|
||||
)?;
|
||||
|
||||
let generations = PyList::empty(py);
|
||||
for g in &b.generations {
|
||||
let ge = PyDict::new(py);
|
||||
ge.set_item("generation", g.generation)?;
|
||||
ge.set_item("bytes", g.bytes)?;
|
||||
ge.set_item("rows", g.rows)?;
|
||||
generations.append(ge)?;
|
||||
}
|
||||
e.set_item("generations", generations)?;
|
||||
e.set_item("compacting", b.compacting)?;
|
||||
|
||||
e.set_item(
|
||||
"memtables",
|
||||
b.memtables
|
||||
.as_ref()
|
||||
.map(|ms| {
|
||||
let l = PyList::empty(py);
|
||||
for m in ms {
|
||||
let d = PyDict::new(py);
|
||||
d.set_item("generation", m.generation)?;
|
||||
d.set_item("rows", m.rows)?;
|
||||
d.set_item("bytes", m.bytes)?;
|
||||
d.set_item("batches", m.batches)?;
|
||||
d.set_item("indexes", m.indexes.clone())?;
|
||||
l.append(d)?;
|
||||
}
|
||||
PyResult::Ok(l.unbind())
|
||||
})
|
||||
.transpose()?,
|
||||
)?;
|
||||
buckets.append(e)?;
|
||||
}
|
||||
out.set_item("buckets", buckets)?;
|
||||
Ok(out.unbind())
|
||||
}
|
||||
|
||||
#[derive(FromPyObject)]
|
||||
enum PredicateArg {
|
||||
Expr(PyExpr),
|
||||
@@ -185,12 +246,22 @@ impl From<lancedb::table::MergeResult> for MergeResult {
|
||||
}
|
||||
}
|
||||
|
||||
/// Render for `__repr__`, so the default reads as Python's `None` rather than
|
||||
/// Rust's `Some([..])`.
|
||||
fn fmt_maintained(maintained: &Option<Vec<String>>) -> String {
|
||||
match maintained {
|
||||
Some(names) => format!("{:?}", names),
|
||||
None => "None".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Specification selecting Lance's MemWAL LSM-style write path for
|
||||
/// `merge_insert`.
|
||||
///
|
||||
/// Constructed via the `bucket(...)`, `identity(...)`, or `unsharded()`
|
||||
/// classmethods, then optionally chain `with_maintained_indexes(...)` and
|
||||
/// `with_writer_config_defaults(...)`.
|
||||
/// `with_writer_config_defaults(...)`. A fresh spec maintains every index the
|
||||
/// MemWAL supports, resolved on install.
|
||||
#[pyclass(from_py_object)]
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct LsmWriteSpec {
|
||||
@@ -230,11 +301,11 @@ impl LsmWriteSpec {
|
||||
}
|
||||
}
|
||||
|
||||
/// Replace the list of indexes the MemWAL should keep up to date as
|
||||
/// rows are appended. Each name must reference an index that
|
||||
/// already exists on the table at the time `set_lsm_write_spec`
|
||||
/// is called.
|
||||
pub fn with_maintained_indexes(&self, indexes: Vec<String>) -> Self {
|
||||
/// Set which indexes the MemWAL maintains. `None` (the default)
|
||||
/// resolves every supported index on install; a list is verbatim,
|
||||
/// and an empty list maintains nothing.
|
||||
#[pyo3(signature = (indexes))]
|
||||
pub fn with_maintained_indexes(&self, indexes: Option<Vec<String>>) -> Self {
|
||||
Self {
|
||||
inner: self.inner.clone().with_maintained_indexes(indexes),
|
||||
}
|
||||
@@ -256,23 +327,29 @@ impl LsmWriteSpec {
|
||||
maintained_indexes,
|
||||
writer_config_defaults,
|
||||
} => format!(
|
||||
"LsmWriteSpec.bucket(column={:?}, num_buckets={}, maintained_indexes={:?}, writer_config_defaults={:?})",
|
||||
column, num_buckets, maintained_indexes, writer_config_defaults,
|
||||
"LsmWriteSpec.bucket(column={:?}, num_buckets={}, maintained_indexes={}, writer_config_defaults={:?})",
|
||||
column,
|
||||
num_buckets,
|
||||
fmt_maintained(maintained_indexes),
|
||||
writer_config_defaults,
|
||||
),
|
||||
lancedb::table::LsmWriteSpec::Identity {
|
||||
column,
|
||||
maintained_indexes,
|
||||
writer_config_defaults,
|
||||
} => format!(
|
||||
"LsmWriteSpec.identity(column={:?}, maintained_indexes={:?}, writer_config_defaults={:?})",
|
||||
column, maintained_indexes, writer_config_defaults,
|
||||
"LsmWriteSpec.identity(column={:?}, maintained_indexes={}, writer_config_defaults={:?})",
|
||||
column,
|
||||
fmt_maintained(maintained_indexes),
|
||||
writer_config_defaults,
|
||||
),
|
||||
lancedb::table::LsmWriteSpec::Unsharded {
|
||||
maintained_indexes,
|
||||
writer_config_defaults,
|
||||
} => format!(
|
||||
"LsmWriteSpec.unsharded(maintained_indexes={:?}, writer_config_defaults={:?})",
|
||||
maintained_indexes, writer_config_defaults,
|
||||
"LsmWriteSpec.unsharded(maintained_indexes={}, writer_config_defaults={:?})",
|
||||
fmt_maintained(maintained_indexes),
|
||||
writer_config_defaults,
|
||||
),
|
||||
}
|
||||
}
|
||||
@@ -307,10 +384,10 @@ impl LsmWriteSpec {
|
||||
}
|
||||
}
|
||||
|
||||
/// Names of indexes the MemWAL should keep up to date during writes.
|
||||
/// Indexes the MemWAL keeps up to date, or `None` for every supported one.
|
||||
#[getter]
|
||||
pub fn maintained_indexes(&self) -> Vec<String> {
|
||||
self.inner.maintained_indexes().to_vec()
|
||||
pub fn maintained_indexes(&self) -> Option<Vec<String>> {
|
||||
self.inner.maintained_indexes().map(<[String]>::to_vec)
|
||||
}
|
||||
|
||||
/// Default `ShardWriter` configuration recorded by this spec.
|
||||
@@ -502,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,
|
||||
@@ -853,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 {
|
||||
@@ -1339,6 +1556,51 @@ impl Table {
|
||||
})
|
||||
}
|
||||
|
||||
/// Converge the table's LSM write path into its base table.
|
||||
///
|
||||
/// Best-effort: with writes flowing, new rows may land after the last
|
||||
/// pass. Errors if the table stops making progress.
|
||||
pub fn checkpoint_lsm(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
inner.checkpoint_lsm().await.infer_error()
|
||||
})
|
||||
}
|
||||
|
||||
/// Seal every bucket's active memtable into L0.
|
||||
pub fn flush_lsm(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(
|
||||
self_.py(),
|
||||
async move { inner.flush_lsm().await.infer_error() },
|
||||
)
|
||||
}
|
||||
|
||||
/// Trigger a background L0 → base pass per bucket. Returns once the
|
||||
/// passes are dispatched, not once they finish — watch `get_lsm_stats`.
|
||||
pub fn compact_lsm(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
inner.compact_lsm().await.infer_error()
|
||||
})
|
||||
}
|
||||
|
||||
/// Live LSM state, or `None` when the LSM write path is not enabled.
|
||||
#[pyo3(signature = (include_generation_rows=false))]
|
||||
pub fn get_lsm_stats(
|
||||
self_: PyRef<'_, Self>,
|
||||
include_generation_rows: bool,
|
||||
) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let stats = inner
|
||||
.get_lsm_stats(include_generation_rows)
|
||||
.await
|
||||
.infer_error()?;
|
||||
Python::attach(|py| stats.map(|s| lsm_stats_to_py(py, &s)).transpose())
|
||||
})
|
||||
}
|
||||
|
||||
pub fn close_lsm_writers(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
|
||||
Generated
+1
-1
@@ -1998,7 +1998,7 @@ requires-dist = [
|
||||
{ name = "pillow", marker = "extra == 'clip'", specifier = ">=12.1.1" },
|
||||
{ name = "pillow", marker = "extra == 'embeddings'", specifier = ">=12.1.1" },
|
||||
{ name = "pillow", marker = "extra == 'siglip'", specifier = ">=12.1.1" },
|
||||
{ name = "polars", marker = "extra == 'tests'", specifier = ">=0.19,<=1.3.0" },
|
||||
{ name = "polars", marker = "extra == 'tests'", specifier = ">=0.19,<=1.32.3" },
|
||||
{ name = "pre-commit", marker = "extra == 'dev'", specifier = ">=3.5.0" },
|
||||
{ name = "pyarrow", specifier = ">=16" },
|
||||
{ name = "pyarrow", marker = "extra == 'tests'", specifier = "<25" },
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "lancedb"
|
||||
version = "0.37.1-beta.0"
|
||||
version = "0.37.1-beta.1"
|
||||
edition.workspace = true
|
||||
description = "LanceDB: A serverless, low-latency vector database for AI applications"
|
||||
license.workspace = true
|
||||
@@ -12,6 +12,7 @@ rust-version.workspace = true
|
||||
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
|
||||
[dependencies]
|
||||
ahash = { workspace = true }
|
||||
base64 = "0.22"
|
||||
arrow = { workspace = true }
|
||||
arrow-array = { workspace = true }
|
||||
arrow-buffer = { workspace = true }
|
||||
@@ -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 transitive GooseFS SDK until the 0.1.6 compile break is fixed upstream.
|
||||
goosefs-sdk = { version = "=0.1.5", optional = true }
|
||||
moka = { workspace = true }
|
||||
pin-project = { workspace = true }
|
||||
tokio = { version = "1.23", features = ["rt-multi-thread", "sync"] }
|
||||
@@ -75,6 +74,8 @@ reqwest = { version = "0.12.0", default-features = false, features = [
|
||||
"http2",
|
||||
"json",
|
||||
"macos-system-configuration",
|
||||
# Avoid linking OpenSSL into Python wheels, which breaks on FIPS hosts.
|
||||
"rustls-tls-native-roots",
|
||||
"stream",
|
||||
], optional = true }
|
||||
http = { version = "1", optional = true } # Matching what is in reqwest
|
||||
@@ -98,7 +99,8 @@ anyhow = "1"
|
||||
lance-testing = { workspace = true }
|
||||
tempfile = "3.5.0"
|
||||
random_word = { version = "0.4.3", features = ["en"] }
|
||||
tokio = { version = "1.23", features = ["io-util", "macros", "net", "rt-multi-thread", "sync"] }
|
||||
roaring = "0.11.4"
|
||||
tokio = { version = "1.23", features = ["io-util", "macros", "net", "rt-multi-thread", "sync", "test-util"] }
|
||||
uuid = { version = "1.7.0", features = ["v4"] }
|
||||
walkdir = "2"
|
||||
aws-sdk-dynamodb = { version = "1.55.0" }
|
||||
@@ -133,7 +135,6 @@ azure = [
|
||||
]
|
||||
cos = ["lance/tencent", "lance-io/tencent"]
|
||||
goosefs = [
|
||||
"dep:goosefs-sdk",
|
||||
"lance/goosefs",
|
||||
"lance-io/goosefs",
|
||||
"lance-namespace-impls/dir-goosefs",
|
||||
|
||||
@@ -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_encoding::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]
|
||||
|
||||
@@ -28,13 +28,14 @@ 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,
|
||||
db::{OPT_REMOTE_API_KEY, OPT_REMOTE_HOST_OVERRIDE, OPT_REMOTE_REGION},
|
||||
};
|
||||
use lance::io::ObjectStoreParams;
|
||||
pub use lance_encoding::version::LanceFileVersion;
|
||||
pub use lance_file::version::LanceFileVersion;
|
||||
#[cfg(feature = "remote")]
|
||||
use lance_io::object_store::StorageOptions;
|
||||
use lance_io::object_store::{StorageOptionsAccessor, StorageOptionsProvider};
|
||||
@@ -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
|
||||
|
||||
@@ -202,6 +202,17 @@ mod tests {
|
||||
assert_eq!(table.count_rows(None).await.unwrap(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn create_table_in_named_memory_database() {
|
||||
let db = connect("memory://foo").execute().await.unwrap();
|
||||
let batch = record_batch!(("id", Int64, [1, 2, 3])).unwrap();
|
||||
|
||||
let table = db.create_table("my_table", batch).execute().await.unwrap();
|
||||
|
||||
assert_eq!(table.uri().await.unwrap(), "memory://foo/my_table.lance");
|
||||
assert_eq!(table.count_rows(None).await.unwrap(), 3);
|
||||
}
|
||||
|
||||
async fn test_create_table_with_data<T>(data: T)
|
||||
where
|
||||
T: Scannable + 'static,
|
||||
@@ -427,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]
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
@@ -12,7 +12,7 @@ use lance::dataset::refs::Ref;
|
||||
use lance::dataset::{ReadParams, WriteMode, builder::DatasetBuilder};
|
||||
use lance::io::{ObjectStore, ObjectStoreParams, WrappingObjectStore};
|
||||
use lance_datafusion::utils::StreamingWriteSource;
|
||||
use lance_encoding::version::LanceFileVersion;
|
||||
use lance_file::version::LanceFileVersion;
|
||||
use lance_io::object_store::{StorageOptionsAccessor, StorageOptionsProvider};
|
||||
use lance_table::io::commit::commit_handler_from_url;
|
||||
use object_store::local::LocalFileSystem;
|
||||
@@ -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;
|
||||
}
|
||||
@@ -1294,9 +1299,11 @@ mod tests {
|
||||
use crate::connection::ConnectRequest;
|
||||
use crate::data::scannable::Scannable;
|
||||
use crate::database::{CreateTableMode, CreateTableRequest};
|
||||
use crate::table::WriteOptions;
|
||||
use crate::query::QueryRequest;
|
||||
use crate::table::{AnyQuery, WriteOptions};
|
||||
use arrow_array::{Int32Array, RecordBatch, StringArray};
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
use futures::TryStreamExt;
|
||||
use std::path::PathBuf;
|
||||
use tempfile::tempdir;
|
||||
|
||||
@@ -1376,6 +1383,156 @@ mod tests {
|
||||
assert!(!tempdir.path().join("__manifest").exists());
|
||||
}
|
||||
|
||||
/// Regression test for https://github.com/lancedb/lancedb/issues/1600.
|
||||
///
|
||||
/// Opening a table used to create a separate object-store client instead of
|
||||
/// reusing the one that successfully connected to the database. Repeating
|
||||
/// credential discovery made S3 table opens intermittent, especially in AWS
|
||||
/// Lambda, and the failed open was reported as `TableNotFound`.
|
||||
#[tokio::test]
|
||||
async fn test_open_table_reuses_connection_object_store() {
|
||||
let tempdir = tempdir().unwrap();
|
||||
let uri = tempdir.path().to_str().unwrap();
|
||||
let registry = Arc::new(lance_io::object_store::ObjectStoreRegistry::default());
|
||||
let session = Arc::new(lance::session::Session::new(16, 16, registry.clone()));
|
||||
|
||||
let request = 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: Some(session),
|
||||
};
|
||||
let db = ListingDatabase::connect_with_options(&request)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]));
|
||||
db.create_table(CreateTableRequest {
|
||||
name: "test".to_string(),
|
||||
namespace_path: vec![],
|
||||
data: Box::new(RecordBatch::new_empty(schema)) as Box<dyn Scannable>,
|
||||
mode: CreateTableMode::Create,
|
||||
write_options: Default::default(),
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let before_open = registry.stats();
|
||||
for _ in 0..3 {
|
||||
let table = db
|
||||
.open_table(OpenTableRequest {
|
||||
name: "test".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!(table.count_rows(None).await.unwrap(), 0);
|
||||
}
|
||||
|
||||
let after_open = registry.stats();
|
||||
assert_eq!(after_open.misses, before_open.misses);
|
||||
assert!(after_open.hits >= before_open.hits + 3);
|
||||
}
|
||||
|
||||
/// Regression test for https://github.com/lancedb/lancedb/issues/3197.
|
||||
#[cfg(unix)]
|
||||
#[tokio::test]
|
||||
async fn test_open_table_follows_hugging_face_symlinks() {
|
||||
let (tempdir, db) = setup_database().await;
|
||||
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]));
|
||||
db.create_table(CreateTableRequest {
|
||||
name: "test".to_string(),
|
||||
namespace_path: vec![],
|
||||
data: Box::new(
|
||||
RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(vec![1, 2, 3]))])
|
||||
.unwrap(),
|
||||
) as Box<dyn Scannable>,
|
||||
mode: CreateTableMode::Create,
|
||||
write_options: Default::default(),
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let table_dir = tempdir.path().join("test.lance");
|
||||
let versions_dir = table_dir.join("_versions");
|
||||
let manifest_path = std::fs::read_dir(&versions_dir)
|
||||
.unwrap()
|
||||
.map(|entry| entry.unwrap().path())
|
||||
.find(|path| path.extension().is_some_and(|ext| ext == "manifest"))
|
||||
.unwrap();
|
||||
let data_path = std::fs::read_dir(table_dir.join("data"))
|
||||
.unwrap()
|
||||
.map(|entry| entry.unwrap().path())
|
||||
.find(|path| path.extension().is_some_and(|ext| ext == "lance"))
|
||||
.unwrap();
|
||||
|
||||
// Hugging Face snapshots keep dataset objects in a separate blob directory and
|
||||
// expose them through relative symlinks.
|
||||
let blobs_dir = tempdir.path().join("blobs");
|
||||
std::fs::create_dir(&blobs_dir).unwrap();
|
||||
let manifest_blob = "9b603c63d0e692e05d58be25605f2f2064cc781e5ff94fe983a405059547b816";
|
||||
let data_blob = "be64f20e5723bd0a27cfdbdb41cf7d6fad94cd572a71973b717fb8340f4310c5";
|
||||
std::fs::rename(&manifest_path, blobs_dir.join(manifest_blob)).unwrap();
|
||||
std::fs::rename(&data_path, blobs_dir.join(data_blob)).unwrap();
|
||||
std::os::unix::fs::symlink(Path::new("../../blobs").join(manifest_blob), &manifest_path)
|
||||
.unwrap();
|
||||
std::os::unix::fs::symlink(Path::new("../../blobs").join(data_blob), &data_path).unwrap();
|
||||
let symlink_len = std::fs::symlink_metadata(&manifest_path).unwrap().len();
|
||||
let target_len = std::fs::metadata(&manifest_path).unwrap().len();
|
||||
assert_ne!(symlink_len, target_len);
|
||||
|
||||
drop(db);
|
||||
let db = ListingDatabase::connect_with_options(&ConnectRequest {
|
||||
uri: tempdir.path().to_str().unwrap().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,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let table = db
|
||||
.open_table(OpenTableRequest {
|
||||
name: "test".to_string(),
|
||||
namespace_path: vec![],
|
||||
index_cache_size: None,
|
||||
lance_read_params: None,
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
managed_versioning: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let batches = table
|
||||
.query(
|
||||
&AnyQuery::Query(QueryRequest::default()),
|
||||
Default::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
.try_collect::<Vec<_>>()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::<usize>(), 3);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_clone_table_basic() {
|
||||
let (_tempdir, db) = setup_database().await;
|
||||
@@ -2280,7 +2437,7 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_table_uri() {
|
||||
let (_tempdir, db) = setup_database().await;
|
||||
let (_tempdir, mut db) = setup_database().await;
|
||||
|
||||
let mut pb = PathBuf::new();
|
||||
pb.push(db.uri.clone());
|
||||
@@ -2289,6 +2446,18 @@ mod tests {
|
||||
let expected = pb.to_str().unwrap();
|
||||
let uri = db.table_uri("test").ok().unwrap();
|
||||
assert_eq!(uri, expected);
|
||||
|
||||
// URI paths always use forward slashes, even on Windows. Using
|
||||
// `Path::join` here used to produce `az://container/prefix\\test.lance`,
|
||||
// which Azure treated as a different object from the table returned by
|
||||
// `table_names` (https://github.com/lancedb/lancedb/issues/1072).
|
||||
for base_uri in ["az://container/prefix", "az://container/prefix/"] {
|
||||
db.uri = base_uri.to_string();
|
||||
assert_eq!(
|
||||
db.table_uri("test").unwrap(),
|
||||
"az://container/prefix/test.lance"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// Regression: connecting via a URL-style URI (which goes through
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -201,7 +202,7 @@ impl LanceNamespaceDatabase {
|
||||
&self,
|
||||
request: &DbCreateTableRequest,
|
||||
) -> Result<(
|
||||
Option<lance_encoding::version::LanceFileVersion>,
|
||||
Option<lance_file::version::LanceFileVersion>,
|
||||
Option<bool>,
|
||||
Option<bool>,
|
||||
)> {
|
||||
@@ -214,7 +215,7 @@ impl LanceNamespaceDatabase {
|
||||
|
||||
let storage_version_override = storage_options
|
||||
.and_then(|opts| opts.get(OPT_NEW_TABLE_STORAGE_VERSION))
|
||||
.map(|s| s.parse::<lance_encoding::version::LanceFileVersion>())
|
||||
.map(|s| s.parse::<lance_file::version::LanceFileVersion>())
|
||||
.transpose()?;
|
||||
|
||||
let v2_manifest_override = storage_options
|
||||
@@ -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;
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user