mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-27 00:18:31 +00:00
Compare commits
133 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 | |||
| 7357d63e87 | |||
| 624a75edf7 | |||
| c7ea91f3ea | |||
| 8e24dd3828 | |||
| f79dc017c4 | |||
| e6ae93f52a | |||
| 3dd9c598e9 | |||
| 9e26bf3fba | |||
| 93354baf34 | |||
| 05602ec7d5 | |||
| e3b472c212 | |||
| a6418b6cb9 | |||
| dd2b11eda2 | |||
| 5a1015ba72 | |||
| 48945d0658 | |||
| 77208fd464 | |||
| b505dc1315 | |||
| 7dfdfe6401 | |||
| 4dc2d9a0f2 | |||
| 1ad6ce3a4e |
+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
|
||||
|
||||
@@ -92,6 +92,8 @@ Python bindings changes:
|
||||
* Should use `LOOP.run()` to call the corresponding `AsyncTable` method.
|
||||
6. Add concrete sync method to `RemoteTable` class in `python/python/lancedb/remote/table.py`.
|
||||
7. Add unit test in `python/tests/test_table.py`.
|
||||
8. If you added a new public class or module-level function (not just a method on an
|
||||
existing class), expose it in the API reference. See "Python API reference" below.
|
||||
|
||||
TypeScript bindings changes:
|
||||
|
||||
@@ -103,6 +105,33 @@ TypeScript bindings changes:
|
||||
5. Add test in `nodejs/__test__/table.test.ts`.
|
||||
6. Run `npm run docs` to generate TypeScript documentation.
|
||||
|
||||
## Python API reference
|
||||
|
||||
`docs/src/python/python.md` is the entire Python API reference. It is maintained by
|
||||
hand, and anything not listed there is not rendered at all, so new public classes and
|
||||
module-level functions have to be added explicitly. How depends on the module:
|
||||
|
||||
* `lancedb.index`, `lancedb.embeddings`, `lancedb.remote`, and `lancedb.rerankers` are
|
||||
rendered by a single directive each, driven by the module's `__all__`. Add the new
|
||||
name to `__all__` and it appears; forget, and it is silently omitted.
|
||||
* Everything else (`lancedb`, `lancedb.table`, `lancedb.query`, `lancedb.db`, ...) is
|
||||
listed symbol by symbol. Add a `::: lancedb.<module>.<Name>` line to the matching
|
||||
section, and remember that the page separates synchronous and asynchronous APIs.
|
||||
|
||||
Deliberately undocumented: concrete implementations reached through an abstract base
|
||||
(`LanceTable`, `LanceDBConnection`, `RemoteDBConnection`), query base classes already
|
||||
covered by `inherited_members`, and internal helpers.
|
||||
|
||||
Cross-references in docstrings use mkdocstrings syntax, `[text][lancedb.table.Table]`.
|
||||
Plain relative links such as `[Table](Table)` do not resolve. To check your work:
|
||||
|
||||
```shell
|
||||
pip install -r docs/requirements.txt
|
||||
cd docs && PYTHONPATH=. mkdocs build
|
||||
```
|
||||
|
||||
The docs site only builds on pushes to `main`, so this is not covered by PR CI.
|
||||
|
||||
## Review Guidelines
|
||||
|
||||
Please consider the following when reviewing code contributions.
|
||||
|
||||
Generated
+328
-307
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.0.0-beta.5", default-features = false, "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-core = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datagen = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-file = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-io = { "version" = "=10.0.0-beta.5", default-features = false, "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-index = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-linalg = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace-impls = { "version" = "=10.0.0-beta.5", default-features = false, "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-table = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-testing = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datafusion = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-encoding = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-arrow = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "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"
|
||||
|
||||
@@ -51,6 +51,11 @@ plugins:
|
||||
paths: [../python/python]
|
||||
options:
|
||||
docstring_style: numpy
|
||||
docstring_options:
|
||||
# Attributes documented in a `Parameters` section, and pydantic
|
||||
# dataclasses whose `__init__` griffe cannot see statically, both
|
||||
# trip this check. It reports nothing actionable here.
|
||||
warn_unknown_params: false
|
||||
heading_level: 3
|
||||
show_signature_annotations: true
|
||||
show_root_heading: true
|
||||
|
||||
@@ -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>
|
||||
```
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# Contributing to LanceDB Typescript
|
||||
|
||||
This document outlines the process for contributing to LanceDB Typescript.
|
||||
For general contribution guidelines, see [CONTRIBUTING.md](../CONTRIBUTING.md).
|
||||
For general contribution guidelines, see [CONTRIBUTING.md](https://github.com/lancedb/lancedb/blob/main/CONTRIBUTING.md).
|
||||
|
||||
## Project layout
|
||||
|
||||
|
||||
@@ -25,6 +25,27 @@ the underlying connection has been closed.
|
||||
|
||||
## Methods
|
||||
|
||||
### cancelJob()
|
||||
|
||||
```ts
|
||||
abstract cancelJob(jobId): Promise<boolean>
|
||||
```
|
||||
|
||||
Request cancellation of a server-side job by id.
|
||||
|
||||
Resolves to true if the server accepted the cancellation, false if no
|
||||
such job exists. Cancelling an already-terminal job is a no-op success.
|
||||
|
||||
#### Parameters
|
||||
|
||||
* **jobId**: `string`
|
||||
|
||||
#### Returns
|
||||
|
||||
`Promise`<`boolean`>
|
||||
|
||||
***
|
||||
|
||||
### cloneTable()
|
||||
|
||||
```ts
|
||||
@@ -365,6 +386,26 @@ Drop an existing table.
|
||||
|
||||
***
|
||||
|
||||
### getJob()
|
||||
|
||||
```ts
|
||||
abstract getJob(jobId): Promise<null | JobDescription>
|
||||
```
|
||||
|
||||
Describe a single server-side job by id.
|
||||
|
||||
Resolves to `null` when the server has no such job.
|
||||
|
||||
#### Parameters
|
||||
|
||||
* **jobId**: `string`
|
||||
|
||||
#### Returns
|
||||
|
||||
`Promise`<`null` \| [`JobDescription`](../interfaces/JobDescription.md)>
|
||||
|
||||
***
|
||||
|
||||
### isOpen()
|
||||
|
||||
```ts
|
||||
@@ -379,6 +420,62 @@ Return true if the connection has not been closed
|
||||
|
||||
***
|
||||
|
||||
### job()
|
||||
|
||||
```ts
|
||||
abstract job(jobId): Job
|
||||
```
|
||||
|
||||
A [Job](Job.md) handle for a server-side job by id.
|
||||
|
||||
The handle is constructed without a server round trip; an unknown id
|
||||
surfaces when the handle is used. Dropping the handle has no effect on
|
||||
the job itself.
|
||||
|
||||
#### Parameters
|
||||
|
||||
* **jobId**: `string`
|
||||
|
||||
#### Returns
|
||||
|
||||
[`Job`](Job.md)
|
||||
|
||||
***
|
||||
|
||||
### jobHistory()
|
||||
|
||||
```ts
|
||||
abstract jobHistory(jobId?): Promise<Table<any>>
|
||||
```
|
||||
|
||||
The lifecycle event history of a server-side job, as an Arrow table.
|
||||
|
||||
Lists history across all jobs when `jobId` is omitted.
|
||||
|
||||
#### Parameters
|
||||
|
||||
* **jobId?**: `string`
|
||||
|
||||
#### Returns
|
||||
|
||||
`Promise`<`Table`<`any`>>
|
||||
|
||||
***
|
||||
|
||||
### listJobs()
|
||||
|
||||
```ts
|
||||
abstract listJobs(): Promise<JobInfo[]>
|
||||
```
|
||||
|
||||
List server-side jobs across the database's tables.
|
||||
|
||||
#### Returns
|
||||
|
||||
`Promise`<[`JobInfo`](../interfaces/JobInfo.md)[]>
|
||||
|
||||
***
|
||||
|
||||
### listNamespaces()
|
||||
|
||||
```ts
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
[**@lancedb/lancedb**](../README.md) • **Docs**
|
||||
|
||||
***
|
||||
|
||||
[@lancedb/lancedb](../globals.md) / Job
|
||||
|
||||
# Class: Job
|
||||
|
||||
A handle to an operation that may still be running.
|
||||
|
||||
## Constructors
|
||||
|
||||
### new Job()
|
||||
|
||||
```ts
|
||||
new Job(): Job
|
||||
```
|
||||
|
||||
#### Returns
|
||||
|
||||
[`Job`](Job.md)
|
||||
|
||||
## Accessors
|
||||
|
||||
### id
|
||||
|
||||
```ts
|
||||
get id(): null | string
|
||||
```
|
||||
|
||||
Identifies the operation on the server that is running it. Operations
|
||||
that run in this process have no server id. The value is opaque.
|
||||
|
||||
#### Returns
|
||||
|
||||
`null` \| `string`
|
||||
|
||||
## Methods
|
||||
|
||||
### cancel()
|
||||
|
||||
```ts
|
||||
cancel(): Promise<void>
|
||||
```
|
||||
|
||||
Request cancellation. Cancelling a finished operation is a no-op.
|
||||
|
||||
#### Returns
|
||||
|
||||
`Promise`<`void`>
|
||||
|
||||
***
|
||||
|
||||
### status()
|
||||
|
||||
```ts
|
||||
status(): Promise<string>
|
||||
```
|
||||
|
||||
The operation's current lifecycle state: "running", "finished",
|
||||
"failed", or "cancelled".
|
||||
|
||||
A point snapshot; unlike [Job.wait](Job.md#wait) it does not block or reject
|
||||
on a terminal failure state. States a newer server reports that this
|
||||
client version does not know pass through as-is.
|
||||
|
||||
#### Returns
|
||||
|
||||
`Promise`<`string`>
|
||||
|
||||
***
|
||||
|
||||
### wait()
|
||||
|
||||
```ts
|
||||
wait(): Promise<void>
|
||||
```
|
||||
|
||||
Wait until the operation reaches a terminal state.
|
||||
|
||||
#### Returns
|
||||
|
||||
`Promise`<`void`>
|
||||
@@ -295,6 +295,29 @@ await table.createIndex("my_float_col");
|
||||
|
||||
***
|
||||
|
||||
### createIndexAsync()
|
||||
|
||||
```ts
|
||||
abstract createIndexAsync(column, options?): Promise<Job>
|
||||
```
|
||||
|
||||
Create an index, returning a handle to the indexing job.
|
||||
|
||||
The job may already be complete when returned; callers must not assume
|
||||
the index exists until [Job.wait](Job.md#wait) resolves.
|
||||
|
||||
#### Parameters
|
||||
|
||||
* **column**: `string`
|
||||
|
||||
* **options?**: `Partial`<[`IndexOptions`](../interfaces/IndexOptions.md)>
|
||||
|
||||
#### Returns
|
||||
|
||||
`Promise`<[`Job`](Job.md)>
|
||||
|
||||
***
|
||||
|
||||
### currentBranch()
|
||||
|
||||
```ts
|
||||
@@ -408,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
|
||||
|
||||
@@ -783,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)
|
||||
|
||||
@@ -25,6 +25,7 @@
|
||||
- [Connection](classes/Connection.md)
|
||||
- [HeaderProvider](classes/HeaderProvider.md)
|
||||
- [Index](classes/Index.md)
|
||||
- [Job](classes/Job.md)
|
||||
- [MakeArrowTableOptions](classes/MakeArrowTableOptions.md)
|
||||
- [MatchQuery](classes/MatchQuery.md)
|
||||
- [MergeInsertBuilder](classes/MergeInsertBuilder.md)
|
||||
@@ -88,6 +89,9 @@
|
||||
- [IvfFlatOptions](interfaces/IvfFlatOptions.md)
|
||||
- [IvfPqOptions](interfaces/IvfPqOptions.md)
|
||||
- [IvfRqOptions](interfaces/IvfRqOptions.md)
|
||||
- [JobDescription](interfaces/JobDescription.md)
|
||||
- [JobFailureInfo](interfaces/JobFailureInfo.md)
|
||||
- [JobInfo](interfaces/JobInfo.md)
|
||||
- [ListNamespacesOptions](interfaces/ListNamespacesOptions.md)
|
||||
- [ListNamespacesResponse](interfaces/ListNamespacesResponse.md)
|
||||
- [LsmWriteSpec](interfaces/LsmWriteSpec.md)
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
[**@lancedb/lancedb**](../README.md) • **Docs**
|
||||
|
||||
***
|
||||
|
||||
[@lancedb/lancedb](../globals.md) / JobDescription
|
||||
|
||||
# Interface: JobDescription
|
||||
|
||||
A described job from `Connection.getJob`.
|
||||
|
||||
## Properties
|
||||
|
||||
### creationMs
|
||||
|
||||
```ts
|
||||
creationMs: number;
|
||||
```
|
||||
|
||||
When the job was created, in milliseconds since the epoch.
|
||||
|
||||
***
|
||||
|
||||
### failure?
|
||||
|
||||
```ts
|
||||
optional failure: JobFailureInfo;
|
||||
```
|
||||
|
||||
Why the job failed, when the job is failed and the server reports a
|
||||
reason.
|
||||
|
||||
***
|
||||
|
||||
### jobId
|
||||
|
||||
```ts
|
||||
jobId: string;
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### jobType
|
||||
|
||||
```ts
|
||||
jobType: string;
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### specJson?
|
||||
|
||||
```ts
|
||||
optional specJson: string;
|
||||
```
|
||||
|
||||
The job-type-specific specification as a JSON string, when present.
|
||||
|
||||
***
|
||||
|
||||
### state
|
||||
|
||||
```ts
|
||||
state: string;
|
||||
```
|
||||
|
||||
Lifecycle state: "running", "finished", "failed", or "cancelled".
|
||||
@@ -0,0 +1,33 @@
|
||||
[**@lancedb/lancedb**](../README.md) • **Docs**
|
||||
|
||||
***
|
||||
|
||||
[@lancedb/lancedb](../globals.md) / JobFailureInfo
|
||||
|
||||
# Interface: JobFailureInfo
|
||||
|
||||
The server's account of why a job failed.
|
||||
|
||||
## Properties
|
||||
|
||||
### message?
|
||||
|
||||
```ts
|
||||
optional message: string;
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### phase?
|
||||
|
||||
```ts
|
||||
optional phase: string;
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### retryable?
|
||||
|
||||
```ts
|
||||
optional retryable: boolean;
|
||||
```
|
||||
@@ -0,0 +1,58 @@
|
||||
[**@lancedb/lancedb**](../README.md) • **Docs**
|
||||
|
||||
***
|
||||
|
||||
[@lancedb/lancedb](../globals.md) / JobInfo
|
||||
|
||||
# Interface: JobInfo
|
||||
|
||||
A row from `Connection.listJobs`: one server-side job.
|
||||
|
||||
## Properties
|
||||
|
||||
### createdAtMillis
|
||||
|
||||
```ts
|
||||
createdAtMillis: number;
|
||||
```
|
||||
|
||||
When the job was created, in milliseconds since the epoch.
|
||||
|
||||
***
|
||||
|
||||
### jobId
|
||||
|
||||
```ts
|
||||
jobId: string;
|
||||
```
|
||||
|
||||
The job id -- what `Connection.getJob` and `Connection.cancelJob`
|
||||
accept.
|
||||
|
||||
***
|
||||
|
||||
### jobType
|
||||
|
||||
```ts
|
||||
jobType: string;
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### state
|
||||
|
||||
```ts
|
||||
state: string;
|
||||
```
|
||||
|
||||
Lifecycle state: "running", "finished", "failed", or "cancelled".
|
||||
|
||||
***
|
||||
|
||||
### table
|
||||
|
||||
```ts
|
||||
table: string;
|
||||
```
|
||||
|
||||
The table the job runs against, without URI or namespace.
|
||||
@@ -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.
|
||||
|
||||
+114
-49
@@ -26,6 +26,18 @@ is also an [asynchronous API client](#connections-asynchronous).
|
||||
|
||||
::: lancedb.db.DBConnection
|
||||
|
||||
::: lancedb.Session
|
||||
|
||||
## Namespaces (Synchronous)
|
||||
|
||||
A namespace-backed connection resolves tables through a
|
||||
[Lance namespace](https://lance-format.github.io/lance-namespace/) service instead of
|
||||
listing a storage directory.
|
||||
|
||||
::: lancedb.connect_namespace
|
||||
|
||||
::: lancedb.namespace.LanceNamespaceDBConnection
|
||||
|
||||
## Tables (Synchronous)
|
||||
|
||||
::: lancedb.table.Table
|
||||
@@ -34,8 +46,12 @@ is also an [asynchronous API client](#connections-asynchronous).
|
||||
|
||||
::: lancedb.table.FragmentSummaryStats
|
||||
|
||||
::: lancedb.table.TableStatistics
|
||||
|
||||
::: lancedb.table.Tags
|
||||
|
||||
::: lancedb.table.Branches
|
||||
|
||||
## Expressions
|
||||
|
||||
Type-safe expression builder for filters and projections. Use these instead
|
||||
@@ -62,29 +78,46 @@ of raw SQL strings with [where][lancedb.query.LanceQueryBuilder.where] and
|
||||
|
||||
::: lancedb.query.LanceHybridQueryBuilder
|
||||
|
||||
::: lancedb.query.LanceEmptyQueryBuilder
|
||||
|
||||
::: lancedb.query.LanceTakeQueryBuilder
|
||||
|
||||
## Full text queries
|
||||
|
||||
Structured full text queries can be passed to
|
||||
[Table.search][lancedb.table.Table.search] or
|
||||
[AsyncTable.search][lancedb.table.AsyncTable.search] in place of a query string,
|
||||
and combined with [BooleanQuery][lancedb.query.BooleanQuery].
|
||||
|
||||
::: lancedb.query.FullTextQuery
|
||||
|
||||
::: lancedb.query.MatchQuery
|
||||
|
||||
::: lancedb.query.PhraseQuery
|
||||
|
||||
::: lancedb.query.BoostQuery
|
||||
|
||||
::: lancedb.query.MultiMatchQuery
|
||||
|
||||
::: lancedb.query.BooleanQuery
|
||||
|
||||
::: lancedb.query.FullTextOperator
|
||||
|
||||
::: lancedb.query.Occur
|
||||
|
||||
## Embeddings
|
||||
|
||||
::: lancedb.embeddings.registry.EmbeddingFunctionRegistry
|
||||
|
||||
::: lancedb.embeddings.base.EmbeddingFunctionConfig
|
||||
|
||||
::: lancedb.embeddings.base.EmbeddingFunction
|
||||
|
||||
::: lancedb.embeddings.base.TextEmbeddingFunction
|
||||
|
||||
::: lancedb.embeddings.sentence_transformers.SentenceTransformerEmbeddings
|
||||
|
||||
::: lancedb.embeddings.openai.OpenAIEmbeddings
|
||||
|
||||
::: lancedb.embeddings.open_clip.OpenClipEmbeddings
|
||||
::: lancedb.embeddings
|
||||
options:
|
||||
show_root_heading: false
|
||||
show_root_toc_entry: false
|
||||
|
||||
## Remote configuration
|
||||
|
||||
::: lancedb.remote.ClientConfig
|
||||
|
||||
::: lancedb.remote.TimeoutConfig
|
||||
|
||||
::: lancedb.remote.RetryConfig
|
||||
::: lancedb.remote
|
||||
options:
|
||||
show_root_heading: false
|
||||
show_root_toc_entry: false
|
||||
|
||||
## Context
|
||||
|
||||
@@ -122,7 +155,22 @@ tokens = list(lancedb.tokenize("acme makes searchable data",
|
||||
custom_stop_words=["acme"]))
|
||||
```
|
||||
|
||||
::: lancedb.index.FTS
|
||||
::: lancedb.tokenize
|
||||
|
||||
::: lancedb.FtsToken
|
||||
|
||||
## Blobs
|
||||
|
||||
Blob columns store large binary values out of line so they can be read lazily
|
||||
instead of being materialized with the rest of the row.
|
||||
|
||||
::: lancedb.blob
|
||||
|
||||
::: lancedb.BlobType
|
||||
|
||||
::: lancedb._blob.BlobFile
|
||||
options:
|
||||
show_root_full_path: false
|
||||
|
||||
## Utilities
|
||||
|
||||
@@ -130,6 +178,14 @@ tokens = list(lancedb.tokenize("acme makes searchable data",
|
||||
|
||||
::: lancedb.merge.LanceMergeInsertBuilder
|
||||
|
||||
::: lancedb.otel.instrument_lancedb_metrics
|
||||
|
||||
## Exceptions
|
||||
|
||||
::: lancedb.exceptions.MissingValueError
|
||||
|
||||
::: lancedb.exceptions.MissingColumnError
|
||||
|
||||
## Integrations
|
||||
|
||||
## Pydantic
|
||||
@@ -138,19 +194,30 @@ tokens = list(lancedb.tokenize("acme makes searchable data",
|
||||
|
||||
::: lancedb.pydantic.vector
|
||||
|
||||
::: lancedb.pydantic.Vector
|
||||
|
||||
::: lancedb.pydantic.MultiVector
|
||||
|
||||
::: lancedb.pydantic.LanceModel
|
||||
|
||||
## PyTorch
|
||||
|
||||
::: lancedb.streaming.StreamingDataset
|
||||
|
||||
::: lancedb.permutation.permutation_builder
|
||||
|
||||
::: lancedb.permutation.PermutationBuilder
|
||||
|
||||
::: lancedb.permutation.Permutation
|
||||
|
||||
::: lancedb.permutation.Transforms
|
||||
|
||||
## Reranking
|
||||
|
||||
::: lancedb.rerankers.linear_combination.LinearCombinationReranker
|
||||
|
||||
::: lancedb.rerankers.cohere.CohereReranker
|
||||
|
||||
::: lancedb.rerankers.colbert.ColbertReranker
|
||||
|
||||
::: lancedb.rerankers.cross_encoder.CrossEncoderReranker
|
||||
|
||||
::: lancedb.rerankers.openai.OpenaiReranker
|
||||
::: lancedb.rerankers
|
||||
options:
|
||||
show_root_heading: false
|
||||
show_root_toc_entry: false
|
||||
|
||||
## Connections (Asynchronous)
|
||||
|
||||
@@ -161,6 +228,12 @@ can be used to create, list, or open tables.
|
||||
|
||||
::: lancedb.db.AsyncConnection
|
||||
|
||||
## Namespaces (Asynchronous)
|
||||
|
||||
::: lancedb.connect_namespace_async
|
||||
|
||||
::: lancedb.namespace.AsyncLanceNamespaceDBConnection
|
||||
|
||||
## Tables (Asynchronous)
|
||||
|
||||
Table hold your actual data as a collection of records / rows.
|
||||
@@ -169,32 +242,20 @@ Table hold your actual data as a collection of records / rows.
|
||||
|
||||
::: lancedb.table.AsyncTags
|
||||
|
||||
::: lancedb.table.AsyncBranches
|
||||
|
||||
## Indices (Asynchronous)
|
||||
|
||||
Indices can be created on a table to speed up queries. This section
|
||||
lists the indices that LanceDb supports.
|
||||
|
||||
::: lancedb.index.BTree
|
||||
|
||||
::: lancedb.index.Bitmap
|
||||
|
||||
::: lancedb.index.LabelList
|
||||
|
||||
::: lancedb.index.FTS
|
||||
|
||||
::: lancedb.index.IvfPq
|
||||
|
||||
::: lancedb.index.HnswPq
|
||||
|
||||
::: lancedb.index.HnswSq
|
||||
|
||||
::: lancedb.index.IvfFlat
|
||||
|
||||
::: lancedb.index.IvfSq
|
||||
|
||||
::: lancedb.index.IvfRq
|
||||
|
||||
::: lancedb.index.HnswFlat
|
||||
::: lancedb.index
|
||||
options:
|
||||
show_root_heading: false
|
||||
show_root_toc_entry: false
|
||||
# `lang_mapping` is defined in the module rather than imported, so it is
|
||||
# picked up despite not being in `__all__`. It is an internal lookup table.
|
||||
filters: ["!^_", "!^lang_mapping$"]
|
||||
|
||||
::: lancedb.table.IndexStatistics
|
||||
|
||||
@@ -222,3 +283,7 @@ rows nearest to a query vector and can be created with the
|
||||
::: lancedb.query.AsyncHybridQuery
|
||||
options:
|
||||
inherited_members: true
|
||||
|
||||
::: lancedb.query.AsyncTakeQuery
|
||||
options:
|
||||
inherited_members: true
|
||||
|
||||
@@ -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.0.0-beta.5</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,7 +1,7 @@
|
||||
# Contributing to LanceDB Typescript
|
||||
|
||||
This document outlines the process for contributing to LanceDB Typescript.
|
||||
For general contribution guidelines, see [CONTRIBUTING.md](../CONTRIBUTING.md).
|
||||
For general contribution guidelines, see [CONTRIBUTING.md](https://github.com/lancedb/lancedb/blob/main/CONTRIBUTING.md).
|
||||
|
||||
## Project layout
|
||||
|
||||
|
||||
+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) => {
|
||||
@@ -877,3 +909,96 @@ describe("remote connection", () => {
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("remote connection jobs surface", () => {
|
||||
it("lists, describes, cancels, and reads history", async () => {
|
||||
const { tableFromArrays, tableToIPC } = await import("apache-arrow");
|
||||
const eventsTable = tableFromArrays({ state: ["created", "succeeded"] });
|
||||
const eventsBody = Buffer.from(tableToIPC(eventsTable, "stream"));
|
||||
|
||||
await withMockDatabase(
|
||||
(req, res) => {
|
||||
let body = "";
|
||||
req.on("data", (chunk) => {
|
||||
body += chunk;
|
||||
});
|
||||
req.on("end", () => {
|
||||
const payload = body.length > 0 ? JSON.parse(body) : {};
|
||||
if (req.url === "/v1/jobs/list") {
|
||||
if (payload["page_token"] === undefined) {
|
||||
res
|
||||
.writeHead(200, { "Content-Type": "application/json" })
|
||||
.end(
|
||||
'{"jobs": [{"job_id": "job-1", "table": "t1", ' +
|
||||
'"job_type": "create_index", "state": "in_progress", ' +
|
||||
'"created_at_millis": 1000}], "page_token": "next"}',
|
||||
);
|
||||
} else {
|
||||
res
|
||||
.writeHead(200, { "Content-Type": "application/json" })
|
||||
.end(
|
||||
'{"jobs": [{"job_id": "job-2", "table": "t2", ' +
|
||||
'"job_type": "create_index", "state": "succeeded", ' +
|
||||
'"created_at_millis": 2000}]}',
|
||||
);
|
||||
}
|
||||
} else if (req.url === "/v1/jobs/describe") {
|
||||
if (payload["job_id"] !== "job-1") {
|
||||
res.writeHead(404).end("no such job");
|
||||
return;
|
||||
}
|
||||
res
|
||||
.writeHead(200, { "Content-Type": "application/json" })
|
||||
.end(
|
||||
'{"job_id": "job-1", "job_type": "create_index", ' +
|
||||
'"job_state": "FAILED", "creation_ms": 1000, ' +
|
||||
'"spec": {"column": "vec"}, "failure": {"phase": "execute", ' +
|
||||
'"message": "worker died", "retryable": true}}',
|
||||
);
|
||||
} else if (req.url === "/v1/jobs/cancel") {
|
||||
if (payload["job_id"] !== "job-1") {
|
||||
res.writeHead(404).end("no such job");
|
||||
return;
|
||||
}
|
||||
res
|
||||
.writeHead(200, { "Content-Type": "application/json" })
|
||||
.end('{"job_id": "job-1"}');
|
||||
} else if (req.url === "/v1/jobs/query_events") {
|
||||
res
|
||||
.writeHead(200, {
|
||||
"Content-Type": "application/vnd.apache.arrow.stream",
|
||||
})
|
||||
.end(eventsBody);
|
||||
} else {
|
||||
res.writeHead(404).end();
|
||||
}
|
||||
});
|
||||
},
|
||||
async (db) => {
|
||||
const jobs = await db.listJobs();
|
||||
expect(jobs.map((job) => job.jobId)).toEqual(["job-1", "job-2"]);
|
||||
expect(jobs[0].state).toEqual("running");
|
||||
expect(jobs[1].state).toEqual("finished");
|
||||
|
||||
const description = await db.getJob("job-1");
|
||||
expect(description?.state).toEqual("failed");
|
||||
expect(JSON.parse(description?.specJson ?? "")).toEqual({
|
||||
column: "vec",
|
||||
});
|
||||
expect(description?.failure?.message).toEqual("worker died");
|
||||
expect(await db.getJob("missing")).toBeNull();
|
||||
|
||||
expect(await db.cancelJob("job-1")).toBe(true);
|
||||
expect(await db.cancelJob("missing")).toBe(false);
|
||||
|
||||
const history = await db.jobHistory("job-1");
|
||||
expect(history.numRows).toEqual(2);
|
||||
|
||||
const job = db.job("job-1");
|
||||
expect(job.id).toEqual("job-1");
|
||||
expect(await job.status()).toEqual("failed");
|
||||
await expect(job.wait()).rejects.toThrow("worker died");
|
||||
},
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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 () => {
|
||||
@@ -851,7 +897,11 @@ describe("When creating an index", () => {
|
||||
afterEach(() => tmpDir.removeCallback());
|
||||
|
||||
it("should create a vector index on vector columns", async () => {
|
||||
await tbl.createIndex("vec");
|
||||
const job = await tbl.createIndexAsync("vec");
|
||||
expect(job.id).toBeNull();
|
||||
await job.wait();
|
||||
// Cancelling a job that already finished succeeds and does nothing.
|
||||
await job.cancel();
|
||||
|
||||
// check index directory
|
||||
const indexDir = path.join(tmpDir.name, "test.lance", "_indices");
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
import { tableFromIPC } from "apache-arrow";
|
||||
import {
|
||||
Data,
|
||||
SchemaLike,
|
||||
@@ -20,6 +21,9 @@ import type {
|
||||
CreateNamespaceResponse,
|
||||
DescribeNamespaceResponse,
|
||||
DropNamespaceResponse,
|
||||
Job,
|
||||
JobDescription,
|
||||
JobInfo,
|
||||
ListNamespacesResponse,
|
||||
} from "./native";
|
||||
export type {
|
||||
@@ -436,6 +440,40 @@ export abstract class Connection {
|
||||
newName: string,
|
||||
options?: RenameTableOptions,
|
||||
): Promise<void>;
|
||||
|
||||
/**
|
||||
* A {@link Job} handle for a server-side job by id.
|
||||
*
|
||||
* The handle is constructed without a server round trip; an unknown id
|
||||
* surfaces when the handle is used. Dropping the handle has no effect on
|
||||
* the job itself.
|
||||
*/
|
||||
abstract job(jobId: string): Job;
|
||||
|
||||
/** List server-side jobs across the database's tables. */
|
||||
abstract listJobs(): Promise<JobInfo[]>;
|
||||
|
||||
/**
|
||||
* Describe a single server-side job by id.
|
||||
*
|
||||
* Resolves to `null` when the server has no such job.
|
||||
*/
|
||||
abstract getJob(jobId: string): Promise<JobDescription | null>;
|
||||
|
||||
/**
|
||||
* Request cancellation of a server-side job by id.
|
||||
*
|
||||
* Resolves to true if the server accepted the cancellation, false if no
|
||||
* such job exists. Cancelling an already-terminal job is a no-op success.
|
||||
*/
|
||||
abstract cancelJob(jobId: string): Promise<boolean>;
|
||||
|
||||
/**
|
||||
* The lifecycle event history of a server-side job, as an Arrow table.
|
||||
*
|
||||
* Lists history across all jobs when `jobId` is omitted.
|
||||
*/
|
||||
abstract jobHistory(jobId?: string): Promise<ArrowTable>;
|
||||
}
|
||||
|
||||
/** @hideconstructor */
|
||||
@@ -722,6 +760,30 @@ export class LocalConnection extends Connection {
|
||||
options?.newNamespacePath,
|
||||
);
|
||||
}
|
||||
|
||||
job(jobId: string): Job {
|
||||
return this.inner.job(jobId);
|
||||
}
|
||||
|
||||
async listJobs(): Promise<JobInfo[]> {
|
||||
return this.inner.listJobs();
|
||||
}
|
||||
|
||||
async getJob(jobId: string): Promise<JobDescription | null> {
|
||||
return this.inner.getJob(jobId);
|
||||
}
|
||||
|
||||
async cancelJob(jobId: string): Promise<boolean> {
|
||||
return this.inner.cancelJob(jobId);
|
||||
}
|
||||
|
||||
async jobHistory(jobId?: string): Promise<ArrowTable> {
|
||||
const buf = await this.inner.jobHistory(jobId);
|
||||
if (buf.length === 0) {
|
||||
return new ArrowTable();
|
||||
}
|
||||
return tableFromIPC(buf);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -85,7 +85,13 @@ export {
|
||||
RenameTableOptions,
|
||||
} from "./connection";
|
||||
|
||||
export { Session } from "./native.js";
|
||||
export {
|
||||
Job,
|
||||
JobDescription,
|
||||
JobFailureInfo,
|
||||
JobInfo,
|
||||
Session,
|
||||
} from "./native.js";
|
||||
|
||||
export {
|
||||
ExecutableQuery,
|
||||
|
||||
+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 = {
|
||||
|
||||
+42
-4
@@ -30,6 +30,7 @@ import {
|
||||
DropColumnsResult,
|
||||
IndexConfig,
|
||||
IndexStatistics,
|
||||
Job,
|
||||
Branches as NativeBranches,
|
||||
OptimizeStats,
|
||||
TableStatistics,
|
||||
@@ -196,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>;
|
||||
@@ -358,6 +363,17 @@ export abstract class Table {
|
||||
options?: Partial<IndexOptions>,
|
||||
): Promise<void>;
|
||||
|
||||
/**
|
||||
* Create an index, returning a handle to the indexing job.
|
||||
*
|
||||
* The job may already be complete when returned; callers must not assume
|
||||
* the index exists until {@link Job.wait} resolves.
|
||||
*/
|
||||
abstract createIndexAsync(
|
||||
column: string,
|
||||
options?: Partial<IndexOptions>,
|
||||
): Promise<Job>;
|
||||
|
||||
/**
|
||||
* Drop an index from the table.
|
||||
*
|
||||
@@ -583,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
|
||||
@@ -610,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>;
|
||||
@@ -940,6 +962,22 @@ export class LocalTable extends Table {
|
||||
);
|
||||
}
|
||||
|
||||
async createIndexAsync(
|
||||
column: string,
|
||||
options?: Partial<IndexOptions>,
|
||||
): Promise<Job> {
|
||||
// biome-ignore lint/suspicious/noExplicitAny: skip
|
||||
const nativeIndex = (options?.config as any)?.inner;
|
||||
return await this.inner.createIndexAsync(
|
||||
nativeIndex,
|
||||
column,
|
||||
options?.replace,
|
||||
options?.waitTimeoutSeconds,
|
||||
options?.name,
|
||||
options?.train,
|
||||
);
|
||||
}
|
||||
|
||||
async dropIndex(name: string): Promise<void> {
|
||||
await this.inner.dropIndex(name);
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -340,6 +340,69 @@ impl Connection {
|
||||
self.get_inner()?.drop_all_tables(&ns).await.default_error()
|
||||
}
|
||||
|
||||
/// A `Job` handle for a server-side job by id.
|
||||
///
|
||||
/// The handle is constructed without a server round trip; an unknown id
|
||||
/// surfaces when the handle is used.
|
||||
#[napi]
|
||||
pub fn job(&self, job_id: String) -> napi::Result<crate::job::Job> {
|
||||
let job = self.get_inner()?.job(job_id).default_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
}
|
||||
|
||||
/// List server-side jobs across the database's tables.
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn list_jobs(&self) -> napi::Result<Vec<crate::job::JobInfo>> {
|
||||
let jobs = self.get_inner()?.list_jobs().await.default_error()?;
|
||||
Ok(jobs.into_iter().map(Into::into).collect())
|
||||
}
|
||||
|
||||
/// Describe a single server-side job by id. `null` when the server has
|
||||
/// no such job.
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn get_job(
|
||||
&self,
|
||||
job_id: String,
|
||||
) -> napi::Result<Option<crate::job::JobDescription>> {
|
||||
let description = self.get_inner()?.get_job(&job_id).await.default_error()?;
|
||||
Ok(description.map(Into::into))
|
||||
}
|
||||
|
||||
/// Request cancellation of a server-side job by id. Returns true if the
|
||||
/// server accepted the cancellation, false if no such job exists.
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn cancel_job(&self, job_id: String) -> napi::Result<bool> {
|
||||
self.get_inner()?.cancel_job(&job_id).await.default_error()
|
||||
}
|
||||
|
||||
/// The lifecycle event history of a server-side job (all jobs when
|
||||
/// `job_id` is null), as an Arrow IPC stream buffer. Empty when there is
|
||||
/// no history.
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn job_history(&self, job_id: Option<String>) -> napi::Result<Buffer> {
|
||||
let batches = self
|
||||
.get_inner()?
|
||||
.job_history(job_id.as_deref())
|
||||
.await
|
||||
.default_error()?;
|
||||
let Some(first) = batches.first() else {
|
||||
return Ok(Buffer::from(Vec::<u8>::new()));
|
||||
};
|
||||
let mut out = Vec::new();
|
||||
let mut writer = arrow_ipc::writer::StreamWriter::try_new(&mut out, &first.schema())
|
||||
.map_err(|e| napi::Error::from_reason(e.to_string()))?;
|
||||
for batch in &batches {
|
||||
writer
|
||||
.write(batch)
|
||||
.map_err(|e| napi::Error::from_reason(e.to_string()))?;
|
||||
}
|
||||
writer
|
||||
.finish()
|
||||
.map_err(|e| napi::Error::from_reason(e.to_string()))?;
|
||||
drop(writer);
|
||||
Ok(Buffer::from(out))
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
/// Describe a namespace and return its properties.
|
||||
pub async fn describe_namespace(
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use napi_derive::napi;
|
||||
|
||||
use crate::error::NapiErrorExt;
|
||||
|
||||
/// A handle to an operation that may still be running.
|
||||
#[napi]
|
||||
pub struct Job {
|
||||
inner: Arc<lancedb::Job>,
|
||||
}
|
||||
|
||||
impl Job {
|
||||
pub(crate) fn new(inner: lancedb::Job) -> Self {
|
||||
Self {
|
||||
inner: Arc::new(inner),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[napi]
|
||||
impl Job {
|
||||
/// Identifies the operation on the server that is running it. Operations
|
||||
/// that run in this process have no server id. The value is opaque.
|
||||
#[napi(getter)]
|
||||
pub fn id(&self) -> Option<String> {
|
||||
self.inner.id().map(str::to_string)
|
||||
}
|
||||
|
||||
/// The operation's current lifecycle state: "running", "finished",
|
||||
/// "failed", or "cancelled".
|
||||
///
|
||||
/// A point snapshot; unlike {@link Job.wait} it does not block or reject
|
||||
/// on a terminal failure state. States a newer server reports that this
|
||||
/// client version does not know pass through as-is.
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn status(&self) -> napi::Result<String> {
|
||||
self.inner.status().await.default_error()
|
||||
}
|
||||
|
||||
/// Wait until the operation reaches a terminal state.
|
||||
///
|
||||
/// Jobs that complete without a resource result resolve successfully.
|
||||
/// Resource results are not exposed on this binding yet; unsupported
|
||||
/// success results reject with a generic error.
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn wait(&self) -> napi::Result<()> {
|
||||
match self.inner.wait().await.default_error()? {
|
||||
lancedb::JobResult::None => Ok(()),
|
||||
// JobResult is non_exhaustive; Function and future variants fail closed.
|
||||
_ => Err(napi::Error::from_reason(
|
||||
"unsupported job result".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Request cancellation. Cancelling a finished operation is a no-op.
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn cancel(&self) -> napi::Result<()> {
|
||||
self.inner.cancel().await.default_error()
|
||||
}
|
||||
}
|
||||
|
||||
/// A row from `Connection.listJobs`: one server-side job.
|
||||
#[napi(object)]
|
||||
pub struct JobInfo {
|
||||
/// The job id -- what `Connection.getJob` and `Connection.cancelJob`
|
||||
/// accept.
|
||||
pub job_id: String,
|
||||
/// The table the job runs against, without URI or namespace.
|
||||
pub table: String,
|
||||
pub job_type: String,
|
||||
/// Lifecycle state: "running", "finished", "failed", or "cancelled".
|
||||
pub state: String,
|
||||
/// When the job was created, in milliseconds since the epoch.
|
||||
pub created_at_millis: i64,
|
||||
}
|
||||
|
||||
impl From<lancedb::database::JobInfo> for JobInfo {
|
||||
fn from(info: lancedb::database::JobInfo) -> Self {
|
||||
Self {
|
||||
job_id: info.job_id,
|
||||
table: info.table,
|
||||
job_type: info.job_type,
|
||||
state: info.state,
|
||||
created_at_millis: info.created_at_millis,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// The server's account of why a job failed.
|
||||
#[napi(object)]
|
||||
pub struct JobFailureInfo {
|
||||
pub phase: Option<String>,
|
||||
pub message: Option<String>,
|
||||
pub retryable: Option<bool>,
|
||||
}
|
||||
|
||||
/// A described job from `Connection.getJob`.
|
||||
#[napi(object)]
|
||||
pub struct JobDescription {
|
||||
pub job_id: String,
|
||||
pub job_type: String,
|
||||
/// Lifecycle state: "running", "finished", "failed", or "cancelled".
|
||||
pub state: String,
|
||||
/// When the job was created, in milliseconds since the epoch.
|
||||
pub creation_ms: i64,
|
||||
/// The job-type-specific specification as a JSON string, when present.
|
||||
pub spec_json: Option<String>,
|
||||
/// Why the job failed, when the job is failed and the server reports a
|
||||
/// reason.
|
||||
pub failure: Option<JobFailureInfo>,
|
||||
}
|
||||
|
||||
impl From<lancedb::database::JobDescription> for JobDescription {
|
||||
fn from(description: lancedb::database::JobDescription) -> Self {
|
||||
Self {
|
||||
job_id: description.job_id,
|
||||
job_type: description.job_type,
|
||||
state: description.state,
|
||||
creation_ms: description.creation_ms,
|
||||
spec_json: (!description.spec.is_null()).then(|| description.spec.to_string()),
|
||||
failure: description.failure.map(|failure| JobFailureInfo {
|
||||
phase: failure.phase,
|
||||
message: failure.message,
|
||||
retryable: failure.retryable,
|
||||
}),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -11,6 +11,7 @@ mod error;
|
||||
mod header;
|
||||
mod index;
|
||||
mod iterator;
|
||||
mod job;
|
||||
pub mod merge;
|
||||
pub mod otel;
|
||||
pub mod permutation;
|
||||
|
||||
+49
-9
@@ -168,6 +168,39 @@ impl Table {
|
||||
builder.execute().await.default_error()
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn create_index_async(
|
||||
&self,
|
||||
index: Option<&Index>,
|
||||
column: String,
|
||||
replace: Option<bool>,
|
||||
wait_timeout_s: Option<i64>,
|
||||
name: Option<String>,
|
||||
train: Option<bool>,
|
||||
) -> napi::Result<crate::job::Job> {
|
||||
let lancedb_index = if let Some(index) = index {
|
||||
index.consume()?
|
||||
} else {
|
||||
lancedb::index::Index::Auto
|
||||
};
|
||||
let mut builder = self.inner_ref()?.create_index(&[column], lancedb_index);
|
||||
if let Some(replace) = replace {
|
||||
builder = builder.replace(replace);
|
||||
}
|
||||
if let Some(timeout) = wait_timeout_s {
|
||||
builder =
|
||||
builder.wait_timeout(std::time::Duration::from_secs(timeout.try_into().unwrap()));
|
||||
}
|
||||
if let Some(name) = name {
|
||||
builder = builder.name(name);
|
||||
}
|
||||
if let Some(train) = train {
|
||||
builder = builder.train(train);
|
||||
}
|
||||
let job = builder.execute_async().await.default_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn drop_index(&self, index_name: String) -> napi::Result<()> {
|
||||
self.inner_ref()?
|
||||
@@ -306,7 +339,9 @@ impl Table {
|
||||
let transforms = NewColumnTransform::SqlExpressions(transforms);
|
||||
let res = self
|
||||
.inner_ref()?
|
||||
.add_columns(transforms, None)
|
||||
.add_columns()
|
||||
.transform(transforms)
|
||||
.execute()
|
||||
.await
|
||||
.default_error()?;
|
||||
Ok(res.into())
|
||||
@@ -323,7 +358,9 @@ impl Table {
|
||||
let transforms = NewColumnTransform::AllNulls(schema);
|
||||
let res = self
|
||||
.inner_ref()?
|
||||
.add_columns(transforms, None)
|
||||
.add_columns()
|
||||
.transform(transforms)
|
||||
.execute()
|
||||
.await
|
||||
.default_error()?;
|
||||
Ok(res.into())
|
||||
@@ -735,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>>,
|
||||
@@ -745,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" => {
|
||||
@@ -772,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))
|
||||
}
|
||||
}
|
||||
@@ -790,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 {
|
||||
@@ -801,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 {
|
||||
@@ -811,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),
|
||||
},
|
||||
}
|
||||
@@ -1006,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
|
||||
@@ -20,8 +21,10 @@ from .remote import ClientConfig
|
||||
from .remote.db import RemoteDBConnection
|
||||
from .expr import Expr, col, lit, func
|
||||
from .schema import blob, vector, BlobType
|
||||
from .job import AsyncJob, Job
|
||||
from .table import AsyncTable, Table
|
||||
from .types import BaseTokenizerType
|
||||
from ._udf import FunctionCapability, udf
|
||||
from ._lancedb import Session
|
||||
from .namespace import (
|
||||
connect_namespace,
|
||||
@@ -500,11 +503,14 @@ __all__ = [
|
||||
"connect_namespace",
|
||||
"connect_namespace_async",
|
||||
"AsyncConnection",
|
||||
"AsyncJob",
|
||||
"AsyncLanceNamespaceDBConnection",
|
||||
"AsyncTable",
|
||||
"FtsToken",
|
||||
"col",
|
||||
"Expr",
|
||||
"Function",
|
||||
"FunctionCapability",
|
||||
"func",
|
||||
"lit",
|
||||
"URI",
|
||||
@@ -513,10 +519,12 @@ __all__ = [
|
||||
"BlobType",
|
||||
"vector",
|
||||
"DBConnection",
|
||||
"Job",
|
||||
"LanceDBConnection",
|
||||
"LanceNamespaceDBConnection",
|
||||
"RemoteDBConnection",
|
||||
"Session",
|
||||
"Table",
|
||||
"udf",
|
||||
"__version__",
|
||||
]
|
||||
|
||||
@@ -14,14 +14,10 @@ import pyarrow as pa
|
||||
from .expr import Expr
|
||||
from .schema import blob_v2_column_paths
|
||||
from .types import BlobMode, QueryProjection, QueryProjectionSpec
|
||||
from .util import get_uri_scheme
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from _typeshed import WriteableBuffer
|
||||
|
||||
from .remote.table import RemoteTable
|
||||
from .table import AsyncTable, Table
|
||||
|
||||
BLOB_MODE_TO_HANDLING = {
|
||||
"lazy": "blobs_descriptions",
|
||||
"bytes": "all_binary",
|
||||
@@ -104,22 +100,6 @@ def validate_blob_mode(blob_mode: BlobMode) -> None:
|
||||
raise ValueError(f"blob_mode must be one of {modes}, got {blob_mode!r}")
|
||||
|
||||
|
||||
def supports_blob_auto_row_id(table: Table | AsyncTable | RemoteTable) -> bool:
|
||||
"""Blob auto row-id applies to native tables, not LanceDB Cloud."""
|
||||
from .remote.table import RemoteTable
|
||||
|
||||
if isinstance(table, RemoteTable):
|
||||
return False
|
||||
|
||||
inner = getattr(table, "_inner", None)
|
||||
if inner is not None:
|
||||
uri = inner.database().uri
|
||||
if isinstance(uri, str) and get_uri_scheme(uri) == "db":
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def projection_includes_blob_column(
|
||||
projection: QueryProjection,
|
||||
blob_columns: Iterable[str],
|
||||
@@ -164,16 +144,14 @@ def v2_projection_needs_row_id(
|
||||
|
||||
|
||||
def blob_auto_row_id_for_scan(
|
||||
table: Table | AsyncTable | RemoteTable,
|
||||
schema: pa.Schema,
|
||||
projection: QueryProjection,
|
||||
*,
|
||||
with_row_id: bool | None,
|
||||
) -> bool:
|
||||
"""Auto row-id only applies when the caller said nothing about row ids."""
|
||||
if with_row_id is not None:
|
||||
return False
|
||||
if not supports_blob_auto_row_id(table):
|
||||
return False
|
||||
return v2_projection_needs_row_id(schema, projection, with_row_id=False)
|
||||
|
||||
|
||||
@@ -186,6 +164,11 @@ def finalize_blob_query_table(
|
||||
) -> pa.Table:
|
||||
if user_requested_row_id or not blob_auto_row_id:
|
||||
return tbl
|
||||
if "_rowid" not in tbl.column_names:
|
||||
# A backend that ignores the row-id request leaves nothing to stash. Hand
|
||||
# back the projection as-is so fetch_blobs raises the error that names the
|
||||
# ways to supply row ids, rather than failing here about a hidden column.
|
||||
return tbl
|
||||
return stash_auto_row_ids(tbl, blob_paths)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -146,6 +146,23 @@ class Connection(object):
|
||||
start_after: Optional[str],
|
||||
limit: Optional[int],
|
||||
) -> list[str]: ... # Deprecated: Use list_tables instead
|
||||
def job(self, job_id: str) -> Job: ...
|
||||
async def list_jobs(self) -> List[JobInfo]: ...
|
||||
async def get_job(self, job_id: str) -> Optional[JobDescription]: ...
|
||||
async def cancel_job(self, job_id: str) -> bool: ...
|
||||
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,
|
||||
@@ -209,6 +226,85 @@ class BlobFile:
|
||||
def read_range(self, offset: int, length: int) -> bytes: ...
|
||||
def read_up_to(self, length: int) -> bytes: ...
|
||||
|
||||
class Function:
|
||||
@property
|
||||
def id(self) -> str: ...
|
||||
@property
|
||||
def parameters(self) -> tuple[tuple[str, pa.DataType], ...]: ...
|
||||
@property
|
||||
def output_type(self) -> pa.DataType: ...
|
||||
@property
|
||||
def output_nullable(self) -> bool: ...
|
||||
def __call__(self, **kwargs: Any) -> "_FunctionCall": ...
|
||||
|
||||
class _FunctionCall:
|
||||
"""Private unresolved Function call authoring value (FF-028)."""
|
||||
|
||||
...
|
||||
|
||||
class _FunctionDefinition:
|
||||
"""Private owner of the Rust FunctionDefinition registration input."""
|
||||
|
||||
def _to_json(self) -> str: ...
|
||||
|
||||
def _new_function_definition(
|
||||
*,
|
||||
parameters: list[tuple[str, pa.DataType]],
|
||||
output_type: pa.DataType,
|
||||
output_nullable: bool,
|
||||
module: str,
|
||||
callable_name: str,
|
||||
source: str,
|
||||
python: str,
|
||||
packages: list[str],
|
||||
capabilities: list[tuple[str, str, Optional[str]]],
|
||||
) -> _FunctionDefinition: ...
|
||||
|
||||
class Job:
|
||||
@property
|
||||
def id(self) -> Optional[str]: ...
|
||||
async def status(self) -> str: ...
|
||||
async def wait(self) -> Optional[Function]: ...
|
||||
async def cancel(self) -> None: ...
|
||||
|
||||
class JobInfo:
|
||||
@property
|
||||
def job_id(self) -> str: ...
|
||||
@property
|
||||
def table(self) -> str: ...
|
||||
@property
|
||||
def job_type(self) -> str: ...
|
||||
@property
|
||||
def state(self) -> str: ...
|
||||
@property
|
||||
def created_at_millis(self) -> int: ...
|
||||
|
||||
class JobFailureInfo:
|
||||
@property
|
||||
def phase(self) -> Optional[str]: ...
|
||||
@property
|
||||
def message(self) -> Optional[str]: ...
|
||||
@property
|
||||
def retryable(self) -> Optional[bool]: ...
|
||||
@property
|
||||
def error_code(self) -> Optional[str]: ...
|
||||
|
||||
class JobDescription:
|
||||
@property
|
||||
def job_id(self) -> str: ...
|
||||
@property
|
||||
def job_type(self) -> str: ...
|
||||
@property
|
||||
def state(self) -> str: ...
|
||||
@property
|
||||
def creation_ms(self) -> int: ...
|
||||
@property
|
||||
def spec_json(self) -> Optional[str]: ...
|
||||
@property
|
||||
def failure(self) -> Optional[JobFailureInfo]: ...
|
||||
@property
|
||||
def result(self) -> Optional[Function]: ...
|
||||
|
||||
class Table:
|
||||
def name(self) -> str: ...
|
||||
def __repr__(self) -> str: ...
|
||||
@@ -248,6 +344,38 @@ class Table:
|
||||
name: Optional[str],
|
||||
train: Optional[bool],
|
||||
): ...
|
||||
async def create_index_async(
|
||||
self,
|
||||
column: str,
|
||||
index: Union[
|
||||
IvfFlat,
|
||||
IvfSq,
|
||||
IvfPq,
|
||||
HnswPq,
|
||||
HnswSq,
|
||||
HnswFlat,
|
||||
BTree,
|
||||
Bitmap,
|
||||
LabelList,
|
||||
Fm,
|
||||
FTS,
|
||||
],
|
||||
replace: Optional[bool],
|
||||
wait_timeout: Optional[object],
|
||||
*,
|
||||
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]): ...
|
||||
@@ -285,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: ...
|
||||
@@ -579,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
|
||||
@@ -596,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,
|
||||
)
|
||||
+311
-10
@@ -45,6 +45,7 @@ from lance_namespace.errors import NamespaceNotEmptyError, TableNotFoundError
|
||||
|
||||
from . import __version__
|
||||
from ._lancedb import connect as lancedb_connect # type: ignore
|
||||
from .job import AsyncJob, Job
|
||||
from .table import (
|
||||
AsyncTable,
|
||||
LanceTable,
|
||||
@@ -62,7 +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
|
||||
@@ -178,6 +184,51 @@ class DBConnection(EnforceOverrides):
|
||||
"Namespace operations are not supported for this connection type"
|
||||
)
|
||||
|
||||
def namespace_exists(self, namespace_id: List[str]) -> bool:
|
||||
"""Check if a namespace exists.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
namespace_id: List[str]
|
||||
The namespace identifier to check.
|
||||
|
||||
Returns
|
||||
-------
|
||||
bool
|
||||
True if the namespace exists, False otherwise.
|
||||
|
||||
Raises
|
||||
------
|
||||
NotImplementedError
|
||||
If the connection type does not support namespace operations.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"Namespace operations are not supported for this connection type"
|
||||
)
|
||||
|
||||
def table_exists(self, table_id: List[str]) -> bool:
|
||||
"""Check if a table exists.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
table_id: List[str]
|
||||
The table identifier to check (full path including namespace
|
||||
segments and table name).
|
||||
|
||||
Returns
|
||||
-------
|
||||
bool
|
||||
True if the table exists, False otherwise.
|
||||
|
||||
Raises
|
||||
------
|
||||
NotImplementedError
|
||||
If the connection type does not support namespace operations.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"Namespace operations are not supported for this connection type"
|
||||
)
|
||||
|
||||
def list_tables(
|
||||
self,
|
||||
namespace_path: Optional[List[str]] = None,
|
||||
@@ -359,7 +410,7 @@ class DBConnection(EnforceOverrides):
|
||||
|
||||
Data is converted to Arrow before being written to disk. For maximum
|
||||
control over how data is saved, either provide the PyArrow schema to
|
||||
convert to or else provide a [PyArrow Table](pyarrow.Table) directly.
|
||||
convert to or else provide a [PyArrow Table][pyarrow.Table] directly.
|
||||
|
||||
>>> import pyarrow as pa
|
||||
>>> custom_schema = pa.schema([
|
||||
@@ -563,6 +614,111 @@ class DBConnection(EnforceOverrides):
|
||||
"""
|
||||
raise NotImplementedError("serialize is not supported for this connection type")
|
||||
|
||||
def job(self, job_id: str) -> Job:
|
||||
"""A [Job][lancedb.job.Job] handle for a server-side job by id.
|
||||
|
||||
The handle is constructed without a server round trip; an unknown id
|
||||
surfaces when the handle is used. Dropping the handle has no effect
|
||||
on the job itself.
|
||||
"""
|
||||
raise NotImplementedError("job is not supported for this connection type")
|
||||
|
||||
def list_jobs(self) -> List[JobInfo]:
|
||||
"""List server-side jobs across the database's tables."""
|
||||
raise NotImplementedError("list_jobs is not supported for this connection type")
|
||||
|
||||
def get_job(self, job_id: str) -> Optional[JobDescription]:
|
||||
"""Describe a single server-side job by id.
|
||||
|
||||
Returns None when the server has no such job.
|
||||
"""
|
||||
raise NotImplementedError("get_job is not supported for this connection type")
|
||||
|
||||
def cancel_job(self, job_id: str) -> bool:
|
||||
"""Request cancellation of a server-side job by id.
|
||||
|
||||
Returns True if the server accepted the cancellation, False if no
|
||||
such job exists. Cancelling an already-terminal job is a no-op
|
||||
success.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"cancel_job is not supported for this connection type"
|
||||
)
|
||||
|
||||
def job_history(self, job_id: Optional[str] = None) -> List[pa.RecordBatch]:
|
||||
"""The lifecycle event history of a server-side job, as Arrow batches.
|
||||
|
||||
Lists history across all jobs when `job_id` is None.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"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):
|
||||
"""
|
||||
@@ -620,6 +776,9 @@ class LanceDBConnection(DBConnection):
|
||||
self._namespace_client_properties = namespace_client_properties
|
||||
if _inner is not None:
|
||||
self._conn = _inner
|
||||
# Native-derived wrappers resolve this in their async reconstruction
|
||||
# path so construction never synchronously re-enters LOOP.
|
||||
self._read_consistency_interval = read_consistency_interval
|
||||
self._cached_namespace_client = None
|
||||
return
|
||||
|
||||
@@ -669,11 +828,14 @@ class LanceDBConnection(DBConnection):
|
||||
# storage_options. Also, this class really shouldn't be holding any state
|
||||
# beyond _conn.
|
||||
self._conn = AsyncConnection(LOOP.run(do_connect()))
|
||||
# Keep property access synchronous so debugger introspection cannot wait on
|
||||
# the background loop while that thread is suspended at a breakpoint.
|
||||
self._read_consistency_interval = read_consistency_interval
|
||||
self._cached_namespace_client: Optional[LanceNamespace] = None
|
||||
|
||||
@property
|
||||
def read_consistency_interval(self) -> Optional[timedelta]:
|
||||
return LOOP.run(self._conn.get_read_consistency_interval())
|
||||
return self._read_consistency_interval
|
||||
|
||||
@property
|
||||
def session(self) -> Optional[Session]:
|
||||
@@ -684,15 +846,19 @@ class LanceDBConnection(DBConnection):
|
||||
return self._conn.uri
|
||||
|
||||
@classmethod
|
||||
def from_inner(cls, inner: LanceDbConnection):
|
||||
return cls(None, _inner=inner)
|
||||
def from_inner(
|
||||
cls,
|
||||
inner: LanceDbConnection,
|
||||
read_consistency_interval: Optional[timedelta],
|
||||
):
|
||||
return cls(
|
||||
None,
|
||||
read_consistency_interval=read_consistency_interval,
|
||||
_inner=inner,
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
val = f"{self.__class__.__name__}(uri={self._conn.uri!r}"
|
||||
if self.read_consistency_interval is not None:
|
||||
val += f", read_consistency_interval={repr(self.read_consistency_interval)}"
|
||||
val += ")"
|
||||
return val
|
||||
return f"{self.__class__.__name__}(uri={self._conn.uri!r})"
|
||||
|
||||
@override
|
||||
def serialize(self) -> str:
|
||||
@@ -1129,6 +1295,75 @@ class LanceDBConnection(DBConnection):
|
||||
)
|
||||
)
|
||||
|
||||
@override
|
||||
def job(self, job_id: str) -> Job:
|
||||
"""A [Job][lancedb.job.Job] handle for a server-side job by id.
|
||||
|
||||
The handle is constructed without a server round trip; an unknown id
|
||||
surfaces when the handle is used. Dropping the handle has no effect
|
||||
on the job itself.
|
||||
"""
|
||||
return Job(self._conn.job(job_id))
|
||||
|
||||
@override
|
||||
def list_jobs(self) -> List[JobInfo]:
|
||||
"""List server-side jobs across the database's tables."""
|
||||
return LOOP.run(self._conn.list_jobs())
|
||||
|
||||
@override
|
||||
def get_job(self, job_id: str) -> Optional[JobDescription]:
|
||||
"""Describe a single server-side job by id.
|
||||
|
||||
Returns None when the server has no such job.
|
||||
"""
|
||||
return LOOP.run(self._conn.get_job(job_id))
|
||||
|
||||
@override
|
||||
def cancel_job(self, job_id: str) -> bool:
|
||||
"""Request cancellation of a server-side job by id.
|
||||
|
||||
Returns True if the server accepted the cancellation, False if no
|
||||
such job exists. Cancelling an already-terminal job is a no-op
|
||||
success.
|
||||
"""
|
||||
return LOOP.run(self._conn.cancel_job(job_id))
|
||||
|
||||
@override
|
||||
def job_history(self, job_id: Optional[str] = None) -> List[pa.RecordBatch]:
|
||||
"""The lifecycle event history of a server-side job, as Arrow batches.
|
||||
|
||||
Lists history across all jobs when `job_id` is None.
|
||||
"""
|
||||
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.
|
||||
@@ -1529,7 +1764,7 @@ class AsyncConnection(object):
|
||||
|
||||
Data is converted to Arrow before being written to disk. For maximum
|
||||
control over how data is saved, either provide the PyArrow schema to
|
||||
convert to or else provide a [PyArrow Table](pyarrow.Table) directly.
|
||||
convert to or else provide a [PyArrow Table][pyarrow.Table] directly.
|
||||
|
||||
>>> import pyarrow as pa
|
||||
>>> custom_schema = pa.schema([
|
||||
@@ -1838,6 +2073,72 @@ class AsyncConnection(object):
|
||||
namespace_path = []
|
||||
await self._inner.drop_all_tables(namespace_path=namespace_path)
|
||||
|
||||
def job(self, job_id: str) -> AsyncJob:
|
||||
"""An [AsyncJob][lancedb.job.AsyncJob] handle for a server-side job
|
||||
by id.
|
||||
|
||||
The handle is constructed without a server round trip; an unknown id
|
||||
surfaces when the handle is used. Dropping the handle has no effect
|
||||
on the job itself.
|
||||
"""
|
||||
return AsyncJob(self._inner.job(job_id))
|
||||
|
||||
async def list_jobs(self) -> List[JobInfo]:
|
||||
"""List server-side jobs across the database's tables."""
|
||||
return await self._inner.list_jobs()
|
||||
|
||||
async def get_job(self, job_id: str) -> Optional[JobDescription]:
|
||||
"""Describe a single server-side job by id.
|
||||
|
||||
Returns None when the server has no such job.
|
||||
"""
|
||||
return await self._inner.get_job(job_id)
|
||||
|
||||
async def cancel_job(self, job_id: str) -> bool:
|
||||
"""Request cancellation of a server-side job by id.
|
||||
|
||||
Returns True if the server accepted the cancellation, False if no
|
||||
such job exists. Cancelling an already-terminal job is a no-op
|
||||
success.
|
||||
"""
|
||||
return await self._inner.cancel_job(job_id)
|
||||
|
||||
async def job_history(self, job_id: Optional[str] = None) -> List[pa.RecordBatch]:
|
||||
"""The lifecycle event history of a server-side job, as Arrow batches.
|
||||
|
||||
Lists history across all jobs when `job_id` is None.
|
||||
"""
|
||||
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.
|
||||
|
||||
|
||||
@@ -21,3 +21,32 @@ from .watsonx import WatsonxEmbeddings
|
||||
from .voyageai import VoyageAIEmbeddingFunction
|
||||
from .colpali import ColPaliEmbeddings
|
||||
from .siglip import SigLipEmbeddings
|
||||
|
||||
# The API reference renders this package with a single mkdocstrings directive,
|
||||
# which only picks up names listed here. New embedding functions must be added
|
||||
# to both the imports above and this list, or they will silently go undocumented.
|
||||
__all__ = [
|
||||
"EmbeddingFunction",
|
||||
"EmbeddingFunctionConfig",
|
||||
"TextEmbeddingFunction",
|
||||
"EmbeddingFunctionRegistry",
|
||||
"get_registry",
|
||||
"register",
|
||||
"SentenceTransformerEmbeddings",
|
||||
"OpenAIEmbeddings",
|
||||
"OpenClipEmbeddings",
|
||||
"BedRockText",
|
||||
"CohereEmbeddingFunction",
|
||||
"GeminiText",
|
||||
"GteEmbeddings",
|
||||
"InstructorEmbeddingFunction",
|
||||
"JinaEmbeddings",
|
||||
"OllamaEmbeddings",
|
||||
"TransformersEmbeddingFunction",
|
||||
"ColbertEmbeddings",
|
||||
"VoyageAIEmbeddingFunction",
|
||||
"WatsonxEmbeddings",
|
||||
"ColPaliEmbeddings",
|
||||
"ImageBindEmbeddings",
|
||||
"SigLipEmbeddings",
|
||||
]
|
||||
|
||||
@@ -21,20 +21,20 @@ class BedRockText(TextEmbeddingFunction):
|
||||
"""
|
||||
Parameters
|
||||
----------
|
||||
name: str, default "amazon.titan-embed-text-v1"
|
||||
name : str, default "amazon.titan-embed-text-v1"
|
||||
The model ID of the bedrock model to use. Supported models for are:
|
||||
- amazon.titan-embed-text-v1
|
||||
- cohere.embed-english-v3
|
||||
- cohere.embed-multilingual-v3
|
||||
region: str, default "us-east-1"
|
||||
region : str, default "us-east-1"
|
||||
Optional name of the AWS Region in which the service should be called.
|
||||
profile_name: str, default None
|
||||
profile_name : str, default None
|
||||
Optional name of the AWS profile to use for calling the Bedrock service.
|
||||
If not specified, the default profile will be used.
|
||||
assumed_role: str, default None
|
||||
assumed_role : str, default None
|
||||
Optional ARN of an AWS IAM role to assume for calling the Bedrock service.
|
||||
If not specified, the current active credentials will be used.
|
||||
role_session_name: str, default "lancedb-embeddings"
|
||||
role_session_name : str, default "lancedb-embeddings"
|
||||
Optional name of the AWS IAM role session to use for calling the Bedrock
|
||||
service. If not specified, "lancedb-embeddings" name will be used.
|
||||
|
||||
|
||||
@@ -22,7 +22,7 @@ class CohereEmbeddingFunction(TextEmbeddingFunction):
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name: str, default "embed-multilingual-v2.0"
|
||||
name : str, default "embed-multilingual-v2.0"
|
||||
The name of the model to use. List of acceptable models:
|
||||
|
||||
* embed-english-v3.0
|
||||
@@ -33,12 +33,14 @@ class CohereEmbeddingFunction(TextEmbeddingFunction):
|
||||
* embed-english-light-v2.0
|
||||
* embed-multilingual-v2.0
|
||||
|
||||
source_input_type: str, default "search_document"
|
||||
source_input_type : str, default "search_document"
|
||||
The input type for the source column in the database
|
||||
|
||||
query_input_type: str, default "search_query"
|
||||
query_input_type : str, default "search_query"
|
||||
The input type for the query column in the database
|
||||
|
||||
Notes
|
||||
-----
|
||||
Cohere supports following input types:
|
||||
|
||||
| Input Type | Description |
|
||||
|
||||
@@ -44,7 +44,7 @@ class ColPaliEmbeddings(EmbeddingFunction):
|
||||
The token pooling strategy to use, by default "hierarchical".
|
||||
- "hierarchical": Progressively pools tokens to reduce sequence length.
|
||||
- "lambda": A simpler pooling that uses a custom `pooling_func`.
|
||||
pooling_func: typing.Callable, optional
|
||||
pooling_func : typing.Callable, optional
|
||||
A function to use for pooling when `pooling_strategy` is "lambda".
|
||||
pool_factor : int
|
||||
Factor to reduce sequence length if token pooling is enabled (default 2).
|
||||
@@ -52,7 +52,7 @@ class ColPaliEmbeddings(EmbeddingFunction):
|
||||
Quantization configuration for the model. (default None, bitsandbytes needed)
|
||||
batch_size : int
|
||||
Batch size for processing inputs (default 2).
|
||||
offload_folder: str, optional
|
||||
offload_folder : str, optional
|
||||
Folder to offload model weights if using CPU offloading (default None). This is
|
||||
useful for large models that do not fit in memory.
|
||||
"""
|
||||
|
||||
@@ -48,16 +48,16 @@ class GeminiText(TextEmbeddingFunction):
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name: str, default "gemini-embedding-001"
|
||||
name : str, default "gemini-embedding-001"
|
||||
The name of the model to use. Supported models include:
|
||||
- "gemini-embedding-001" (768 dimensions)
|
||||
|
||||
Note: The legacy "models/embedding-001" format is also supported but
|
||||
"gemini-embedding-001" is recommended.
|
||||
|
||||
query_task_type: str, default "retrieval_query"
|
||||
query_task_type : str, default "retrieval_query"
|
||||
Sets the task type for the queries.
|
||||
source_task_type: str, default "retrieval_document"
|
||||
source_task_type : str, default "retrieval_document"
|
||||
Sets the task type for ingestion.
|
||||
|
||||
Examples
|
||||
|
||||
@@ -26,13 +26,13 @@ class GteEmbeddings(TextEmbeddingFunction):
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name: str, default "thenlper/gte-large"
|
||||
name : str, default "thenlper/gte-large"
|
||||
The name of the model to use.
|
||||
device: str, default "cpu"
|
||||
device : str, default "cpu"
|
||||
Sets the device type for the model.
|
||||
normalize: str, default "True"
|
||||
normalize : str, default "True"
|
||||
Controls normalize param in encode function for the transformer.
|
||||
mlx: bool, default False
|
||||
mlx : bool, default False
|
||||
Controls which model to use. False for gte-large,True for the mlx version.
|
||||
|
||||
Examples
|
||||
|
||||
@@ -35,23 +35,23 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name: str
|
||||
name : str
|
||||
The name of the model to use. Available models are listed at
|
||||
https://github.com/xlang-ai/instructor-embedding#model-list;
|
||||
The default model is hkunlp/instructor-base
|
||||
batch_size: int, default 32
|
||||
batch_size : int, default 32
|
||||
The batch size to use when generating embeddings
|
||||
device: str, default "cpu"
|
||||
device : str, default "cpu"
|
||||
The device to use when generating embeddings
|
||||
show_progress_bar: bool, default True
|
||||
show_progress_bar : bool, default True
|
||||
Whether to show a progress bar when generating embeddings
|
||||
normalize_embeddings: bool, default True
|
||||
normalize_embeddings : bool, default True
|
||||
Whether to normalize the embeddings
|
||||
quantize: bool, default False
|
||||
quantize : bool, default False
|
||||
Whether to quantize the model
|
||||
source_instruction: str, default "represent the document for retrieval"
|
||||
source_instruction : str, default "represent the document for retrieval"
|
||||
The instruction for the source column
|
||||
query_instruction: str, default "represent the document for retrieving the most
|
||||
query_instruction : str, default "represent the document for retrieving the most
|
||||
similar documents"
|
||||
The instruction for the query
|
||||
|
||||
@@ -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]])
|
||||
|
||||
@@ -40,10 +40,10 @@ class JinaEmbeddings(EmbeddingFunction):
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name: str, default "jina-clip-v1". Note that some models support both image
|
||||
name : str, default "jina-clip-v1". Note that some models support both image
|
||||
and text embeddings and some just text embedding
|
||||
|
||||
api_key: str, default None
|
||||
api_key : str, default None
|
||||
The api key to access Jina API. If you pass None, you can set JINA_API_KEY
|
||||
environment variable
|
||||
|
||||
@@ -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)))
|
||||
|
||||
@@ -21,13 +21,13 @@ class SentenceTransformerEmbeddings(TextEmbeddingFunction):
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name: str, default "all-MiniLM-L6-v2"
|
||||
name : str, default "all-MiniLM-L6-v2"
|
||||
The name of the model to use.
|
||||
device: str, default "cpu"
|
||||
device : str, default "cpu"
|
||||
The device to use for the model
|
||||
normalize: bool, default True
|
||||
normalize : bool, default True
|
||||
Whether to normalize the embeddings
|
||||
trust_remote_code: bool, default True
|
||||
trust_remote_code : bool, default True
|
||||
Whether to trust the remote code
|
||||
"""
|
||||
|
||||
|
||||
@@ -167,7 +167,7 @@ class VoyageAIEmbeddingFunction(EmbeddingFunction):
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name: str
|
||||
name : str
|
||||
The name of the model to use. List of acceptable models:
|
||||
|
||||
* voyage-4 (1024 dims, general-purpose and multilingual retrieval)
|
||||
@@ -185,7 +185,7 @@ class VoyageAIEmbeddingFunction(EmbeddingFunction):
|
||||
* voyage-law-2
|
||||
* voyage-code-2
|
||||
|
||||
output_dimension: int, optional
|
||||
output_dimension : int, optional
|
||||
The output dimension for models that support flexible dimensions.
|
||||
Currently only voyage-multimodal-3.5 supports this feature.
|
||||
Valid options: 256, 512, 1024 (default), 2048.
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
|
||||
"""Custom exception handling"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class MissingValueError(ValueError):
|
||||
"""Exception raised when a required value is missing."""
|
||||
@@ -23,3 +25,50 @@ class MissingColumnError(KeyError):
|
||||
return (
|
||||
f"Error: Column '{self.column_name}' does not exist in the DataFrame object"
|
||||
)
|
||||
|
||||
|
||||
class JobFailedError(RuntimeError):
|
||||
"""Exception raised when an asynchronous job reaches the failed state.
|
||||
|
||||
``error_code`` is the optional exact category string projected from the
|
||||
native job failure when the backend supplied one. The RuntimeError
|
||||
message remains the existing diagnostic text and must not be used to
|
||||
recover or override the code.
|
||||
"""
|
||||
|
||||
__slots__ = ("_error_code",)
|
||||
|
||||
def __init__(self, message: str, error_code: Optional[str] = None) -> None:
|
||||
super().__init__(message)
|
||||
self._error_code = error_code
|
||||
|
||||
@property
|
||||
def error_code(self) -> Optional[str]:
|
||||
"""Exact job failure error category string, when supplied."""
|
||||
return self._error_code
|
||||
|
||||
|
||||
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
|
||||
|
||||
@@ -219,7 +219,7 @@ class HnswPq:
|
||||
distance has a range of (-∞, ∞). If the vectors are normalized (i.e. their
|
||||
l2 norm is 1), then dot distance is equivalent to the cosine distance.
|
||||
|
||||
num_partitions, default sqrt(num_rows)
|
||||
num_partitions: int, default sqrt(num_rows)
|
||||
|
||||
The number of IVF partitions to create.
|
||||
|
||||
@@ -228,7 +228,7 @@ class HnswPq:
|
||||
will require too much memory. Each partition becomes its own HNSW graph, so
|
||||
setting this value higher reduces the peak memory use of training.
|
||||
|
||||
num_sub_vectors, default is vector dimension / 16
|
||||
num_sub_vectors: int, default is vector dimension / 16
|
||||
|
||||
Number of sub-vectors of PQ.
|
||||
|
||||
@@ -244,13 +244,13 @@ class HnswPq:
|
||||
If the dimension is not visible by 8 then we use 1 subvector. This is not
|
||||
ideal and will likely result in poor performance.
|
||||
|
||||
num_bits: int, default 8
|
||||
num_bits: int, default 8
|
||||
Number of bits to encode each sub-vector.
|
||||
|
||||
This value controls how much the sub-vectors are compressed. The more bits
|
||||
the more accurate the index but the slower search. Only 4 and 8 are supported.
|
||||
|
||||
max_iterations, default 50
|
||||
max_iterations: int, default 50
|
||||
|
||||
Max iterations to train kmeans.
|
||||
|
||||
@@ -263,7 +263,7 @@ class HnswPq:
|
||||
those cases it is unlikely that setting this larger will lead to the index
|
||||
converging anyways.
|
||||
|
||||
sample_rate, default 256
|
||||
sample_rate: int, default 256
|
||||
|
||||
The rate used to calculate the number of training vectors for kmeans.
|
||||
|
||||
@@ -279,14 +279,14 @@ class HnswPq:
|
||||
Increasing this value might improve the quality of the index but in
|
||||
most cases the default should be sufficient.
|
||||
|
||||
m, default 20
|
||||
m: int, default 20
|
||||
|
||||
The number of neighbors to select for each vector in the HNSW graph.
|
||||
|
||||
This value controls the tradeoff between search speed and accuracy.
|
||||
The higher the value the more accurate the search but the slower it will be.
|
||||
|
||||
ef_construction, default 300
|
||||
ef_construction: int, default 300
|
||||
|
||||
The number of candidates to evaluate during the construction of the HNSW graph.
|
||||
|
||||
@@ -297,7 +297,7 @@ class HnswPq:
|
||||
This value should be set to a value that is not less than `ef` in the
|
||||
search phase.
|
||||
|
||||
target_partition_size, default is 1,048,576
|
||||
target_partition_size: int, default is 1,048,576
|
||||
|
||||
The target size of each partition.
|
||||
|
||||
@@ -351,7 +351,7 @@ class HnswSq:
|
||||
distance has a range of (-∞, ∞). If the vectors are normalized (i.e. their
|
||||
l2 norm is 1), then dot distance is equivalent to the cosine distance.
|
||||
|
||||
num_partitions, default sqrt(num_rows)
|
||||
num_partitions: int, default sqrt(num_rows)
|
||||
|
||||
The number of IVF partitions to create.
|
||||
|
||||
@@ -360,7 +360,7 @@ class HnswSq:
|
||||
will require too much memory. Each partition becomes its own HNSW graph, so
|
||||
setting this value higher reduces the peak memory use of training.
|
||||
|
||||
max_iterations, default 50
|
||||
max_iterations: int, default 50
|
||||
|
||||
Max iterations to train kmeans.
|
||||
|
||||
@@ -373,7 +373,7 @@ class HnswSq:
|
||||
In those cases it is unlikely that setting this larger will lead to
|
||||
the index converging anyways.
|
||||
|
||||
sample_rate, default 256
|
||||
sample_rate: int, default 256
|
||||
|
||||
The rate used to calculate the number of training vectors for kmeans.
|
||||
|
||||
@@ -389,14 +389,14 @@ class HnswSq:
|
||||
Increasing this value might improve the quality of the index but in
|
||||
most cases the default should be sufficient.
|
||||
|
||||
m, default 20
|
||||
m: int, default 20
|
||||
|
||||
The number of neighbors to select for each vector in the HNSW graph.
|
||||
|
||||
This value controls the tradeoff between search speed and accuracy.
|
||||
The higher the value the more accurate the search but the slower it will be.
|
||||
|
||||
ef_construction, default 300
|
||||
ef_construction: int, default 300
|
||||
|
||||
The number of candidates to evaluate during the construction of the HNSW graph.
|
||||
|
||||
@@ -407,7 +407,7 @@ class HnswSq:
|
||||
This value should be set to a value that is not less than `ef` in the search
|
||||
phase.
|
||||
|
||||
target_partition_size, default is 1,048,576
|
||||
target_partition_size: int, default is 1,048,576
|
||||
|
||||
The target size of each partition.
|
||||
|
||||
@@ -460,7 +460,7 @@ class HnswFlat:
|
||||
distance has a range of (-∞, ∞). If the vectors are normalized (i.e. their
|
||||
l2 norm is 1), then dot distance is equivalent to the cosine distance.
|
||||
|
||||
num_partitions, default sqrt(num_rows)
|
||||
num_partitions: int, default sqrt(num_rows)
|
||||
|
||||
The number of IVF partitions to create.
|
||||
|
||||
@@ -470,18 +470,18 @@ class HnswFlat:
|
||||
graph, so setting this value higher reduces the peak memory use of
|
||||
training.
|
||||
|
||||
max_iterations, default 50
|
||||
max_iterations: int, default 50
|
||||
|
||||
Max iterations to train kmeans.
|
||||
|
||||
When training an IVF index we use kmeans to calculate the partitions.
|
||||
This parameter controls how many iterations of kmeans to run.
|
||||
|
||||
sample_rate, default 256
|
||||
sample_rate: int, default 256
|
||||
|
||||
The rate used to calculate the number of training vectors for kmeans.
|
||||
|
||||
m, default 20
|
||||
m: int, default 20
|
||||
|
||||
The number of neighbors to select for each vector in the HNSW graph.
|
||||
|
||||
@@ -489,7 +489,7 @@ class HnswFlat:
|
||||
The higher the value the more accurate the search but the slower it
|
||||
will be.
|
||||
|
||||
ef_construction, default 300
|
||||
ef_construction: int, default 300
|
||||
|
||||
The number of candidates to evaluate during the construction of the HNSW
|
||||
graph.
|
||||
@@ -501,7 +501,7 @@ class HnswFlat:
|
||||
than 500. This value should be set to a value that is not less than `ef`
|
||||
in the search phase.
|
||||
|
||||
target_partition_size, default is 1,048,576
|
||||
target_partition_size: int, default is 1,048,576
|
||||
|
||||
The target size of each partition.
|
||||
"""
|
||||
@@ -605,7 +605,7 @@ class IvfFlat:
|
||||
|
||||
The default value is 256.
|
||||
|
||||
target_partition_size, default is 8192
|
||||
target_partition_size: int, default is 8192
|
||||
|
||||
The target size of each partition.
|
||||
|
||||
@@ -769,7 +769,7 @@ class IvfPq:
|
||||
|
||||
The default value is 256.
|
||||
|
||||
target_partition_size, default is 8192
|
||||
target_partition_size: int, default is 8192
|
||||
|
||||
The target size of each partition.
|
||||
|
||||
@@ -830,7 +830,7 @@ class IvfRq:
|
||||
sample_rate: int, default 256
|
||||
Controls the number of training vectors: sample_rate * num_partitions.
|
||||
|
||||
target_partition_size, default is 8192
|
||||
target_partition_size: int, default is 8192
|
||||
Target size of each partition.
|
||||
"""
|
||||
|
||||
@@ -845,6 +845,9 @@ class IvfRq:
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
# The API reference renders this module with a single mkdocstrings directive,
|
||||
# which only picks up names listed here. New public names must be added to this
|
||||
# list, or they will silently go undocumented.
|
||||
__all__ = [
|
||||
"BTree",
|
||||
"IvfPq",
|
||||
|
||||
@@ -0,0 +1,114 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Handles to operations a server may run asynchronously."""
|
||||
|
||||
import asyncio
|
||||
from datetime import timedelta
|
||||
from typing import Optional
|
||||
|
||||
from lancedb.background_loop import LOOP
|
||||
|
||||
from . import _lancedb
|
||||
from ._lancedb import Function
|
||||
|
||||
|
||||
class AsyncJob:
|
||||
"""A handle to an operation that may still be running.
|
||||
|
||||
The operation may already be complete when the handle is created.
|
||||
"""
|
||||
|
||||
def __init__(self, inner: Optional["_lancedb.Job"]):
|
||||
self._inner = inner
|
||||
|
||||
@property
|
||||
def id(self) -> Optional[str]:
|
||||
"""Identifies the operation on the server that is running it.
|
||||
|
||||
Returned for correlating with server logs or the jobs API. Operations
|
||||
that run in this process have no server id and return `None`. The value
|
||||
is opaque: parsing it or storing it to resume the job later is not
|
||||
supported.
|
||||
"""
|
||||
return self._inner.id if self._inner is not None else None
|
||||
|
||||
async def status(self) -> str:
|
||||
"""The operation's current lifecycle state: "running", "finished",
|
||||
"failed", or "cancelled".
|
||||
|
||||
A point snapshot; unlike `wait` it does not block or raise on a
|
||||
terminal failure state. States a newer server reports that this
|
||||
client version does not know pass through as-is.
|
||||
"""
|
||||
if self._inner is None:
|
||||
return "finished"
|
||||
return await self._inner.status()
|
||||
|
||||
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 None
|
||||
if timeout is None:
|
||||
return await self._inner.wait()
|
||||
else:
|
||||
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."""
|
||||
if self._inner is None:
|
||||
return
|
||||
await self._inner.cancel()
|
||||
|
||||
|
||||
class Job:
|
||||
"""Synchronous counterpart of `AsyncJob`."""
|
||||
|
||||
def __init__(self, inner: Optional[AsyncJob]):
|
||||
self._inner = inner
|
||||
|
||||
@property
|
||||
def id(self) -> Optional[str]:
|
||||
"""Identifies the operation on the server that is running it.
|
||||
|
||||
See :attr:`AsyncJob.id`.
|
||||
"""
|
||||
return self._inner.id if self._inner is not None else None
|
||||
|
||||
def status(self) -> str:
|
||||
"""The operation's current lifecycle state: "running", "finished",
|
||||
"failed", or "cancelled".
|
||||
|
||||
See :meth:`AsyncJob.status`.
|
||||
"""
|
||||
if self._inner is None:
|
||||
return "finished"
|
||||
return LOOP.run(self._inner.status())
|
||||
|
||||
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 None
|
||||
return LOOP.run(self._inner.wait(timeout))
|
||||
|
||||
def cancel(self):
|
||||
"""Request cancellation. Cancelling a finished operation is a no-op."""
|
||||
if self._inner is None:
|
||||
return
|
||||
LOOP.run(self._inner.cancel())
|
||||
@@ -92,8 +92,10 @@ class LanceMergeInsertBuilder(object):
|
||||
self._when_not_matched_by_source_delete = True
|
||||
if isinstance(condition, Expr):
|
||||
self._when_not_matched_by_source_condition_expr = condition._inner
|
||||
elif condition is not None:
|
||||
self._when_not_matched_by_source_condition = None
|
||||
else:
|
||||
self._when_not_matched_by_source_condition = condition
|
||||
self._when_not_matched_by_source_condition_expr = None
|
||||
return self
|
||||
|
||||
def use_index(self, use_index: bool) -> LanceMergeInsertBuilder:
|
||||
|
||||
@@ -38,7 +38,11 @@ from lance_namespace_urllib3_client.models.query_table_request_vector import (
|
||||
QueryTableRequestVector,
|
||||
)
|
||||
from lance_namespace_urllib3_client.models.string_fts_query import StringFtsQuery
|
||||
from lance_namespace.errors import NamespaceNotEmptyError, TableNotFoundError
|
||||
from lance_namespace.errors import (
|
||||
NamespaceNotEmptyError,
|
||||
NamespaceNotFoundError,
|
||||
TableNotFoundError,
|
||||
)
|
||||
from lancedb._lancedb import (
|
||||
connect_namespace as _connect_namespace,
|
||||
connect_namespace_client as _connect_namespace_client,
|
||||
@@ -53,6 +57,8 @@ from lance_namespace import (
|
||||
DropNamespaceResponse,
|
||||
ListNamespacesResponse,
|
||||
ListTablesResponse,
|
||||
NamespaceExistsRequest,
|
||||
TableExistsRequest,
|
||||
)
|
||||
from lancedb.table import AsyncTable, LanceTable, Table
|
||||
from lancedb.util import validate_table_name
|
||||
@@ -780,6 +786,51 @@ class LanceNamespaceDBConnection(DBConnection):
|
||||
"""
|
||||
return LOOP.run(self._inner.describe_namespace(namespace_path))
|
||||
|
||||
@override
|
||||
def namespace_exists(self, namespace_id: List[str]) -> bool:
|
||||
"""
|
||||
Check if a namespace exists.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
namespace_id : List[str]
|
||||
The namespace identifier to check.
|
||||
|
||||
Returns
|
||||
-------
|
||||
bool
|
||||
True if the namespace exists, False otherwise.
|
||||
"""
|
||||
request = NamespaceExistsRequest(id=namespace_id)
|
||||
try:
|
||||
self._namespace_client.namespace_exists(request)
|
||||
return True
|
||||
except NamespaceNotFoundError:
|
||||
return False
|
||||
|
||||
@override
|
||||
def table_exists(self, table_id: List[str]) -> bool:
|
||||
"""
|
||||
Check if a table exists.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
table_id : List[str]
|
||||
The table identifier to check (full path including namespace
|
||||
segments and table name).
|
||||
|
||||
Returns
|
||||
-------
|
||||
bool
|
||||
True if the table exists, False otherwise.
|
||||
"""
|
||||
request = TableExistsRequest(id=table_id)
|
||||
try:
|
||||
self._namespace_client.table_exists(request)
|
||||
return True
|
||||
except TableNotFoundError:
|
||||
return False
|
||||
|
||||
@override
|
||||
def list_tables(
|
||||
self,
|
||||
@@ -1233,6 +1284,49 @@ class AsyncLanceNamespaceDBConnection:
|
||||
"""
|
||||
return await self._inner.describe_namespace(namespace_path)
|
||||
|
||||
async def namespace_exists(self, namespace_id: List[str]) -> bool:
|
||||
"""
|
||||
Check if a namespace exists.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
namespace_id : List[str]
|
||||
The namespace identifier to check.
|
||||
|
||||
Returns
|
||||
-------
|
||||
bool
|
||||
True if the namespace exists, False otherwise.
|
||||
"""
|
||||
request = NamespaceExistsRequest(id=namespace_id)
|
||||
try:
|
||||
self._namespace_client.namespace_exists(request)
|
||||
return True
|
||||
except NamespaceNotFoundError:
|
||||
return False
|
||||
|
||||
async def table_exists(self, table_id: List[str]) -> bool:
|
||||
"""
|
||||
Check if a table exists.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
table_id : List[str]
|
||||
The table identifier to check (full path including namespace
|
||||
segments and table name).
|
||||
|
||||
Returns
|
||||
-------
|
||||
bool
|
||||
True if the table exists, False otherwise.
|
||||
"""
|
||||
request = TableExistsRequest(id=table_id)
|
||||
try:
|
||||
self._namespace_client.table_exists(request)
|
||||
return True
|
||||
except TableNotFoundError:
|
||||
return False
|
||||
|
||||
async def list_tables(
|
||||
self,
|
||||
namespace_path: Optional[List[str]] = None,
|
||||
|
||||
@@ -226,7 +226,7 @@ class PermutationBuilder:
|
||||
|
||||
async def do_execute():
|
||||
inner_tbl = await self._async.execute()
|
||||
return LanceTable.from_inner(inner_tbl)
|
||||
return await LanceTable.from_inner(inner_tbl)
|
||||
|
||||
return LOOP.run(do_execute())
|
||||
|
||||
@@ -438,7 +438,8 @@ class Permutation:
|
||||
_reader: Optional[PermutationReader] = None,
|
||||
):
|
||||
"""
|
||||
Internal constructor. Use [from_tables](#from_tables) instead.
|
||||
Internal constructor. Use
|
||||
[from_tables][lancedb.permutation.Permutation.from_tables] instead.
|
||||
"""
|
||||
assert base_table is not None, "base_table is required"
|
||||
assert selection is not None, "selection is required"
|
||||
@@ -985,8 +986,9 @@ class Permutation:
|
||||
types. Conversion of strings, lists, and structs will require creating python
|
||||
objects and this is not zero-copy.
|
||||
|
||||
For custom formatting, use [with_transform](#with_transform) which overrides
|
||||
this method.
|
||||
For custom formatting, use
|
||||
[with_transform][lancedb.permutation.Permutation.with_transform] which
|
||||
overrides this method.
|
||||
"""
|
||||
assert format is not None, "format is required"
|
||||
if format == "python":
|
||||
@@ -1061,7 +1063,8 @@ class Permutation:
|
||||
Note: this method returns a new permutation and does not modify `self`
|
||||
It is provided for compatibility with the huggingface Dataset API.
|
||||
|
||||
Use [with_skip](#with_skip) instead to avoid confusion.
|
||||
Use [with_skip][lancedb.permutation.Permutation.with_skip] instead to
|
||||
avoid confusion.
|
||||
"""
|
||||
return self.with_skip(skip)
|
||||
|
||||
@@ -1084,7 +1087,8 @@ class Permutation:
|
||||
Note: this method returns a new permutation and does not modify `self`
|
||||
It is provided for compatibility with the huggingface Dataset API.
|
||||
|
||||
Use [with_take](#with_take) instead to avoid confusion.
|
||||
Use [with_take][lancedb.permutation.Permutation.with_take] instead to
|
||||
avoid confusion.
|
||||
"""
|
||||
return self.with_take(limit)
|
||||
|
||||
@@ -1107,7 +1111,8 @@ class Permutation:
|
||||
Note: this method returns a new permutation and does not modify `self`
|
||||
It is provided for compatibility with the huggingface Dataset API.
|
||||
|
||||
Use [with_repeat](#with_repeat) instead to avoid confusion.
|
||||
Use [with_repeat][lancedb.permutation.Permutation.with_repeat] instead
|
||||
to avoid confusion.
|
||||
"""
|
||||
return self.with_repeat(times)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -52,7 +52,6 @@ from ._blob import (
|
||||
finalize_blob_query_table,
|
||||
replace_v2_blob_columns_with_bytes,
|
||||
replace_v2_blob_columns_with_bytes_sync,
|
||||
supports_blob_auto_row_id,
|
||||
validate_blob_mode,
|
||||
)
|
||||
from .types import BlobMode, QueryProjection
|
||||
@@ -651,7 +650,8 @@ class Query(pydantic.BaseModel):
|
||||
distance_type : Optional[str]
|
||||
the distance type to use for vector search
|
||||
|
||||
This can be l2 (default), cosine and dot. See [metric definitions][search] for
|
||||
This can be l2 (default), cosine and dot. See
|
||||
[metric definitions](https://lancedb.com/docs/search/vector-search/) for
|
||||
more details.
|
||||
|
||||
If this is not a vector search this will be None.
|
||||
@@ -664,8 +664,9 @@ class Query(pydantic.BaseModel):
|
||||
|
||||
- A higher number makes search more accurate but also slower.
|
||||
|
||||
- See discussion in [Querying an ANN Index][querying-an-ann-index] for
|
||||
tuning advice.
|
||||
- See discussion in
|
||||
[Querying an ANN Index](https://lancedb.com/docs/indexing/)
|
||||
for tuning advice.
|
||||
|
||||
Will be None if this is not a vector search.
|
||||
refine_factor : Optional[int]
|
||||
@@ -673,8 +674,9 @@ class Query(pydantic.BaseModel):
|
||||
|
||||
- A higher number makes search more accurate but also slower.
|
||||
|
||||
- See discussion in [Querying an ANN Index][querying-an-ann-index] for
|
||||
tuning advice.
|
||||
- See discussion in
|
||||
[Querying an ANN Index](https://lancedb.com/docs/indexing/)
|
||||
for tuning advice.
|
||||
|
||||
Will be None if this is not a vector search.
|
||||
lower_bound : Optional[float]
|
||||
@@ -1277,10 +1279,7 @@ class LanceQueryBuilder(ABC):
|
||||
return self._with_row_id is True
|
||||
|
||||
def _blob_auto_row_id_enabled(self) -> bool:
|
||||
if not supports_blob_auto_row_id(self._table):
|
||||
return False
|
||||
return blob_auto_row_id_for_scan(
|
||||
self._table,
|
||||
self._table.schema,
|
||||
self._columns,
|
||||
with_row_id=self._with_row_id,
|
||||
@@ -1651,8 +1650,8 @@ class LanceVectorQueryBuilder(LanceQueryBuilder):
|
||||
Higher values will yield better recall (more likely to find vectors if
|
||||
they exist) at the expense of latency.
|
||||
|
||||
See discussion in [Querying an ANN Index][querying-an-ann-index] for
|
||||
tuning advice.
|
||||
See discussion in [Querying an ANN Index](https://lancedb.com/docs/indexing/)
|
||||
for tuning advice.
|
||||
|
||||
This method sets both the minimum and maximum number of probes to the same
|
||||
value. See `minimum_nprobes` and `maximum_nprobes` for more fine-grained
|
||||
@@ -1752,8 +1751,8 @@ class LanceVectorQueryBuilder(LanceQueryBuilder):
|
||||
As an example, a refine factor of 2 will sample 2x as many vectors as
|
||||
requested, re-ranks them, and returns the top half most relevant results.
|
||||
|
||||
See discussion in [Querying an ANN Index][querying-an-ann-index] for
|
||||
tuning advice.
|
||||
See discussion in [Querying an ANN Index](https://lancedb.com/docs/indexing/)
|
||||
for tuning advice.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
@@ -2698,7 +2697,7 @@ class LanceHybridQueryBuilder(LanceQueryBuilder):
|
||||
self._fts_query.phrase_query(True)
|
||||
if self._distance_type:
|
||||
self._vector_query.metric(self._distance_type)
|
||||
if self._minimum_nprobes:
|
||||
if self._minimum_nprobes is not None:
|
||||
self._vector_query.minimum_nprobes(self._minimum_nprobes)
|
||||
if self._maximum_nprobes is not None:
|
||||
self._vector_query.maximum_nprobes(self._maximum_nprobes)
|
||||
@@ -2771,7 +2770,7 @@ class AsyncQueryBase(object):
|
||||
)
|
||||
|
||||
async def _maybe_add_blob_row_id(self) -> None:
|
||||
if self._table is None or not supports_blob_auto_row_id(self._table):
|
||||
if self._table is None:
|
||||
self._blob_auto_row_id = False
|
||||
self._blob_paths = ()
|
||||
return
|
||||
@@ -2779,7 +2778,6 @@ class AsyncQueryBase(object):
|
||||
req = self._inner.to_query_request()
|
||||
schema = await self._table.schema()
|
||||
self._blob_auto_row_id = blob_auto_row_id_for_scan(
|
||||
self._table,
|
||||
schema,
|
||||
req.select,
|
||||
with_row_id=self._with_row_id,
|
||||
@@ -3031,7 +3029,6 @@ class AsyncQueryBase(object):
|
||||
|
||||
schema = await self._table.schema()
|
||||
blob_auto_row_id = blob_auto_row_id_for_scan(
|
||||
self._table,
|
||||
schema,
|
||||
query.columns,
|
||||
with_row_id=self._with_row_id,
|
||||
@@ -3379,8 +3376,9 @@ class AsyncQuery(AsyncStandardQuery):
|
||||
are various ANN search parameters that will let you fine tune your recall
|
||||
accuracy vs search latency.
|
||||
|
||||
Vector searches always have a [limit][]. If `limit` has not been called then
|
||||
a default `limit` of 10 will be used.
|
||||
Vector searches always have a
|
||||
[limit][lancedb.query.AsyncVectorQuery.limit]. If `limit` has not been
|
||||
called then a default `limit` of 10 will be used.
|
||||
|
||||
Typically, a single vector is passed in as the query. However, you can also
|
||||
pass in multiple vectors. When multiple vectors are passed in, if the vector
|
||||
@@ -3511,8 +3509,9 @@ class AsyncFTSQuery(AsyncStandardQuery):
|
||||
are various ANN search parameters that will let you fine tune your recall
|
||||
accuracy vs search latency.
|
||||
|
||||
Hybrid searches always have a [limit][]. If `limit` has not been called then
|
||||
a default `limit` of 10 will be used.
|
||||
Hybrid searches always have a
|
||||
[limit][lancedb.query.AsyncHybridQuery.limit]. If `limit` has not been
|
||||
called then a default `limit` of 10 will be used.
|
||||
|
||||
Typically, a single vector is passed in as the query. However, you can also
|
||||
pass in multiple vectors. This can be useful if you want to find the nearest
|
||||
@@ -3875,10 +3874,9 @@ class AsyncHybridQuery(AsyncStandardQuery, AsyncVectorQueryBase):
|
||||
req = fts_query._inner.to_query_request()
|
||||
blob_auto_row_id = False
|
||||
blob_paths: tuple[str, ...] = ()
|
||||
if self._table is not None and supports_blob_auto_row_id(self._table):
|
||||
if self._table is not None:
|
||||
schema = await self._table.schema()
|
||||
blob_auto_row_id = blob_auto_row_id_for_scan(
|
||||
self._table,
|
||||
schema,
|
||||
req.select,
|
||||
with_row_id=self._with_row_id,
|
||||
|
||||
@@ -11,6 +11,9 @@ from lancedb import __version__
|
||||
from .header import HeaderProvider
|
||||
from .oauth import OAuthConfig, OAuthFlowType
|
||||
|
||||
# The API reference renders this module with a single mkdocstrings directive,
|
||||
# which only picks up names listed here. New public names must be added to this
|
||||
# list, or they will silently go undocumented.
|
||||
__all__ = [
|
||||
"TimeoutConfig",
|
||||
"RetryConfig",
|
||||
|
||||
@@ -7,7 +7,7 @@ import json
|
||||
import logging
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
import sys
|
||||
from typing import Any, Dict, Iterable, List, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, Dict, Iterable, List, Optional, Union
|
||||
from urllib.parse import urlparse
|
||||
import warnings
|
||||
|
||||
@@ -23,6 +23,12 @@ import pyarrow as pa
|
||||
|
||||
from ..common import DATA
|
||||
from ..db import DBConnection, LOOP
|
||||
from ..job import Job
|
||||
|
||||
if TYPE_CHECKING:
|
||||
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,
|
||||
@@ -415,6 +421,11 @@ class RemoteDBConnection(DBConnection):
|
||||
|
||||
if namespace_path is None:
|
||||
namespace_path = []
|
||||
if storage_options is not None:
|
||||
logging.info(
|
||||
"storage_options is ignored in LanceDb Cloud"
|
||||
" (storage is managed; set storage_options on connect() instead)"
|
||||
)
|
||||
if index_cache_size is not None:
|
||||
logging.info(
|
||||
"index_cache_size is ignored in LanceDb Cloud"
|
||||
@@ -684,6 +695,75 @@ class RemoteDBConnection(DBConnection):
|
||||
)
|
||||
)
|
||||
|
||||
@override
|
||||
def job(self, job_id: str) -> Job:
|
||||
"""A [Job][lancedb.job.Job] handle for a server-side job by id.
|
||||
|
||||
The handle is constructed without a server round trip; an unknown id
|
||||
surfaces when the handle is used. Dropping the handle has no effect
|
||||
on the job itself.
|
||||
"""
|
||||
return Job(self._conn.job(job_id))
|
||||
|
||||
@override
|
||||
def list_jobs(self) -> List["JobInfo"]:
|
||||
"""List server-side jobs across the database's tables."""
|
||||
return LOOP.run(self._conn.list_jobs())
|
||||
|
||||
@override
|
||||
def get_job(self, job_id: str) -> Optional["JobDescription"]:
|
||||
"""Describe a single server-side job by id.
|
||||
|
||||
Returns None when the server has no such job.
|
||||
"""
|
||||
return LOOP.run(self._conn.get_job(job_id))
|
||||
|
||||
@override
|
||||
def cancel_job(self, job_id: str) -> bool:
|
||||
"""Request cancellation of a server-side job by id.
|
||||
|
||||
Returns True if the server accepted the cancellation, False if no
|
||||
such job exists. Cancelling an already-terminal job is a no-op
|
||||
success.
|
||||
"""
|
||||
return LOOP.run(self._conn.cancel_job(job_id))
|
||||
|
||||
@override
|
||||
def job_history(self, job_id: Optional[str] = None) -> List[pa.RecordBatch]:
|
||||
"""The lifecycle event history of a server-side job, as Arrow batches.
|
||||
|
||||
Lists history across all jobs when `job_id` is None.
|
||||
"""
|
||||
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.
|
||||
|
||||
@@ -53,9 +53,9 @@ class RetryError(LanceDBClientError):
|
||||
"""An error that occurs when the client has exceeded the maximum number of retries.
|
||||
|
||||
The retry strategy can be adjusted by setting the
|
||||
[retry_config](lancedb.remote.ClientConfig.retry_config) in the client
|
||||
[retry_config][lancedb.remote.ClientConfig.retry_config] in the client
|
||||
configuration. This is passed in the `client_config` argument of
|
||||
[connect](lancedb.connect) and [connect_async](lancedb.connect_async).
|
||||
[connect][lancedb.connect] and [connect_async][lancedb.connect_async].
|
||||
|
||||
The __cause__ attribute of this exception will be the last exception that
|
||||
caused the retry to fail. It will be an
|
||||
|
||||
@@ -7,6 +7,7 @@ import logging
|
||||
from functools import cached_property
|
||||
import os
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
@@ -20,6 +21,7 @@ from typing import (
|
||||
import warnings
|
||||
|
||||
from lancedb import __version__
|
||||
from lancedb._blob import BlobFile
|
||||
|
||||
from lancedb._lancedb import (
|
||||
AddColumnsResult,
|
||||
@@ -47,6 +49,7 @@ from lancedb.index import (
|
||||
IvfSq,
|
||||
LabelList,
|
||||
)
|
||||
from lancedb.job import Job
|
||||
from lancedb.remote.db import LOOP
|
||||
from lancedb.table import IndexConfigType, KNOWN_METRICS
|
||||
import pyarrow as pa
|
||||
@@ -65,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__(
|
||||
@@ -540,6 +546,73 @@ class RemoteTable(Table):
|
||||
)
|
||||
)
|
||||
|
||||
def create_index_async(
|
||||
self,
|
||||
column: str,
|
||||
*,
|
||||
config: IndexConfigType,
|
||||
replace: Optional[bool] = None,
|
||||
wait_timeout: Optional[timedelta] = None,
|
||||
name: Optional[str] = None,
|
||||
train: bool = True,
|
||||
) -> Job:
|
||||
"""Create an index, returning a handle to the indexing job.
|
||||
|
||||
The job may already be complete when returned; callers must not assume
|
||||
the index exists until :meth:`Job.wait` returns.
|
||||
"""
|
||||
return Job(
|
||||
LOOP.run(
|
||||
self._table.create_index_async(
|
||||
column,
|
||||
replace=replace,
|
||||
config=config,
|
||||
wait_timeout=wait_timeout,
|
||||
name=name,
|
||||
train=train,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
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,
|
||||
@@ -580,8 +653,9 @@ class RemoteTable(Table):
|
||||
progress: Optional[Union[bool, Callable, Any]] = None,
|
||||
write_parallelism: Optional[int] = None,
|
||||
) -> AddResult:
|
||||
"""Add more data to the [Table](Table). It has the same API signature as
|
||||
the OSS version.
|
||||
"""Add more data to the [Table][lancedb.table.Table].
|
||||
|
||||
It has the same API signature as the OSS version.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
@@ -641,7 +715,8 @@ class RemoteTable(Table):
|
||||
fast_search: bool = False,
|
||||
) -> LanceVectorQueryBuilder:
|
||||
"""Create a search query to find the nearest neighbors
|
||||
of the given query vector. We currently support [vector search][search]
|
||||
of the given query vector. We currently support
|
||||
[vector search](https://lancedb.com/docs/search/vector-search/)
|
||||
|
||||
All query options are defined in
|
||||
[LanceVectorQueryBuilder][lancedb.query.LanceVectorQueryBuilder].
|
||||
@@ -1037,22 +1112,22 @@ class RemoteTable(Table):
|
||||
)
|
||||
|
||||
def blob_columns(self) -> list[str]:
|
||||
raise NotImplementedError(
|
||||
"blob_columns() is not yet supported on the LanceDB Cloud"
|
||||
)
|
||||
return LOOP.run(self._table.blob_columns())
|
||||
|
||||
def fetch_blobs(self, column: str, row_ids) -> pa.LargeBinaryArray:
|
||||
raise NotImplementedError("fetch_blobs() is not supported on LanceDB Cloud")
|
||||
def fetch_blobs(
|
||||
self, column: str, row_ids: Union[list[int], pa.Table]
|
||||
) -> pa.LargeBinaryArray:
|
||||
return LOOP.run(self._table.fetch_blobs(column, row_ids))
|
||||
|
||||
def fetch_blob_ranges(self, column: str, requests) -> pa.LargeBinaryArray:
|
||||
raise NotImplementedError(
|
||||
"fetch_blob_ranges() is not supported on LanceDB Cloud"
|
||||
)
|
||||
|
||||
def fetch_blob_files(self, column: str, row_ids):
|
||||
raise NotImplementedError(
|
||||
"fetch_blob_files() is not supported on LanceDB Cloud"
|
||||
)
|
||||
def fetch_blob_files(
|
||||
self, column: str, row_ids: Union[list[int], pa.Table]
|
||||
) -> "list[Optional[BlobFile]]":
|
||||
return LOOP.run(self._table.fetch_blob_files(column, row_ids))
|
||||
|
||||
def head(self, n=5) -> pa.Table:
|
||||
"""
|
||||
|
||||
@@ -14,6 +14,9 @@ from .answerdotai import AnswerdotaiRerankers
|
||||
from .voyageai import VoyageAIReranker
|
||||
from .watsonx import WatsonxReranker
|
||||
|
||||
# The API reference renders this module with a single mkdocstrings directive,
|
||||
# which only picks up names listed here. New public names must be added to this
|
||||
# list, or they will silently go undocumented.
|
||||
__all__ = [
|
||||
"Reranker",
|
||||
"CrossEncoderReranker",
|
||||
|
||||
@@ -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
|
||||
|
||||
+362
-39
@@ -40,6 +40,7 @@ from ._blob import (
|
||||
from .types import BlobMode
|
||||
from lancedb.arrow import peek_reader
|
||||
from lancedb.background_loop import LOOP, embedding_executor
|
||||
from lancedb.job import AsyncJob, Job
|
||||
from .dependencies import (
|
||||
_check_for_hugging_face,
|
||||
_check_for_lance,
|
||||
@@ -107,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",
|
||||
@@ -179,6 +185,7 @@ if TYPE_CHECKING:
|
||||
LsmWriteSpec,
|
||||
MergeResult,
|
||||
UpdateResult,
|
||||
_FunctionCall,
|
||||
)
|
||||
from .index import IndexConfig
|
||||
import pandas
|
||||
@@ -863,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
|
||||
|
||||
@@ -977,6 +990,61 @@ class Table(ABC):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def create_index_async(
|
||||
self,
|
||||
column: str,
|
||||
*,
|
||||
config: IndexConfigType,
|
||||
replace: Optional[bool] = None,
|
||||
wait_timeout: Optional[timedelta] = None,
|
||||
name: Optional[str] = None,
|
||||
train: bool = True,
|
||||
) -> Job:
|
||||
"""Create an index, returning a handle to the indexing job.
|
||||
|
||||
Takes the same arguments as :meth:`create_index`. The job may already
|
||||
be complete when returned; callers must not assume the index exists
|
||||
until :meth:`Job.wait` returns.
|
||||
"""
|
||||
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.
|
||||
@@ -1211,7 +1279,7 @@ class Table(ABC):
|
||||
progress: Optional[Union[bool, Callable, Any]] = None,
|
||||
write_parallelism: Optional[int] = None,
|
||||
) -> AddResult:
|
||||
"""Add more data to the [Table](Table).
|
||||
"""Add more data to the [Table][lancedb.table.Table].
|
||||
|
||||
Parameters
|
||||
----------
|
||||
@@ -1343,8 +1411,8 @@ class Table(ABC):
|
||||
fts_columns: Optional[Union[str, List[str]]] = None,
|
||||
) -> LanceQueryBuilder:
|
||||
"""Create a search query to find the nearest neighbors
|
||||
of the given query vector. We currently support [vector search][search]
|
||||
and [full-text search][experimental-full-text-search].
|
||||
of the given query vector. We currently support [vector search](https://lancedb.com/docs/search/vector-search/)
|
||||
and [full-text search](https://lancedb.com/docs/search/full-text-search/).
|
||||
|
||||
All query options are defined in
|
||||
[LanceQueryBuilder][lancedb.query.LanceQueryBuilder].
|
||||
@@ -1574,8 +1642,10 @@ class Table(ABC):
|
||||
"""Open lazy, seekable :class:`~lancedb._blob.BlobFile` handles.
|
||||
|
||||
Prefer this over :meth:`fetch_blobs` for large payloads. ``row_ids`` is
|
||||
a ``list[int]`` or query ``pyarrow.Table`` with ``_rowid`` (or stashed
|
||||
row-id metadata). Null rows are ``None``. Local tables only.
|
||||
a ``list[int]`` or a query ``pyarrow.Table`` carrying row identity via
|
||||
``_rowid`` or a ``_lance_row_id`` field on the blob descriptor. Null
|
||||
rows are ``None``. Remote tables require LanceDB Cloud server 0.5.0 or
|
||||
newer.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
@@ -1778,7 +1848,7 @@ class Table(ABC):
|
||||
for faster reads.
|
||||
|
||||
Arguments are passed onto Lance's
|
||||
[compact_files][lance.dataset.DatasetOptimizer.compact_files].
|
||||
`lance.dataset.DatasetOptimizer.compact_files`.
|
||||
For most cases, the default should be fine.
|
||||
|
||||
See Also
|
||||
@@ -1832,6 +1902,8 @@ class Table(ABC):
|
||||
retrain: bool, default False
|
||||
This parameter is no longer used and is deprecated.
|
||||
|
||||
Notes
|
||||
-----
|
||||
The frequency an application should call optimize is based on the frequency of
|
||||
data modifications. If data is frequently added, deleted, or updated then
|
||||
optimize should be run frequently. A good rule of thumb is to run optimize if
|
||||
@@ -1986,15 +2058,14 @@ class Table(ABC):
|
||||
change permanent you can use the `[Self::restore]` method.
|
||||
|
||||
Any operation that modifies the table will fail while the table is in a checked
|
||||
out state.
|
||||
out state. To return the table to a normal state use
|
||||
`[Self::checkout_latest]`.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
version: int | str,
|
||||
The version to check out. A version number (`int`) or a tag
|
||||
(`str`) can be provided.
|
||||
|
||||
To return the table to a normal state use `[Self::checkout_latest]`
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
@@ -2160,11 +2231,15 @@ class LanceTable(Table):
|
||||
return self.name
|
||||
|
||||
@classmethod
|
||||
def from_inner(cls, tbl: LanceDBTable):
|
||||
from .db import LanceDBConnection
|
||||
async def from_inner(cls, tbl: LanceDBTable):
|
||||
from .db import AsyncConnection, LanceDBConnection
|
||||
|
||||
async_tbl = AsyncTable(tbl)
|
||||
conn = LanceDBConnection.from_inner(tbl.database())
|
||||
inner_conn = tbl.database()
|
||||
read_consistency_interval = await AsyncConnection(
|
||||
inner_conn
|
||||
).get_read_consistency_interval()
|
||||
conn = LanceDBConnection.from_inner(inner_conn, read_consistency_interval)
|
||||
return cls(
|
||||
conn,
|
||||
async_tbl.name,
|
||||
@@ -2468,13 +2543,7 @@ class LanceTable(Table):
|
||||
return LOOP.run(self._table.count_rows(filter))
|
||||
|
||||
def __repr__(self) -> str:
|
||||
val = f"{self.__class__.__name__}(name={self.name!r}"
|
||||
if self._conn.read_consistency_interval is not None:
|
||||
val += ", read_consistency_interval={!r}".format(
|
||||
self._conn.read_consistency_interval
|
||||
)
|
||||
val += f", _conn={self._conn!r})"
|
||||
return val
|
||||
return f"{self.__class__.__name__}(name={self.name!r}, _conn={self._conn!r})"
|
||||
|
||||
def __str__(self) -> str:
|
||||
return self.__repr__()
|
||||
@@ -2549,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
|
||||
-------
|
||||
@@ -2557,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
|
||||
@@ -2783,6 +2859,71 @@ class LanceTable(Table):
|
||||
)
|
||||
)
|
||||
|
||||
def create_index_async(
|
||||
self,
|
||||
column: str,
|
||||
*,
|
||||
config: IndexConfigType,
|
||||
replace: Optional[bool] = None,
|
||||
wait_timeout: Optional[timedelta] = None,
|
||||
name: Optional[str] = None,
|
||||
train: bool = True,
|
||||
) -> Job:
|
||||
"""Create an index, returning a handle to the indexing job.
|
||||
|
||||
The job may already be complete when returned; callers must not assume
|
||||
the index exists until :meth:`Job.wait` returns.
|
||||
"""
|
||||
return Job(
|
||||
LOOP.run(
|
||||
self._table.create_index_async(
|
||||
column,
|
||||
replace=replace,
|
||||
config=config,
|
||||
wait_timeout=wait_timeout,
|
||||
name=name,
|
||||
train=train,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
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,
|
||||
@@ -3387,8 +3528,8 @@ class LanceTable(Table):
|
||||
fts_columns: Optional[Union[str, List[str]]] = None,
|
||||
) -> LanceQueryBuilder:
|
||||
"""Create a search query to find the nearest neighbors
|
||||
of the given query vector. We currently support [vector search][search]
|
||||
and [full-text search][search].
|
||||
of the given query vector. We currently support [vector search](https://lancedb.com/docs/search/vector-search/)
|
||||
and [full-text search](https://lancedb.com/docs/search/full-text-search/).
|
||||
|
||||
Examples
|
||||
--------
|
||||
@@ -3418,8 +3559,9 @@ class LanceTable(Table):
|
||||
- *default None*.
|
||||
Acceptable types are: list, np.ndarray, PIL.Image.Image
|
||||
|
||||
- If None then the select/[where][sql]/limit clauses are applied
|
||||
to filter the table
|
||||
- If None then the
|
||||
select/[where][lancedb.query.LanceQueryBuilder.where]/limit clauses
|
||||
are applied to filter the table
|
||||
vector_column_name: str, optional
|
||||
The name of the vector column to search.
|
||||
|
||||
@@ -3813,6 +3955,8 @@ class LanceTable(Table):
|
||||
retrain: bool, default False
|
||||
This parameter is no longer used and is deprecated.
|
||||
|
||||
Notes
|
||||
-----
|
||||
The frequency an application should call optimize is based on the frequency of
|
||||
data modifications. If data is frequently added, deleted, or updated then
|
||||
optimize should be run frequently. A good rule of thumb is to run optimize if
|
||||
@@ -3907,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]."""
|
||||
@@ -4585,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
|
||||
@@ -4611,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.
|
||||
|
||||
@@ -4691,7 +4925,7 @@ class AsyncTable:
|
||||
Parameters
|
||||
----------
|
||||
**kwargs
|
||||
Forwarded to [`lance.dataset`][lance.dataset].
|
||||
Forwarded to `lance.dataset`.
|
||||
|
||||
Returns
|
||||
-------
|
||||
@@ -4867,6 +5101,90 @@ class AsyncTable:
|
||||
)
|
||||
raise e
|
||||
|
||||
async def create_index_async(
|
||||
self,
|
||||
column: str,
|
||||
*,
|
||||
replace: Optional[bool] = None,
|
||||
config: Optional[
|
||||
Union[
|
||||
IvfFlat,
|
||||
IvfPq,
|
||||
IvfRq,
|
||||
HnswPq,
|
||||
HnswSq,
|
||||
HnswFlat,
|
||||
BTree,
|
||||
Bitmap,
|
||||
LabelList,
|
||||
Fm,
|
||||
FTS,
|
||||
]
|
||||
] = None,
|
||||
wait_timeout: Optional[timedelta] = None,
|
||||
name: Optional[str] = None,
|
||||
train: bool = True,
|
||||
) -> AsyncJob:
|
||||
"""Create an index, returning a handle to the indexing job.
|
||||
|
||||
Takes the same arguments as :meth:`create_index`. The job may already
|
||||
be complete when returned; callers must not assume the index exists
|
||||
until :meth:`AsyncJob.wait` resolves.
|
||||
"""
|
||||
job = await self._inner.create_index_async(
|
||||
column,
|
||||
index=config,
|
||||
replace=replace,
|
||||
wait_timeout=wait_timeout,
|
||||
name=name,
|
||||
train=train,
|
||||
)
|
||||
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.
|
||||
@@ -5010,7 +5328,7 @@ class AsyncTable:
|
||||
progress: Optional[Union[bool, Callable, Any]] = None,
|
||||
write_parallelism: Optional[int] = None,
|
||||
) -> AddResult:
|
||||
"""Add more data to the [Table](Table).
|
||||
"""Add more data to the [AsyncTable][lancedb.table.AsyncTable].
|
||||
|
||||
Parameters
|
||||
----------
|
||||
@@ -5212,8 +5530,8 @@ class AsyncTable:
|
||||
fts_columns: Optional[Union[str, List[str]]] = None,
|
||||
) -> Union[AsyncHybridQuery, AsyncFTSQuery, AsyncVectorQuery]:
|
||||
"""Create a search query to find the nearest neighbors
|
||||
of the given query vector. We currently support [vector search][search]
|
||||
and [full-text search][experimental-full-text-search].
|
||||
of the given query vector. We currently support [vector search](https://lancedb.com/docs/search/vector-search/)
|
||||
and [full-text search](https://lancedb.com/docs/search/full-text-search/).
|
||||
|
||||
All query options are defined in [AsyncQuery][lancedb.query.AsyncQuery].
|
||||
|
||||
@@ -5774,15 +6092,14 @@ class AsyncTable:
|
||||
change permanent you can use the `[Self::restore]` method.
|
||||
|
||||
Any operation that modifies the table will fail while the table is in a checked
|
||||
out state.
|
||||
out state. To return the table to a normal state use
|
||||
`[Self::checkout_latest]`.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
version: int | str,
|
||||
The version to check out. A version number (`int`) or a tag
|
||||
(`str`) can be provided.
|
||||
|
||||
To return the table to a normal state use `[Self::checkout_latest]`
|
||||
"""
|
||||
try:
|
||||
await self._inner.checkout(version)
|
||||
@@ -5966,6 +6283,8 @@ class AsyncTable:
|
||||
retrain: bool, default False
|
||||
This parameter is no longer used and is deprecated.
|
||||
|
||||
Notes
|
||||
-----
|
||||
The frequency an application should call optimize is based on the frequency of
|
||||
data modifications. If data is frequently added, deleted, or updated then
|
||||
optimize should be run frequently. A good rule of thumb is to run optimize if
|
||||
@@ -6141,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
|
||||
@@ -6346,6 +6667,8 @@ class Branches:
|
||||
dry_run: bool, default False
|
||||
When True, only preview. When False, attempt the merge.
|
||||
|
||||
Notes
|
||||
-----
|
||||
A rejected merge returns ``status="rejected"`` instead of raising.
|
||||
"""
|
||||
return LOOP.run(self._table.branches.merge(from_branch, dry_run))
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -226,13 +226,13 @@ def test_fetch_blob_ranges_validates_requests():
|
||||
table = _blob_table("range_validation", [{"id": 1, "image": b"abc"}])
|
||||
row_id = _row_ids_by_id(table)[1]
|
||||
|
||||
with pytest.raises(RuntimeError, match="exceeds blob size"):
|
||||
with pytest.raises(ValueError, match="exceeds blob size"):
|
||||
table.fetch_blob_ranges("image", [(row_id, 2, 2)])
|
||||
|
||||
with pytest.raises(RuntimeError, match="offset \\+ length overflowed"):
|
||||
with pytest.raises(ValueError, match="offset \\+ length overflowed"):
|
||||
table.fetch_blob_ranges("image", [(row_id, 2**64 - 1, 1)])
|
||||
|
||||
with pytest.raises(ValueError, match="row ids"):
|
||||
with pytest.raises(ValueError, match="row IDs"):
|
||||
table.fetch_blob_ranges("image", [(2**64 - 1, 0, 1)])
|
||||
|
||||
|
||||
|
||||
@@ -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,6 +68,44 @@ def test_basic(tmp_path):
|
||||
assert db.open_table("test").name == db["test"].name
|
||||
|
||||
|
||||
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("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})"
|
||||
|
||||
|
||||
def test_read_consistency_interval_does_not_use_background_loop(tmp_path, monkeypatch):
|
||||
from lancedb.background_loop import LOOP
|
||||
from lancedb.db import LanceDBConnection
|
||||
|
||||
consistency_interval = timedelta(seconds=5)
|
||||
db = lancedb.connect(tmp_path, read_consistency_interval=consistency_interval)
|
||||
db_from_inner = LanceDBConnection.from_inner(db._inner, consistency_interval)
|
||||
|
||||
def fail_run(*args, **kwargs):
|
||||
raise AssertionError("properties should not use the Python background loop")
|
||||
|
||||
monkeypatch.setattr(LOOP, "run", fail_run)
|
||||
|
||||
assert db.read_consistency_interval == consistency_interval
|
||||
assert db_from_inner.read_consistency_interval == consistency_interval
|
||||
|
||||
|
||||
def test_ingest_pd(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
|
||||
|
||||
@@ -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
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user